diff --git a/apps/application/flow/common.py b/apps/application/flow/common.py index d7520cf690c..26d3e4229fa 100644 --- a/apps/application/flow/common.py +++ b/apps/application/flow/common.py @@ -248,33 +248,56 @@ def is_valid_start_node(self): def is_valid_model_params(self): node_list = [node for node in self.nodes if ( node.type == 'ai-chat-node' or node.type == 'question-node' or node.type == 'parameter-extraction-node')] + + model_ids = [] + nodes_requiring_validation = [] for node in node_list: if (node.properties.get('node_data', {}).get('model_id_type') or 'custom') == 'reference': continue - model = QuerySet(Model).filter(id=node.properties.get('node_data', {}).get('model_id')).first() - if model is None: - raise ValidationError(ErrorDetail( - _('The node {node} model does not exist').format(node=node.properties.get("stepName")))) - credential = get_model_credential(model.provider, model.model_type, model.model_name) - model_params_setting = node.properties.get('node_data', {}).get('model_params_setting') - model_params_setting_form = credential.get_model_params_setting_form( - model.model_name) - if model_params_setting is None: - model_params_setting = model_params_setting_form.get_default_form_data() - node.properties.get('node_data', {})['model_params_setting'] = model_params_setting - if node.properties.get('status', 200) != 200: - raise ValidationError( - ErrorDetail(_("Node {node} is unavailable").format(node=node.properties.get("stepName")))) + model_id = node.properties.get('node_data', {}).get('model_id') + if model_id: + model_ids.append(model_id) + nodes_requiring_validation.append((node, model_id)) + if model_ids: + # 先批量获取数据 + models_map = {str(model.id): model for model in QuerySet(Model).filter(id__in=model_ids)} + # 再循环从map中取数据进行处理 + for node, model_id in nodes_requiring_validation: + model = models_map.get(str(model_id)) + if model is None: + raise ValidationError(ErrorDetail( + _('The node {node} model does not exist').format(node=node.properties.get("stepName")))) + credential = get_model_credential(model.provider, model.model_type, model.model_name) + model_params_setting = node.properties.get('node_data', {}).get('model_params_setting') + model_params_setting_form = credential.get_model_params_setting_form( + model.model_name) + if model_params_setting is None: + model_params_setting = model_params_setting_form.get_default_form_data() + node.properties.get('node_data', {})['model_params_setting'] = model_params_setting + if node.properties.get('status', 200) != 200: + raise ValidationError( + ErrorDetail(_("Node {node} is unavailable").format(node=node.properties.get("stepName")))) + node_list = [node for node in self.nodes if (node.type == 'function-lib-node')] + + function_lib_ids = [] + nodes_requiring_tool_validation = [] for node in node_list: function_lib_id = node.properties.get('node_data', {}).get('function_lib_id') if function_lib_id is None: raise ValidationError(ErrorDetail( _('The library ID of node {node} cannot be empty').format(node=node.properties.get("stepName")))) - f_lib = QuerySet(Tool).filter(id=function_lib_id).first() - if f_lib is None: - raise ValidationError(ErrorDetail(_("The function library for node {node} is not available").format( - node=node.properties.get("stepName")))) + function_lib_ids.append(function_lib_id) + nodes_requiring_tool_validation.append((node, function_lib_id)) + if function_lib_ids: + # 先批量获取数据 + tools_map = {str(tool.id): tool for tool in QuerySet(Tool).filter(id__in=function_lib_ids)} + # 再循环从map中取数据进行处理 + for node, function_lib_id in nodes_requiring_tool_validation: + f_lib = tools_map.get(str(function_lib_id)) + if f_lib is None: + raise ValidationError(ErrorDetail(_("The function library for node {node} is not available").format( + node=node.properties.get("stepName")))) def is_valid_base_node(self): base_node_list = [node for node in self.nodes if node.id == 'base-node'] diff --git a/apps/application/flow/step_node/ai_chat_step_node/impl/base_chat_node.py b/apps/application/flow/step_node/ai_chat_step_node/impl/base_chat_node.py index 43ee372ef3e..e3a3ca6ac10 100644 --- a/apps/application/flow/step_node/ai_chat_step_node/impl/base_chat_node.py +++ b/apps/application/flow/step_node/ai_chat_step_node/impl/base_chat_node.py @@ -60,8 +60,8 @@ def write_context_stream(node_variable: Dict, workflow_variable: Dict, node: INo @param workflow: 工作流管理器 """ response = node_variable.get("result") - answer = "" - reasoning_content = "" + answer_chunks = [] + reasoning_content_chunks = [] model_setting = node.context.get( "model_setting", {"reasoning_content_enable": False, "reasoning_content_end": "", "reasoning_content_start": ""}, @@ -81,10 +81,10 @@ def write_context_stream(node_variable: Dict, workflow_variable: Dict, node: INo reasoning_content_chunk = chunk.additional_kwargs.get("reasoning_content", "") else: reasoning_content_chunk = reasoning_chunk.get("reasoning_content") - answer += content_chunk + answer_chunks.append(content_chunk) if reasoning_content_chunk is None: reasoning_content_chunk = "" - reasoning_content += reasoning_content_chunk + reasoning_content_chunks.append(reasoning_content_chunk) yield { "content": content_chunk, "reasoning_content": reasoning_content_chunk @@ -93,7 +93,7 @@ def write_context_stream(node_variable: Dict, workflow_variable: Dict, node: INo } reasoning_chunk = reasoning.get_end_reasoning_content() - answer += reasoning_chunk.get("content") + answer_chunks.append(reasoning_chunk.get("content")) reasoning_content_chunk = "" if not response_reasoning_content: reasoning_content_chunk = reasoning_chunk.get("reasoning_content") @@ -101,6 +101,8 @@ def write_context_stream(node_variable: Dict, workflow_variable: Dict, node: INo "content": reasoning_chunk.get("content"), "reasoning_content": reasoning_content_chunk if model_setting.get("reasoning_content_enable", False) else "", } + answer = "".join(answer_chunks) + reasoning_content = "".join(reasoning_content_chunks) _write_context(node_variable, workflow_variable, node, workflow, answer, reasoning_content) @@ -525,12 +527,20 @@ def _process_videos(self, image, video_model): if isinstance(image, str) and image.startswith("http"): videos.append({"type": "video_url", "video_url": {"url": image}}) elif image is not None and len(image) > 0: + # 先批量获取数据 + file_ids = [img["file_id"] for img in image if "file_id" in img] + if file_ids: + files_map = {str(f.id): f for f in QuerySet(File).filter(id__in=file_ids)} + else: + files_map = {} + # 再循环从map中取数据进行处理 for img in image: if "file_id" in img: file_id = img["file_id"] - file = QuerySet(File).filter(id=file_id).first() - url = video_model.upload_file_and_get_url(file.get_bytes(), file.file_name) - videos.append({"type": "video_url", "video_url": {"url": url}}) + file = files_map.get(str(file_id)) + if file: + url = video_model.upload_file_and_get_url(file.get_bytes(), file.file_name) + videos.append({"type": "video_url", "video_url": {"url": url}}) elif "url" in img and img["url"].startswith("http"): videos.append({"type": "video_url", "video_url": {"url": img["url"]}}) return videos @@ -543,16 +553,24 @@ def _process_images(self, image): if isinstance(image, str) and image.startswith("http"): images.append({"type": "image_url", "image_url": {"url": image}}) elif image is not None and len(image) > 0: + # 先批量获取数据 + file_ids = [img["file_id"] for img in image if "file_id" in img] + if file_ids: + files_map = {str(f.id): f for f in QuerySet(File).filter(id__in=file_ids)} + else: + files_map = {} + # 再循环从map中取数据进行处理 for img in image: if "file_id" in img: file_id = img["file_id"] - file = QuerySet(File).filter(id=file_id).first() - image_bytes = file.get_bytes() - base64_image = base64.b64encode(image_bytes).decode("utf-8") - image_format = what(None, image_bytes) - images.append( - {"type": "image_url", "image_url": {"url": f"data:image/{image_format};base64,{base64_image}"}} - ) + file = files_map.get(str(file_id)) + if file: + image_bytes = file.get_bytes() + base64_image = base64.b64encode(image_bytes).decode("utf-8") + image_format = what(None, image_bytes) + images.append( + {"type": "image_url", "image_url": {"url": f"data:image/{image_format};base64,{base64_image}"}} + ) elif "url" in img and img["url"].startswith("http"): images.append({"type": "image_url", "image_url": {"url": img["url"]}}) return images diff --git a/apps/application/serializers/application.py b/apps/application/serializers/application.py index 5c74570911b..d13a48fc230 100644 --- a/apps/application/serializers/application.py +++ b/apps/application/serializers/application.py @@ -69,6 +69,76 @@ from application.serializers.common import update_resource_mapping_by_application +def get_bound_tool_ids(instance: Dict) -> List[str]: + """ + 收集应用配置(含工作流节点)中引用的所有工具id,用于绑定前的权限校验 + """ + tool_ids = set() + for key in ("tool_ids", "skill_tool_ids", "mcp_tool_ids"): + for tool_id in (instance.get(key) or []): + tool_ids.add(str(tool_id)) + if instance.get("mcp_tool_id"): + tool_ids.add(str(instance.get("mcp_tool_id"))) + + def walk(work_flow): + if not work_flow: + return + for node in work_flow.get("nodes", []) or []: + node_data = (node.get("properties") or {}).get("node_data") or {} + for key in ("tool_lib_id", "mcp_tool_id"): + if node_data.get(key): + tool_ids.add(str(node_data.get(key))) + for key in ("mcp_tool_ids", "tool_ids", "skill_tool_ids"): + for tool_id in (node_data.get(key) or []): + tool_ids.add(str(tool_id)) + if node.get("type") == "loop-node": + walk(node_data.get("loop_body")) + + walk(instance.get("work_flow")) + return list(tool_ids) + + +def get_authorized_tool_ids(user_id: str, workspace_id: str, tool_ids: List[str]) -> List[str]: + """ + 返回 tool_ids 中当前用户被授权绑定/使用的工具id。 + 工作空间管理员默认拥有全部工具权限;其他用户必须在 workspace_user_resource_permission + 中存在针对该工具的显式授权记录(默认拒绝)。 + """ + if not tool_ids: + return [] + tool_ids = list({str(t) for t in tool_ids}) + if is_workspace_manage(user_id, workspace_id): + return tool_ids + granted_tool_ids = { + str(permission.target) + for permission in QuerySet(WorkspaceUserResourcePermission).filter( + workspace_id=workspace_id, + user_id=user_id, + auth_target_type=AuthTargetType.TOOL.value, + target__in=tool_ids, + ) + if "VIEW" in permission.permission_list or "ROLE" in permission.permission_list + } + return [tool_id for tool_id in tool_ids if tool_id in granted_tool_ids] + + +def validate_bound_tool_permissions(user_id: str, workspace_id: str, instance: Dict): + """ + 校验应用/工作流中绑定的工具,当前用户是否都有权限使用,防止低权限成员 + 绑定自己被禁止访问的工具,并通过应用/工作流执行绕过工具的单独授权控制。 + """ + tool_ids = get_bound_tool_ids(instance) + if not tool_ids: + return + authorized_tool_ids = set(get_authorized_tool_ids(user_id, workspace_id, tool_ids)) + unauthorized_tool_ids = [tool_id for tool_id in tool_ids if tool_id not in authorized_tool_ids] + if unauthorized_tool_ids: + message = lazy_format( + _("No permission to use tool(s): {tool_ids}"), tool_ids=", ".join(unauthorized_tool_ids) + ) + raise AppApiException(403, str(message)) + + def get_base_node_work_flow(work_flow): node_list = work_flow.get("nodes") base_node_list = [node for node in node_list if node.get("id") == "base-node"] @@ -623,6 +693,7 @@ def insert_workflow(self, instance: Dict): workspace_id = self.data.get("workspace_id") wq = ApplicationCreateSerializer.WorkflowRequest(data=instance) wq.is_valid(raise_exception=True) + validate_bound_tool_permissions(user_id, workspace_id, instance) application_model = wq.to_application_model(user_id, workspace_id, instance) application_model.save() # 插入认证信息 @@ -703,6 +774,20 @@ def import_(self, instance: dict, is_import_tool, with_valid=True): if not exits_tool_id_list.__contains__(tool.get("id")) and not exits_tool_id_list.__contains__(generate_uuid((tool.get("id") + workspace_id or ""))) ] + # 导入包内新建的工具由导入者本人持有,无需校验;仅需校验绑定到已存在工具的引用 + existing_bound_tool_ids = [ + tool_id for tool_id in get_bound_tool_ids(application) if tool_id not in update_tool_map + ] + if existing_bound_tool_ids: + authorized_tool_ids = set(get_authorized_tool_ids(user_id, workspace_id, existing_bound_tool_ids)) + unauthorized_tool_ids = [ + tool_id for tool_id in existing_bound_tool_ids if tool_id not in authorized_tool_ids + ] + if unauthorized_tool_ids: + message = lazy_format( + _("No permission to use tool(s): {tool_ids}"), tool_ids=", ".join(unauthorized_tool_ids) + ) + raise AppApiException(403, str(message)) application_model = self.to_application(application, workspace_id, user_id, update_tool_map, folder_id) tool_model_list = [self.to_tool(f, workspace_id, user_id) for f in tool_list] application_model.save() @@ -1190,6 +1275,8 @@ def edit(self, instance: Dict, with_valid=True): if "work_flow_template" in instance: return self.update_template_workflow(instance, application) + validate_bound_tool_permissions(self.data.get("user_id"), self.data.get("workspace_id"), instance) + if instance.get("model_id") is None or len(instance.get("model_id")) == 0: application.model_id = None else: diff --git a/apps/chat/mcp/tools.py b/apps/chat/mcp/tools.py index 4a3ff972388..d25f0aa8edd 100644 --- a/apps/chat/mcp/tools.py +++ b/apps/chat/mcp/tools.py @@ -3,6 +3,7 @@ import uuid_utils.compat as uuid from django.db.models import QuerySet +from django.utils import timezone from application.models import ApplicationApiKey, Application, ChatUserType, ChatSourceChoices from chat.serializers.chat import ChatSerializers @@ -13,6 +14,8 @@ def __init__(self, auth_header): app_key = QuerySet(ApplicationApiKey).filter(secret_key=auth_header, is_active=True).first() if not app_key: raise PermissionError("Invalid API Key") + if app_key.is_permanent is False and app_key.expire_time < timezone.now(): + raise PermissionError("API Key is expired") self.application = QuerySet(Application).filter(id=app_key.application_id, is_publish=True).first() if not self.application: diff --git a/apps/common/utils/fork.py b/apps/common/utils/fork.py index b4c47ab1114..3f9a17c8dd6 100644 --- a/apps/common/utils/fork.py +++ b/apps/common/utils/fork.py @@ -1,5 +1,7 @@ +import base64 import copy import re +import sys import traceback from functools import reduce from typing import List, Set @@ -10,8 +12,18 @@ from bs4 import BeautifulSoup from common.utils.logger import maxkb_logger +from maxkb.const import CONFIG requests.packages.urllib3.disable_warnings() +_enable_sandbox = bool(int(CONFIG.get("SANDBOX", 0))) + + +class SandboxFetchResponse: + def __init__(self, status_code: int, content: bytes, encoding: str | None, apparent_encoding: str | None): + self.status_code = status_code + self.content = content + self.encoding = encoding + self.apparent_encoding = apparent_encoding class ChildLink: @@ -182,6 +194,41 @@ def get_charset_list(meta_list): charset_list.append(match.group(1)) return charset_list + @staticmethod + def _sandbox_requests_get(base_fork_url: str, headers: dict): + from common.utils.tool_code import ToolExecutor + + response = ToolExecutor().exec_code( + """ +def fetch_url(url, headers): + import base64 + import requests + + requests.packages.urllib3.disable_warnings() + response = requests.get(url, verify=False, headers=headers) + return { + "status_code": response.status_code, + "content": base64.b64encode(response.content).decode("ascii"), + "encoding": response.encoding, + "apparent_encoding": response.apparent_encoding, + } +""", + {"url": base_fork_url, "headers": headers}, + function_name="fetch_url", + ) + return SandboxFetchResponse( + response.get("status_code"), + base64.b64decode(response.get("content")), + response.get("encoding"), + response.get("apparent_encoding"), + ) + + @staticmethod + def requests_get(base_fork_url: str, headers: dict): + if _enable_sandbox and sys.platform.startswith("linux"): + return Fork._sandbox_requests_get(base_fork_url, headers) + return requests.get(base_fork_url, verify=False, headers=headers) + def fork(self): try: @@ -190,7 +237,7 @@ def fork(self): } maxkb_logger.info(f'fork:{self.base_fork_url}') - response = requests.get(self.base_fork_url, verify=False, headers=headers) + response = self.requests_get(self.base_fork_url, headers) if response.status_code != 200: maxkb_logger.error(f"url: {self.base_fork_url} code:{response.status_code}") return Fork.Response.error(f"url: {self.base_fork_url} code:{response.status_code}") diff --git a/apps/knowledge/serializers/document.py b/apps/knowledge/serializers/document.py index 0dda62b017d..9a069572381 100644 --- a/apps/knowledge/serializers/document.py +++ b/apps/knowledge/serializers/document.py @@ -694,7 +694,8 @@ def is_valid(self, *, raise_exception=False): if not query_set.exists(): raise AppApiException(500, _("Knowledge id does not exist")) document_id = self.data.get("document_id") - if not QuerySet(Document).filter(id=document_id).exists(): + knowledge_id = self.data.get("knowledge_id") + if not QuerySet(Document).filter(id=document_id, knowledge_id=knowledge_id).exists(): raise AppApiException(500, _("document id not exist")) def export(self, with_valid=True): diff --git a/apps/knowledge/serializers/paragraph.py b/apps/knowledge/serializers/paragraph.py index 09a30242c77..8630e8adf8c 100644 --- a/apps/knowledge/serializers/paragraph.py +++ b/apps/knowledge/serializers/paragraph.py @@ -117,7 +117,11 @@ def is_valid(self, *, raise_exception=False): query_set = query_set.filter(workspace_id=workspace_id) if not query_set.exists(): raise AppApiException(500, _("Knowledge id does not exist")) - if not QuerySet(Paragraph).filter(id=self.data.get("paragraph_id")).exists(): + if not QuerySet(Paragraph).filter( + id=self.data.get("paragraph_id"), + document_id=self.data.get("document_id"), + knowledge_id=self.data.get("knowledge_id"), + ).exists(): raise AppApiException(500, _("Paragraph id does not exist")) def list(self, with_valid=False): @@ -209,7 +213,11 @@ def is_valid(self, *, raise_exception=True): query_set = query_set.filter(workspace_id=workspace_id) if not query_set.exists(): raise AppApiException(500, _("Knowledge id does not exist")) - if not QuerySet(Paragraph).filter(id=self.data.get("paragraph_id")).exists(): + if not QuerySet(Paragraph).filter( + id=self.data.get("paragraph_id"), + document_id=self.data.get("document_id"), + knowledge_id=self.data.get("knowledge_id"), + ).exists(): raise AppApiException(500, _("Paragraph id does not exist")) @staticmethod diff --git a/apps/models_provider/serializers/model_serializer.py b/apps/models_provider/serializers/model_serializer.py index 7e866eda128..8d0cef2a3e7 100644 --- a/apps/models_provider/serializers/model_serializer.py +++ b/apps/models_provider/serializers/model_serializer.py @@ -6,26 +6,25 @@ from typing import Dict import uuid_utils.compat as uuid -from django.core.cache import cache -from django.db import transaction -from django.db.models import QuerySet -from django.utils.translation import gettext_lazy as _ -from rest_framework import serializers - from common.config.embedding_config import ModelManage from common.constants.cache_version import Cache_Version -from common.constants.permission_constants import ResourcePermission, ResourceAuthType +from common.constants.permission_constants import ResourceAuthType, ResourcePermission from common.database_model_manage.database_model_manage import DatabaseModelManage from common.db.search import native_search from common.exception.app_exception import AppApiException from common.utils.common import get_file_content -from common.utils.rsa_util import rsa_long_encrypt, rsa_long_decrypt +from common.utils.rsa_util import rsa_long_decrypt, rsa_long_encrypt +from django.core.cache import cache +from django.db import transaction +from django.db.models import QuerySet +from django.utils.translation import gettext_lazy as _ from maxkb.conf import PROJECT_DIR -from models_provider.base_model_provider import ValidCode, DownModelChunkStatus +from models_provider.base_model_provider import DownModelChunkStatus, ValidCode from models_provider.constants.model_provider_constants import ModelProvideConstants from models_provider.models import Model, Status from models_provider.tools import get_model_credential -from system_manage.models import WorkspaceUserResourcePermission, AuthTargetType +from rest_framework import serializers +from system_manage.models import AuthTargetType, WorkspaceUserResourcePermission from system_manage.models.resource_mapping import ResourceMapping from system_manage.serializers.resource_mapping_serializers import ResourceMappingSerializer from system_manage.serializers.user_resource_permission import UserResourcePermissionSerializer @@ -44,9 +43,19 @@ class ModelModelSerializer(serializers.ModelSerializer): class Meta: model = Model fields = [ - 'id', 'name', 'status', 'model_type', 'model_name', - 'user', 'provider', 'credential', 'meta', - 'model_params_form', 'workspace_id', 'create_time', 'update_time' + "id", + "name", + "status", + "model_type", + "model_name", + "user", + "provider", + "credential", + "meta", + "model_params_form", + "workspace_id", + "create_time", + "update_time", ] @@ -83,19 +92,15 @@ def pull(model: Model, credential: Dict): status = Status.ERROR message = "" for chunk in down_model_chunk.values(): - if chunk.get('status') == DownModelChunkStatus.success.value: + if chunk.get("status") == DownModelChunkStatus.success.value: status = Status.SUCCESS - elif chunk.get('status') == DownModelChunkStatus.error.value: + elif chunk.get("status") == DownModelChunkStatus.error.value: message = chunk.get("digest") - QuerySet(Model).filter(id=model.id).update( - meta={"down_model_chunk": [], "message": message}, - status=status - ) + QuerySet(Model).filter(id=model.id).update(meta={"down_model_chunk": [], "message": message}, status=status) except Exception as e: QuerySet(Model).filter(id=model.id).update( - meta={"down_model_chunk": [], "message": str(e)}, - status=Status.ERROR + meta={"down_model_chunk": [], "message": str(e)}, status=Status.ERROR ) @@ -104,19 +109,19 @@ class ModelSerializer(serializers.Serializer): def model_to_dict(model: Model): credential = json.loads(rsa_long_decrypt(model.credential)) return { - 'id': str(model.id), - 'provider': model.provider, - 'name': model.name, - 'model_type': model.model_type, - 'model_name': model.model_name, - 'status': model.status, - 'meta': model.meta, - 'credential': ModelProvideConstants[model.provider].value.get_model_credential( - model.model_type, model.model_name - ).encryption_dict(credential), - 'workspace_id': model.workspace_id, - 'nick_name': model.user.nick_name if model.user else '', - 'username': model.user.username if model.user else '' + "id": str(model.id), + "provider": model.provider, + "name": model.name, + "model_type": model.model_type, + "model_name": model.model_name, + "status": model.status, + "meta": model.meta, + "credential": ModelProvideConstants[model.provider] + .value.get_model_credential(model.model_type, model.model_name) + .encryption_dict(credential), + "workspace_id": model.workspace_id, + "nick_name": model.user.nick_name if model.user else "", + "username": model.user.username if model.user else "", } class Operate(serializers.Serializer): @@ -132,44 +137,49 @@ def is_valid(self, *, raise_exception=False): model_query = model_query.filter(workspace_id=workspace_id) model = model_query.first() if model is None: - raise AppApiException(500, _('Model does not exist')) - if model.workspace_id == 'None': - raise AppApiException(500, _('Shared models cannot be deleted or modified')) + raise AppApiException(500, _("Model does not exist")) + if model.workspace_id == "None": + raise AppApiException(500, _("Shared models cannot be deleted or modified")) def one(self, with_valid=False): if with_valid: super().is_valid(raise_exception=True) - model = QuerySet(Model).get( - id=self.data.get('id'), workspace_id=self.data.get('workspace_id', 'None') - ) + model = QuerySet(Model).get(id=self.data.get("id"), workspace_id=self.data.get("workspace_id", "None")) return ModelSerializer.model_to_dict(model) def one_meta(self, with_valid=False): model = None if with_valid: super().is_valid(raise_exception=True) - model = QuerySet(Model).filter(id=self.data.get("id"), - workspace_id=self.data.get('workspace_id', 'None')).first() + model = ( + QuerySet(Model) + .filter(id=self.data.get("id"), workspace_id=self.data.get("workspace_id", "None")) + .first() + ) if model is None: - raise AppApiException(500, _('Model does not exist')) - return {'id': str(model.id), 'provider': model.provider, 'name': model.name, 'model_type': model.model_type, - 'model_name': model.model_name, - 'status': model.status, - 'meta': model.meta, - 'workspace_id': model.workspace_id, - } + raise AppApiException(500, _("Model does not exist")) + return { + "id": str(model.id), + "provider": model.provider, + "name": model.name, + "model_type": model.model_type, + "model_name": model.model_name, + "status": model.status, + "meta": model.meta, + "workspace_id": model.workspace_id, + } def pause_download(self, with_valid=True): if with_valid: self.is_valid(raise_exception=True) - QuerySet(Model).filter(id=self.data.get('id')).update(status=Status.PAUSE_DOWNLOAD) + QuerySet(Model).filter(id=self.data.get("id")).update(status=Status.PAUSE_DOWNLOAD) return True @transaction.atomic def delete(self, with_valid=True): if with_valid: super().is_valid(raise_exception=True) - model_id = self.data.get('id') + model_id = self.data.get("id") model = Model.objects.filter(id=model_id).first() if model is None: return True @@ -198,30 +208,28 @@ def delete(self, with_valid=True): def edit(self, instance: Dict, user_id: str, with_valid=True): if with_valid: super().is_valid(raise_exception=True) - model = QuerySet(Model).filter(id=self.data.get('id')).first() + model = QuerySet(Model).filter(id=self.data.get("id")).first() - credential, model_credential, provider_handler = ModelSerializer.Edit( - data={**instance}).is_valid( - model=model) + credential, model_credential, provider_handler = ModelSerializer.Edit(data={**instance}).is_valid( + model=model + ) try: model.status = Status.SUCCESS - default_params = {item['field']: item['default_value'] for item in model.model_params_form} + default_params = {item["field"]: item["default_value"] for item in model.model_params_form} # 校验模型认证数据 - provider_handler.is_valid_credential(model.model_type, - instance.get("model_name"), - credential, - default_params, - raise_exception=True) + provider_handler.is_valid_credential( + model.model_type, instance.get("model_name"), credential, default_params, raise_exception=True + ) except AppApiException as e: if e.code == ValidCode.model_not_fount: model.status = Status.DOWNLOAD else: raise e - update_keys = ['credential', 'name', 'model_type', 'model_name'] + update_keys = ["credential", "name", "model_type", "model_name"] for update_key in update_keys: if update_key in instance and instance.get(update_key) is not None: - if update_key == 'credential': + if update_key == "credential": model_credential_str = json.dumps(credential) model.__setattr__(update_key, rsa_long_encrypt(model_credential_str)) else: @@ -235,38 +243,35 @@ def edit(self, instance: Dict, user_id: str, with_valid=True): return self.one(with_valid=False) class Edit(serializers.Serializer): - user_id = serializers.CharField(required=False, label=(_('user id'))) + user_id = serializers.CharField(required=False, label=(_("user id"))) - name = serializers.CharField(required=False, max_length=64, - label=(_("model name"))) + name = serializers.CharField(required=False, max_length=64, label=(_("model name"))) model_type = serializers.CharField(required=False, label=(_("model type"))) model_name = serializers.CharField(required=False, label=(_("base model"))) - credential = serializers.DictField(required=False, - label=(_("certification information"))) + credential = serializers.DictField(required=False, label=(_("certification information"))) workspace_id = serializers.CharField(required=False, label=(_("workspace id"))) def is_valid(self, model=None, raise_exception=False): super().is_valid(raise_exception=True) - filter_params = {'workspace_id': model.workspace_id} - if 'name' in self.data and self.data.get('name') is not None: - filter_params['name'] = self.data.get('name') + filter_params = {"workspace_id": model.workspace_id} + if "name" in self.data and self.data.get("name") is not None: + filter_params["name"] = self.data.get("name") if QuerySet(Model).exclude(id=model.id).filter(**filter_params).exists(): - raise AppApiException(500, _('base model【{model_name}】already exists').format( - model_name=self.data.get("name"))) + raise AppApiException( + 500, _("base model【{model_name}】already exists").format(model_name=self.data.get("name")) + ) ModelSerializer.model_to_dict(model) provider = model.provider - model_type = self.data.get('model_type') - model_name = self.data.get( - 'model_name') - credential = self.data.get('credential') + model_type = self.data.get("model_type") + model_name = self.data.get("model_name") + credential = self.data.get("credential") provider_handler = ModelProvideConstants[provider].value - model_credential = ModelProvideConstants[provider].value.get_model_credential(model_type, - model_name) + model_credential = ModelProvideConstants[provider].value.get_model_credential(model_type, model_name) source_model_credential = json.loads(rsa_long_decrypt(model.credential)) source_encryption_model_credential = model_credential.encryption_dict(source_model_credential) if credential is not None: @@ -276,7 +281,7 @@ def is_valid(self, model=None, raise_exception=False): return credential, model_credential, provider_handler class Create(serializers.Serializer): - user_id = serializers.UUIDField(required=True, label=_('user id')) + user_id = serializers.UUIDField(required=True, label=_("user id")) name = serializers.CharField(required=True, max_length=64, label=_("model name")) provider = serializers.CharField(required=True, label=_("provider")) model_type = serializers.CharField(required=True, label=_("model type")) @@ -287,21 +292,21 @@ class Create(serializers.Serializer): def is_valid(self, *, raise_exception=False): super().is_valid(raise_exception=True) - if QuerySet(Model).filter( - name=self.data.get('name'), - workspace_id=self.data.get('workspace_id', 'None') - ).exists(): + if ( + QuerySet(Model) + .filter(name=self.data.get("name"), workspace_id=self.data.get("workspace_id", "None")) + .exists() + ): raise AppApiException( - 500, - _('base model【{model_name}】already exists').format(model_name=self.data.get("name")) + 500, _("base model【{model_name}】already exists").format(model_name=self.data.get("name")) ) - default_params = {item['field']: item['default_value'] for item in self.data.get('model_params_form')} - ModelProvideConstants[self.data.get('provider')].value.is_valid_credential( - self.data.get('model_type'), - self.data.get('model_name'), - self.data.get('credential'), + default_params = {item["field"]: item["default_value"] for item in self.data.get("model_params_form")} + ModelProvideConstants[self.data.get("provider")].value.is_valid_credential( + self.data.get("model_type"), + self.data.get("model_name"), + self.data.get("credential"), default_params, - raise_exception=True + raise_exception=True, ) def insert(self, workspace_id, with_valid=True): @@ -315,28 +320,30 @@ def insert(self, workspace_id, with_valid=True): else: raise e - credential = self.data.get('credential') + credential = self.data.get("credential") model_data = { - 'id': uuid.uuid7(), - 'status': status, - 'user_id': self.data.get('user_id'), - 'name': self.data.get('name'), - 'credential': rsa_long_encrypt(json.dumps(credential)), - 'provider': self.data.get('provider'), - 'model_type': self.data.get('model_type'), - 'model_name': self.data.get('model_name'), - 'model_params_form': self.data.get('model_params_form'), - 'workspace_id': workspace_id + "id": uuid.uuid7(), + "status": status, + "user_id": self.data.get("user_id"), + "name": self.data.get("name"), + "credential": rsa_long_encrypt(json.dumps(credential)), + "provider": self.data.get("provider"), + "model_type": self.data.get("model_type"), + "model_name": self.data.get("model_name"), + "model_params_form": self.data.get("model_params_form"), + "workspace_id": workspace_id, } model = Model(**model_data) try: model.save() - if workspace_id != 'None': - UserResourcePermissionSerializer(data={ - 'workspace_id': workspace_id, - 'user_id': self.data.get('user_id'), - 'auth_target_type': AuthTargetType.MODEL.value - }).auth_resource(str(model.id)) + if workspace_id != "None": + UserResourcePermissionSerializer( + data={ + "workspace_id": workspace_id, + "user_id": self.data.get("user_id"), + "auth_target_type": AuthTargetType.MODEL.value, + } + ).auth_resource(str(model.id)) except Exception as save_error: # 可添加日志记录 raise AppApiException(500, _("Model saving failed")) from save_error @@ -349,12 +356,12 @@ def insert(self, workspace_id, with_valid=True): class Query(serializers.Serializer): user_id = serializers.CharField(required=True, label=_("User ID")) - name = serializers.CharField(required=False, max_length=64, label=_('model name')) - model_type = serializers.CharField(required=False, label=_('model type')) - model_name = serializers.CharField(required=False, label=_('base model')) - provider = serializers.CharField(required=False, label=_('provider')) - create_user = serializers.CharField(required=False, label=_('create user')) - workspace_id = serializers.CharField(required=False, label=_('workspace id')) + name = serializers.CharField(required=False, max_length=64, label=_("model name")) + model_type = serializers.CharField(required=False, label=_("model type")) + model_name = serializers.CharField(required=False, label=_("base model")) + provider = serializers.CharField(required=False, label=_("provider")) + create_user = serializers.CharField(required=False, label=_("create user")) + workspace_id = serializers.CharField(required=False, label=_("workspace id")) @staticmethod def is_x_pack_ee(): @@ -366,15 +373,23 @@ def list(self, workspace_id, with_valid): if with_valid: self.is_valid(raise_exception=True) user_id = self.data.get("user_id") - workspace_manage = is_workspace_manage_permission_read(user_id, workspace_id, 'MODEL:READ') + workspace_manage = is_workspace_manage_permission_read(user_id, workspace_id, "MODEL:READ") query_params = self._build_query_params(workspace_id, workspace_manage, user_id) is_x_pack_ee = self.is_x_pack_ee() - result = native_search(query_params, - select_string=get_file_content( - os.path.join(PROJECT_DIR, "apps", "models_provider", 'sql', - 'list_model.sql' if workspace_manage else ( - 'list_model_user_ee.sql' if is_x_pack_ee else 'list_model_user.sql') - ))) + result = native_search( + query_params, + select_string=get_file_content( + os.path.join( + PROJECT_DIR, + "apps", + "models_provider", + "sql", + "list_model.sql" + if workspace_manage + else ("list_model_user_ee.sql" if is_x_pack_ee else "list_model_user.sql"), + ) + ), + ) return ResourceMappingSerializer().get_resource_count(result) def share_list(self, workspace_id, with_valid=True): @@ -382,24 +397,20 @@ def share_list(self, workspace_id, with_valid=True): self.is_valid(raise_exception=True) user_id = self.data.get("user_id") query_params = self._build_query_params(workspace_id, False, user_id) - result = [ - self._build_model_data( - model - ) for model in query_params.get('model_query_set') - ] + result = [self._build_model_data(model) for model in query_params.get("model_query_set")] return ResourceMappingSerializer().get_resource_count(result) def model_list(self, workspace_id, with_valid=True): if with_valid: self.is_valid(raise_exception=True) user_id = self.data.get("user_id") - workspace_manage = is_workspace_manage_permission_read(user_id, workspace_id, 'MODEL:READ') + workspace_manage = is_workspace_manage_permission_read(user_id, workspace_id, "MODEL:READ") queryset = self._build_query_params(workspace_id, workspace_manage, user_id) get_authorized_model = DatabaseModelManage.get_model("get_authorized_model") shared_queryset = QuerySet(Model).none() if get_authorized_model is not None: - shared_queryset = self._build_query_params('None', False, user_id)['model_query_set'] + shared_queryset = self._build_query_params("None", False, user_id)["model_query_set"] shared_queryset = get_authorized_model(shared_queryset, workspace_id) # 构建共享模型和普通模型列表 @@ -410,95 +421,116 @@ def model_list(self, workspace_id, with_valid=True): queryset, select_string=get_file_content( os.path.join( - PROJECT_DIR, "apps", "models_provider", 'sql', - 'list_model.sql' if workspace_manage else ( - 'list_model_user_ee.sql' if is_x_pack_ee else 'list_model_user.sql') + PROJECT_DIR, + "apps", + "models_provider", + "sql", + "list_model.sql" + if workspace_manage + else ("list_model_user_ee.sql" if is_x_pack_ee else "list_model_user.sql"), ) - ) + ), ) - return { - "shared_model": shared_model, - "model": normal_model - } + return {"shared_model": shared_model, "model": normal_model} def _build_query_params(self, workspace_id, workspace_manage: bool, user_id): queryset = QuerySet(Model) if workspace_id: queryset = queryset.filter(workspace_id=workspace_id) - for field in ['name', 'model_type', 'model_name', 'provider', 'create_user']: + for field in ["name", "model_type", "model_name", "provider", "create_user"]: value = self.data.get(field) if value is not None: - if field == 'name': - queryset = queryset.filter(**{f'{field}__icontains': value}) - elif field == 'create_user': + if field == "name": + queryset = queryset.filter(**{f"{field}__icontains": value}) + elif field == "create_user": queryset = queryset.filter(user_id=value) else: queryset = queryset.filter(**{field: value}) queryset = queryset.order_by("-create_time") - return { - 'model_query_set': queryset, - 'workspace_user_resource_permission_query_set': QuerySet(WorkspaceUserResourcePermission).filter( - auth_target_type="MODEL", - workspace_id=workspace_id, - user_id=user_id)} if ( - not workspace_manage) else { - 'model_query_set': queryset, - } + return ( + { + "model_query_set": queryset, + "workspace_user_resource_permission_query_set": QuerySet(WorkspaceUserResourcePermission).filter( + auth_target_type="MODEL", workspace_id=workspace_id, user_id=user_id + ), + } + if (not workspace_manage) + else { + "model_query_set": queryset, + } + ) def _build_model_data(self, model): return { - 'id': str(model.id), - 'provider': model.provider, - 'name': model.name, - 'model_type': model.model_type, - 'model_name': model.model_name, - 'status': model.status, - 'meta': model.meta, - 'user_id': model.user_id, - 'username': model.user.username, - 'nick_name': model.user.nick_name, + "id": str(model.id), + "provider": model.provider, + "name": model.name, + "model_type": model.model_type, + "model_name": model.model_name, + "status": model.status, + "meta": model.meta, + "user_id": model.user_id, + "username": model.user.username, + "nick_name": model.user.nick_name, } def page(self, current_page, page_size): pass class ModelParams(serializers.Serializer): - id = serializers.UUIDField(required=True, label=_('model id')) + id = serializers.UUIDField(required=True, label=_("model id")) + workspace_id = serializers.UUIDField(required=False, label=_("workspace id")) def is_valid(self, *, raise_exception=False): super().is_valid(raise_exception=True) - model = QuerySet(Model).filter(id=self.data.get("id")).first() + + validated_data = self.validated_data + model = ( + QuerySet(Model) + .filter( + id=validated_data.get("id"), + workspace_id=validated_data.get("workspace_id"), + ) + .first() + ) + if model is None: raise AppApiException(500, _("Model does not exist")) + return model + def get_model_params(self, with_valid=True): + model = None + if with_valid: - self.is_valid(raise_exception=True) - model_id = self.data.get('id') - model = QuerySet(Model).filter(id=model_id).first() - return model.model_params_form + model = self.is_valid(raise_exception=True) + + return model.model_params_form if model else None def save_model_params_form(self, model_params_form, with_valid=True): + model = None if with_valid: - self.is_valid(raise_exception=True) + model = self.is_valid(raise_exception=True) if model_params_form is None: model_params_form = [] - model_id = self.data.get('id') - model = QuerySet(Model).filter(id=model_id).first() if not isinstance(model_params_form, list): - raise AppApiException(500, _('model_params_form must be a list')) + raise AppApiException(500, _("model_params_form must be a list")) # 还需要校验几个字段:label required default_value # 校验每个配置项的必要字段 for index, param in enumerate(model_params_form): if not isinstance(param, dict): - raise AppApiException(500, _('The {index}th item in model_params_form must be a dictionary').format( - index=index)) + raise AppApiException( + 500, _("The {index}th item in model_params_form must be a dictionary").format(index=index) + ) # 校验 label 字段 - if 'label' not in param or param['label'] is None: - raise AppApiException(500, - _('The label field is required for the {index}th item in model_params_form').format( - index=index)) + if "label" not in param or param["label"] is None: + raise AppApiException( + 500, + _("The label field is required for the {index}th item in model_params_form").format( + index=index + ), + ) model.model_params_form = model_params_form model.save() @@ -506,31 +538,31 @@ def save_model_params_form(self, model_params_form, with_valid=True): class WorkspaceSharedModelSerializer(serializers.Serializer): - workspace_id = serializers.CharField(required=True, label=_('workspace id')) - name = serializers.CharField(required=False, max_length=64, label=_('model name')) - model_type = serializers.CharField(required=False, label=_('model type')) - model_name = serializers.CharField(required=False, label=_('base model')) - provider = serializers.CharField(required=False, label=_('provider')) - create_user = serializers.CharField(required=False, label=_('create user')) + workspace_id = serializers.CharField(required=True, label=_("workspace id")) + name = serializers.CharField(required=False, max_length=64, label=_("model name")) + model_type = serializers.CharField(required=False, label=_("model type")) + model_name = serializers.CharField(required=False, label=_("base model")) + provider = serializers.CharField(required=False, label=_("provider")) + create_user = serializers.CharField(required=False, label=_("create user")) def get_share_model_list(self): self.is_valid(raise_exception=True) - workspace_id = self.data.get('workspace_id') + workspace_id = self.data.get("workspace_id") queryset = self._build_queryset(workspace_id) return [ { - 'id': str(model.id), - 'provider': model.provider, - 'name': model.name, - 'model_type': model.model_type, - 'model_name': model.model_name, - 'status': model.status, - 'meta': model.meta, - 'user_id': model.user_id, - 'nick_name': model.user.nick_name, - 'username': model.user.username + "id": str(model.id), + "provider": model.provider, + "name": model.name, + "model_type": model.model_type, + "model_name": model.model_name, + "status": model.status, + "meta": model.meta, + "user_id": model.user_id, + "nick_name": model.user.nick_name, + "username": model.user.username, } for model in queryset.order_by("-create_time") ] @@ -542,12 +574,12 @@ def _build_queryset(self, workspace_id): if get_authorized_model is not None: queryset = get_authorized_model(queryset, workspace_id) - for field in ['name', 'model_type', 'model_name', 'provider', 'create_user']: + for field in ["name", "model_type", "model_name", "provider", "create_user"]: value = self.data.get(field) if value is not None: - if field == 'name': - queryset = queryset.filter(**{f'{field}__icontains': value}) - elif field == 'create_user': + if field == "name": + queryset = queryset.filter(**{f"{field}__icontains": value}) + elif field == "create_user": queryset = queryset.filter(user_id=value) else: queryset = queryset.filter(**{field: value}) diff --git a/apps/models_provider/views/model.py b/apps/models_provider/views/model.py index fe498b20ac1..fa9fb69b80d 100644 --- a/apps/models_provider/views/model.py +++ b/apps/models_provider/views/model.py @@ -195,7 +195,7 @@ class ModelParamsForm(APIView): RoleConstants.USER.get_workspace_role(),) def get(self, request: Request, workspace_id: str, model_id: str): return result.success( - ModelSerializer.ModelParams(data={'id': model_id}).get_model_params()) + ModelSerializer.ModelParams(data={'id': model_id, 'workspace_id': workspace_id}).get_model_params()) @extend_schema(methods=['PUT'], summary=_('Save model parameter form'), @@ -217,7 +217,7 @@ def get(self, request: Request, workspace_id: str, model_id: str): ) def put(self, request: Request, workspace_id: str, model_id: str): return result.success( - ModelSerializer.ModelParams(data={'id': model_id}).save_model_params_form(request.data)) + ModelSerializer.ModelParams(data={'id': model_id, 'workspace_id': workspace_id}).save_model_params_form(request.data)) class ModelMeta(APIView): authentication_classes = [TokenAuth]