Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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
Expand Down
10 changes: 10 additions & 0 deletions apps/application/workflow/nodes/intent_node/intent_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand Down
10 changes: 10 additions & 0 deletions apps/application/workflow/nodes/question_node/question_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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()
Expand Down
12 changes: 12 additions & 0 deletions apps/application/workflow/nodes/reranker_node/reranker_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand All @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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
Expand Down
Loading