From 44445b399fc17bf5ee08e1ec38f82c227305c997 Mon Sep 17 00:00:00 2001 From: hexiaonan-800 Date: Sun, 20 Sep 2026 09:51:40 +0800 Subject: [PATCH] feat: Optimize application workflow publish validation --- apps/application/serializers/application.py | 4 +- apps/application/workflow/common.py | 53 +++++++++++++++---- apps/application/workflow/i_node.py | 7 +-- .../nodes/ai_chat_node/ai_chat_node.py | 12 ++++- .../workflow/nodes/start_node/start_node.py | 1 - 5 files changed, 60 insertions(+), 17 deletions(-) diff --git a/apps/application/serializers/application.py b/apps/application/serializers/application.py index 6c9ef959e85..e0edf3e8a70 100644 --- a/apps/application/serializers/application.py +++ b/apps/application/serializers/application.py @@ -21,7 +21,7 @@ import requests import uuid_utils.compat as uuid -from application.workflow.common import new_instance +from application.workflow.common import new_instance, WorkflowType from application.long_term_memory import schedule_extract_long_term_memory from application.models.application import Application, ApplicationFolder, ApplicationTypeChoices, ApplicationVersion from application.models.application_access_token import ApplicationAccessToken @@ -1278,7 +1278,7 @@ def publish(self, instance, with_valid=True): work_flow = application.work_flow if work_flow is None: raise AppApiException(500, _("work_flow is a required field")) - new_instance(work_flow).is_valid() + new_instance(work_flow).is_valid(workflow_type=WorkflowType.APPLICATION) base_node = get_base_node_work_flow(work_flow) if base_node is not None: node_data = base_node.get("properties").get("node_data") diff --git a/apps/application/workflow/common.py b/apps/application/workflow/common.py index f386939b835..802b2a4ee39 100644 --- a/apps/application/workflow/common.py +++ b/apps/application/workflow/common.py @@ -10,6 +10,8 @@ from enum import Enum from typing import List, Dict +from django.utils.translation import gettext as _ +from common.exception.app_exception import AppApiException from common.utils.common import group_by @@ -106,6 +108,15 @@ def reset_variable(self, prompt: str): return prompt +class WorkflowType(Enum): + # 应用 + APPLICATION = "APPLICATION" + # 知识库 + KNOWLEDGE = "KNOWLEDGE" + # 工具 + TOOL = "TOOL" + + class Workflow: """ 节点列表 @@ -194,17 +205,39 @@ def reset_prompt(self, prompt): prompt = node_field.reset_variable(prompt) return prompt - def is_valid(self): - pass - + def is_valid(self, workflow_type: WorkflowType): + """ + 校验工作流数据:一趟遍历同时统计节点id出现次数、校验每个节点的参数 + """ + start_node_list = [] + for node in self.nodes: + if node.id == "start-node": + start_node_list.append(node) + self.is_valid_node(node, workflow_type) + self.is_valid_start_node(start_node_list) + + def is_valid_start_node(self, start_node_list: List[Node]): + """ + 校验开始节点:有且只有一个 start-node + """ + if len(start_node_list) == 0: + raise AppApiException(500, _("The starting node is required")) + if len(start_node_list) > 1: + raise AppApiException(500, _("There can only be one starting node")) -class WorkflowType(Enum): - # 应用 - APPLICATION = "APPLICATION" - # 知识库 - KNOWLEDGE = "KNOWLEDGE" - # 工具 - TOOL = "TOOL" + def is_valid_node(self, node: Node, workflow_type: WorkflowType = WorkflowType.APPLICATION): + """ + 校验单个节点:交给该节点类型对应的序列化器 + """ + from application.workflow.nodes import node_map + + node_class = node_map.get(node.type, {}).get(workflow_type) + if node_class is None or node_class.serializer_class is None: + return + try: + node_class.serializer_class(data=get_node_parameters(node)).is_valid(raise_exception=True) + except AppApiException as e: + raise AppApiException(500, f"{node.properties.get('stepName')}:{e.message}") def new_instance(flow_obj: Dict, workflow_type: WorkflowType = WorkflowType.APPLICATION): diff --git a/apps/application/workflow/i_node.py b/apps/application/workflow/i_node.py index 07b83ba0c36..0b3a17c0c94 100644 --- a/apps/application/workflow/i_node.py +++ b/apps/application/workflow/i_node.py @@ -42,9 +42,10 @@ class INode: # 序列化校验器 serializer_class: Optional[Type[serializers.Serializer]] = None - @staticmethod - def is_valid(data): - INode.serializer_class(data=data).is_valid(raise_exception=True) + @classmethod + def is_valid(cls, data): + if cls.serializer_class: + cls.serializer_class(data=data).is_valid(raise_exception=True) def __init__(self, node, workflow_manage, get_node_parameters: Callable[[Node], dict]): self.node = node diff --git a/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py b/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py index 86589127f3c..63db8463899 100644 --- a/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py +++ b/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py @@ -41,6 +41,7 @@ from knowledge.models import File from models_provider.models import Model from models_provider.tools import get_model_credential, get_model_instance_by_model_workspace_id +from common.exception.app_exception import AppApiException class AgentCallBack: @@ -98,6 +99,15 @@ class ChatNodeSerializer(serializers.Serializer): image_list = serializers.ListField(required=False, label=_("picture")) vision = serializers.BooleanField(required=False, default=False, label=_("vision")) + def is_valid(self, *, raise_exception=False): + super().is_valid(raise_exception=True) + # reference / default 在运行时才解析,此处只校验自定义模型 + if (self.data.get("model_id_type") or "custom") in ("reference", "default"): + return + model_id = self.data.get("model_id") + if not model_id or not QuerySet(Model).filter(id=model_id).exists(): + raise AppApiException(500, _("The model of the node does not exist")) + def _get_default_model_params_setting(model_id): model = QuerySet(Model).filter(id=model_id).first() @@ -639,7 +649,7 @@ def _get_reference_content(self, fields): def get_details(self, index: int = 0, position: dict = None, old_details: dict = None, **kwargs): details = super().get_details(index, position, old_details, **kwargs) aggregation = AggregationManager() - for m in self.data.get("messages"): + for m in self.data.get("messages") or []: aggregation.aggregate(m) messages = aggregation.get_contents() details.update( diff --git a/apps/application/workflow/nodes/start_node/start_node.py b/apps/application/workflow/nodes/start_node/start_node.py index cffee289ef9..e9a863aef76 100644 --- a/apps/application/workflow/nodes/start_node/start_node.py +++ b/apps/application/workflow/nodes/start_node/start_node.py @@ -36,7 +36,6 @@ class ApplicationSerializer(serializers.Serializer): class StarNode(INode): - serializer_class = ApplicationSerializer supported_workflow_type_list = [WorkflowType.APPLICATION] type = "start-node"