From 8119f050de871e6cdb617c938ce958da7afaf1c4 Mon Sep 17 00:00:00 2001 From: hexiaonan-800 Date: Sun, 20 Sep 2026 11:06:26 +0800 Subject: [PATCH] feat: Application model validation --- .../nodes/image_generate_node/image_generate_node.py | 12 ++++++++++++ .../nodes/image_to_video_node/image_to_video_node.py | 11 +++++++++++ .../image_understand_node/image_understand_node.py | 11 +++++++++++ .../workflow/nodes/intent_node/intent_node.py | 10 ++++++++++ .../parameter_extraction_node.py | 10 ++++++++++ .../workflow/nodes/question_node/question_node.py | 10 ++++++++++ .../workflow/nodes/reranker_node/reranker_node.py | 12 ++++++++++++ .../nodes/speech_to_text_node/speech_to_text_node.py | 11 +++++++++++ .../nodes/text_to_speech_node/text_to_speech_node.py | 12 ++++++++++++ .../nodes/text_to_video_node/text_to_video_node.py | 12 ++++++++++++ .../video_understand_node/video_understand_node.py | 11 +++++++++++ 11 files changed, 122 insertions(+) diff --git a/apps/application/workflow/nodes/image_generate_node/image_generate_node.py b/apps/application/workflow/nodes/image_generate_node/image_generate_node.py index e0c148be4dd..c2776a1c50d 100644 --- a/apps/application/workflow/nodes/image_generate_node/image_generate_node.py +++ b/apps/application/workflow/nodes/image_generate_node/image_generate_node.py @@ -12,6 +12,7 @@ import requests import uuid_utils.compat as uuid +from django.db.models import QuerySet from django.utils.translation import gettext_lazy as _ from langchain_core.messages import HumanMessage, AIMessage from rest_framework import serializers @@ -21,8 +22,10 @@ from application.workflow.message.struct.content import NodeInfo, Position from application.workflow.message.struct.text_content import TextContent from application.workflow.status import Status +from common.exception.app_exception import AppApiException from common.utils.common import bytes_to_uploaded_file from knowledge.models import FileSourceType +from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id from oss.serializers.file import FileSerializer @@ -44,6 +47,15 @@ class ImageGenerateNodeSerializer(serializers.Serializer): is_result = serializers.BooleanField(required=False, label=_("Whether to return content")) model_params_setting = serializers.JSONField(required=False, default=dict, label=_("Model parameter settings")) + 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")) + class ImageGenerateNode(INode): serializer_class = ImageGenerateNodeSerializer diff --git a/apps/application/workflow/nodes/image_to_video_node/image_to_video_node.py b/apps/application/workflow/nodes/image_to_video_node/image_to_video_node.py index 93dfdef542b..624a71ab0e9 100644 --- a/apps/application/workflow/nodes/image_to_video_node/image_to_video_node.py +++ b/apps/application/workflow/nodes/image_to_video_node/image_to_video_node.py @@ -18,7 +18,9 @@ from common.utils.common import bytes_to_uploaded_file from knowledge.models import FileSourceType, File from oss.serializers.file import FileSerializer, mime_types +from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id +from common.exception.app_exception import AppApiException from common.utils.logger import maxkb_logger @@ -41,6 +43,15 @@ class ImageToVideoNodeSerializer(serializers.Serializer): first_frame_url = serializers.ListField(required=True, label=_("First frame url")) last_frame_url = serializers.ListField(required=False, label=_("Last frame url")) + 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")) + class ImageToVideoNode(INode): serializer_class = ImageToVideoNodeSerializer diff --git a/apps/application/workflow/nodes/image_understand_node/image_understand_node.py b/apps/application/workflow/nodes/image_understand_node/image_understand_node.py index f25d528c8a6..20df7753c4e 100644 --- a/apps/application/workflow/nodes/image_understand_node/image_understand_node.py +++ b/apps/application/workflow/nodes/image_understand_node/image_understand_node.py @@ -16,7 +16,9 @@ from application.workflow.status import Status from application.workflow.tools import Reasoning from common.utils.common import guess_image_format +from common.exception.app_exception import AppApiException from knowledge.models import File +from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id @@ -35,6 +37,15 @@ class ImageUnderstandNodeSerializer(serializers.Serializer): model_params_setting = serializers.JSONField(required=False, default=dict, label=_("Model parameter settings")) model_setting = serializers.DictField(required=False, label="Model settings") + 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")) + class ImageUnderstandNode(INode): serializer_class = ImageUnderstandNodeSerializer diff --git a/apps/application/workflow/nodes/intent_node/intent_node.py b/apps/application/workflow/nodes/intent_node/intent_node.py index 8ef3165c5df..bb0e884c2f6 100644 --- a/apps/application/workflow/nodes/intent_node/intent_node.py +++ b/apps/application/workflow/nodes/intent_node/intent_node.py @@ -12,6 +12,7 @@ from application.workflow.common import WorkflowType from application.workflow.i_node import INode from application.workflow.status import Status +from common.exception.app_exception import AppApiException from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id, get_model_credential from .prompt_template import PROMPT_TEMPLATE @@ -34,6 +35,15 @@ class IntentNodeSerializer(serializers.Serializer): model_params_setting = serializers.DictField(required=False, label=_("Model parameter settings")) branch = IntentBranchSerializer(many=True) + 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() diff --git a/apps/application/workflow/nodes/parameter_extraction_node/parameter_extraction_node.py b/apps/application/workflow/nodes/parameter_extraction_node/parameter_extraction_node.py index b33a30936d2..ce3ea8f63f6 100644 --- a/apps/application/workflow/nodes/parameter_extraction_node/parameter_extraction_node.py +++ b/apps/application/workflow/nodes/parameter_extraction_node/parameter_extraction_node.py @@ -18,6 +18,7 @@ from application.workflow.i_node import INode from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id, get_model_credential +from common.exception.app_exception import AppApiException prompt = """ Please strictly process the text according to the following requirements: @@ -48,6 +49,15 @@ class ParameterExtractionNodeSerializer(serializers.Serializer): required=False, child=serializers.CharField(), allow_empty=True, label=_("Reference Field") ) + 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() diff --git a/apps/application/workflow/nodes/question_node/question_node.py b/apps/application/workflow/nodes/question_node/question_node.py index edfa7645d86..7313a1a9f07 100644 --- a/apps/application/workflow/nodes/question_node/question_node.py +++ b/apps/application/workflow/nodes/question_node/question_node.py @@ -18,6 +18,7 @@ from application.workflow.message.struct.content import NodeInfo, Position from application.workflow.message.struct.text_content import TextContent from application.workflow.status import Status +from common.exception.app_exception import AppApiException from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id, get_model_credential @@ -34,6 +35,15 @@ class QuestionNodeSerializer(serializers.Serializer): is_result = serializers.BooleanField(required=False, label=_("Whether to return content")) model_params_setting = serializers.DictField(required=False, label=_("Model parameter settings")) + 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() diff --git a/apps/application/workflow/nodes/reranker_node/reranker_node.py b/apps/application/workflow/nodes/reranker_node/reranker_node.py index ad41f199d84..87802a36941 100644 --- a/apps/application/workflow/nodes/reranker_node/reranker_node.py +++ b/apps/application/workflow/nodes/reranker_node/reranker_node.py @@ -7,12 +7,15 @@ from typing import List +from django.db.models import QuerySet from django.utils.translation import gettext_lazy as _ from langchain_core.documents import Document from rest_framework import serializers from application.workflow.common import WorkflowType from application.workflow.i_node import INode +from common.exception.app_exception import AppApiException +from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id @@ -35,6 +38,15 @@ class RerankerNodeSerializer(serializers.Serializer): required=True, label=_("The results are displayed in the knowledge sources") ) + def is_valid(self, *, raise_exception=False): + super().is_valid(raise_exception=True) + # reference / default 在运行时才解析,此处只校验自定义模型 + if (self.data.get("reranker_model_id_type") or "custom") in ("reference", "default"): + return + model_id = self.data.get("reranker_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 _merge_reranker_list(reranker_list, result=None): if result is None: diff --git a/apps/application/workflow/nodes/speech_to_text_node/speech_to_text_node.py b/apps/application/workflow/nodes/speech_to_text_node/speech_to_text_node.py index f9be89ade3a..37b194d5535 100644 --- a/apps/application/workflow/nodes/speech_to_text_node/speech_to_text_node.py +++ b/apps/application/workflow/nodes/speech_to_text_node/speech_to_text_node.py @@ -19,7 +19,9 @@ from application.workflow.message.struct.text_content import TextContent from application.workflow.status import Status from common.utils.common import split_and_transcribe, any_to_mp3 +from common.exception.app_exception import AppApiException from knowledge.models import File +from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id @@ -33,6 +35,15 @@ class SpeechToTextNodeSerializer(serializers.Serializer): audio_list = serializers.ListField(required=True, label=_("The audio file cannot be empty")) model_params_setting = serializers.DictField(required=False, label=_("Model parameter settings")) + def is_valid(self, *, raise_exception=False): + super().is_valid(raise_exception=True) + # reference / default 在运行时才解析,此处只校验自定义模型 + if (self.data.get("stt_model_id_type") or "custom") in ("reference", "default"): + return + model_id = self.data.get("stt_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 _process_audio_item(audio_item, model): file = QuerySet(File).filter(id=audio_item["file_id"]).first() diff --git a/apps/application/workflow/nodes/text_to_speech_node/text_to_speech_node.py b/apps/application/workflow/nodes/text_to_speech_node/text_to_speech_node.py index 80b0a913292..79e2ad9d8cc 100644 --- a/apps/application/workflow/nodes/text_to_speech_node/text_to_speech_node.py +++ b/apps/application/workflow/nodes/text_to_speech_node/text_to_speech_node.py @@ -9,6 +9,7 @@ import mimetypes from django.core.files.uploadedfile import InMemoryUploadedFile +from django.db.models import QuerySet from django.utils.translation import gettext_lazy as _ from pydub import AudioSegment from rest_framework import serializers @@ -18,8 +19,10 @@ from application.workflow.message.struct.content import NodeInfo, Position from application.workflow.message.struct.text_content import TextContent from application.workflow.status import Status +from common.exception.app_exception import AppApiException from common.utils.common import _remove_empty_lines from knowledge.models import FileSourceType +from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id from oss.serializers.file import FileSerializer @@ -34,6 +37,15 @@ class TextToSpeechNodeSerializer(serializers.Serializer): content_list = serializers.ListField(required=True, label=_("Text content")) model_params_setting = serializers.DictField(required=False, label=_("Model parameter settings")) + def is_valid(self, *, raise_exception=False): + super().is_valid(raise_exception=True) + # reference / default 在运行时才解析,此处只校验自定义模型 + if (self.data.get("tts_model_id_type") or "custom") in ("reference", "default"): + return + model_id = self.data.get("tts_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 _bytes_to_uploaded_file(file_bytes, file_name="generated_audio.mp3"): content_type, _ = mimetypes.guess_type(file_name) diff --git a/apps/application/workflow/nodes/text_to_video_node/text_to_video_node.py b/apps/application/workflow/nodes/text_to_video_node/text_to_video_node.py index 868b4209b4a..c4ad6b989db 100644 --- a/apps/application/workflow/nodes/text_to_video_node/text_to_video_node.py +++ b/apps/application/workflow/nodes/text_to_video_node/text_to_video_node.py @@ -4,6 +4,7 @@ from functools import reduce from typing import List +from django.db.models import QuerySet from django.utils.translation import gettext_lazy as _, gettext from langchain_core.messages import BaseMessage, HumanMessage, AIMessage from rest_framework import serializers @@ -13,9 +14,11 @@ from application.workflow.message.struct.content import NodeInfo, Position from application.workflow.message.struct.text_content import TextContent from application.workflow.status import Status +from common.exception.app_exception import AppApiException from common.utils.common import bytes_to_uploaded_file from knowledge.models import FileSourceType from oss.serializers.file import FileSerializer +from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id from common.utils.logger import maxkb_logger @@ -37,6 +40,15 @@ class TextToVideoNodeSerializer(serializers.Serializer): is_result = serializers.BooleanField(required=False, label=_("Whether to return content")) model_params_setting = serializers.JSONField(required=False, default=dict, label=_("Model parameter settings")) + 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")) + class TextToVideoNode(INode): serializer_class = TextToVideoNodeSerializer diff --git a/apps/application/workflow/nodes/video_understand_node/video_understand_node.py b/apps/application/workflow/nodes/video_understand_node/video_understand_node.py index 9e7706dac2a..8f9d3b3b8ac 100644 --- a/apps/application/workflow/nodes/video_understand_node/video_understand_node.py +++ b/apps/application/workflow/nodes/video_understand_node/video_understand_node.py @@ -14,7 +14,9 @@ from application.workflow.message.struct.text_content import TextContent from application.workflow.status import Status from application.workflow.tools import Reasoning +from common.exception.app_exception import AppApiException from knowledge.models import File +from models_provider.models import Model from models_provider.tools import get_model_instance_by_model_workspace_id @@ -33,6 +35,15 @@ class VideoUnderstandNodeSerializer(serializers.Serializer): model_params_setting = serializers.JSONField(required=False, default=dict, label=_("Model parameter settings")) model_setting = serializers.DictField(required=False, label="Model settings") + 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")) + class VideoUnderstandNode(INode): serializer_class = VideoUnderstandNodeSerializer