From 4f64c1e4c892e55af9a40cbc363bcf88f731e93d Mon Sep 17 00:00:00 2001 From: wangzifei Date: Tue, 22 Sep 2026 19:07:05 +0800 Subject: [PATCH] fix: report conflicting server names when merging referenced MCP tools MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Referencing multiple MCP tools merges their server config dicts by key. Server names are user-defined free text, so two referenced tools defining the same server name silently overwrote each other and only the last server's tools were visible (issue #7120) — every tool worked when referenced alone. Raise an explicit error naming the conflicting servers instead of dropping them silently. Fixes #7120 Signed-off-by: wangzifei --- .../step/chat_step/impl/base_chat_step.py | 1627 +++++++++-------- .../ai_chat_step_node/impl/base_chat_node.py | 1291 ++++++------- 2 files changed, 1470 insertions(+), 1448 deletions(-) diff --git a/apps/application/chat_pipeline/step/chat_step/impl/base_chat_step.py b/apps/application/chat_pipeline/step/chat_step/impl/base_chat_step.py index 62fd55166f4..7acd6cf58b0 100644 --- a/apps/application/chat_pipeline/step/chat_step/impl/base_chat_step.py +++ b/apps/application/chat_pipeline/step/chat_step/impl/base_chat_step.py @@ -1,808 +1,819 @@ -# coding=utf-8 -""" -@project: maxkb -@Author:虎 -@file: base_chat_step.py -@date:2024/1/9 18:25 -@desc: 对话step Base实现 -""" - -import json -import time -import traceback -from typing import List - -import uuid_utils.compat as uuid -from application.chat_pipeline.I_base_chat_pipeline import ParagraphPipelineModel -from application.chat_pipeline.pipeline_manage import PipelineManage -from application.chat_pipeline.step.chat_step.i_chat_step import IChatStep, PostResponseHandler -from application.flow.tools import Reasoning, get_tools, mcp_response_generator -from application.long_term_memory import extract_long_term_memory -from application.models import ( - Application, - ApplicationAccessToken, - ApplicationApiKey, - ApplicationChatUserStats, - ApplicationLongTermMemory, - ChatUserType, -) -from common.exception.app_exception import AppApiException -from common.utils.logger import maxkb_logger -from common.utils.rsa_util import rsa_long_decrypt -from common.utils.shared_resource_auth import filter_authorized_ids, get_runtime_user_id -from common.utils.tool_code import ToolExecutor -from django.db.models import QuerySet -from django.http import StreamingHttpResponse -from django.utils.translation import gettext as _ -from langchain.chat_models.base import BaseChatModel -from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, HumanMessage, SystemMessage -from models_provider.tools import get_model_instance_by_model_workspace_id -from rest_framework import status -from tools.models import Tool, ToolType - - -def add_access_num(chat_user_id=None, chat_user_type=None, application_id=None): - if [ChatUserType.ANONYMOUS_USER.value, ChatUserType.CHAT_USER.value].__contains__( - chat_user_type - ) and application_id is not None: - application_public_access_client = ( - QuerySet(ApplicationChatUserStats) - .filter(chat_user_id=chat_user_id, chat_user_type=chat_user_type, application_id=application_id) - .first() - ) - if application_public_access_client is not None: - application_public_access_client.access_num = application_public_access_client.access_num + 1 - application_public_access_client.intraday_access_num = ( - application_public_access_client.intraday_access_num + 1 - ) - application_public_access_client.save() - - -def write_context(step, manage, request_token, response_token, all_text): - step.context["message_tokens"] = request_token - step.context["answer_tokens"] = response_token - current_time = time.time() - step.context["answer_text"] = all_text - step.context["run_time"] = current_time - step.context["start_time"] - manage.context["run_time"] = current_time - manage.context["start_time"] - manage.context["message_tokens"] = manage.context["message_tokens"] + request_token - manage.context["answer_tokens"] = manage.context["answer_tokens"] + response_token - - -def event_content( - response, - chat_id, - chat_record_id, - paragraph_list: List[ParagraphPipelineModel], - post_response_handler: PostResponseHandler, - manage, - step, - chat_model, - message_list: List[BaseMessage], - problem_text: str, - padding_problem_text: str = None, - chat_user_id=None, - chat_user_type=None, - is_ai_chat: bool = None, - model_setting=None, -): - if model_setting is None: - model_setting = {} - reasoning_content_enable = model_setting.get("reasoning_content_enable", False) - reasoning_content_start = model_setting.get("reasoning_content_start", "") - reasoning_content_end = model_setting.get("reasoning_content_end", "") - reasoning = Reasoning(reasoning_content_start, reasoning_content_end) - all_text = "" - reasoning_content = "" - try: - response_reasoning_content = False - for chunk in response: - reasoning_chunk = reasoning.get_reasoning_content(chunk) - content_chunk = reasoning_chunk.get("content") - if "reasoning_content" in chunk.additional_kwargs: - response_reasoning_content = True - reasoning_content_chunk = chunk.additional_kwargs.get("reasoning_content", "") - else: - reasoning_content_chunk = reasoning_chunk.get("reasoning_content") - content_chunk = reasoning._normalize_content(content_chunk) - all_text += content_chunk - if reasoning_content_chunk is None: - reasoning_content_chunk = "" - reasoning_content += reasoning_content_chunk - yield manage.get_base_to_response().to_stream_chunk_response( - chat_id, - str(chat_record_id), - "ai-chat-node", - [], - content_chunk, - False, - 0, - 0, - { - "node_is_end": False, - "view_type": "many_view", - "node_type": "ai-chat-node", - "real_node_id": "ai-chat-node", - "reasoning_content": reasoning_content_chunk if reasoning_content_enable else "", - }, - ) - reasoning_chunk = reasoning.get_end_reasoning_content() - all_text += reasoning_chunk.get("content") - reasoning_content_chunk = "" - if not response_reasoning_content: - reasoning_content_chunk = reasoning_chunk.get("reasoning_content") - yield manage.get_base_to_response().to_stream_chunk_response( - chat_id, - str(chat_record_id), - "ai-chat-node", - [], - reasoning_chunk.get("content"), - False, - 0, - 0, - { - "node_is_end": False, - "view_type": "many_view", - "node_type": "ai-chat-node", - "real_node_id": "ai-chat-node", - "reasoning_content": reasoning_content_chunk if reasoning_content_enable else "", - }, - ) - # 获取token - if is_ai_chat: - try: - request_token = chat_model.get_num_tokens_from_messages(message_list) - response_token = chat_model.get_num_tokens(all_text) - except Exception as e: - request_token = 0 - response_token = 0 - else: - request_token = 0 - response_token = 0 - write_context(step, manage, request_token, response_token, all_text) - post_response_handler.handler( - chat_id, - chat_record_id, - paragraph_list, - problem_text, - all_text, - manage, - step, - padding_problem_text, - reasoning_content=reasoning_content if reasoning_content_enable else "", - ) - yield manage.get_base_to_response().to_stream_chunk_response( - chat_id, - str(chat_record_id), - "ai-chat-node", - [], - "", - True, - request_token, - response_token, - {"node_is_end": True, "view_type": "many_view", "node_type": "ai-chat-node"}, - ) - if not manage.debug: - add_access_num(chat_user_id, chat_user_type, manage.context.get("application_id")) - except BaseException as e: - if isinstance(e, GeneratorExit): - maxkb_logger.error(f"Generator was closed (client disconnected)") - else: - maxkb_logger.error(f"{str(e)}:{traceback.format_exc()}") - all_text = "Exception:" + str(e) - write_context(step, manage, 0, 0, all_text) - post_response_handler.handler( - chat_id, - chat_record_id, - paragraph_list, - problem_text, - all_text, - manage, - step, - padding_problem_text, - reasoning_content=reasoning_content if reasoning_content_enable else "", - ) - if not manage.debug: - add_access_num(chat_user_id, chat_user_type, manage.context.get("application_id")) - yield manage.get_base_to_response().to_stream_chunk_response( - chat_id, - str(chat_record_id), - "ai-chat-node", - [], - all_text, - False, - 0, - 0, - { - "node_is_end": False, - "view_type": "many_view", - "node_type": "ai-chat-node", - "real_node_id": "ai-chat-node", - "reasoning_content": "", - }, - ) - - -class BaseChatStep(IChatStep): - def execute( - self, - message_list: List[BaseMessage], - chat_id, - problem_text, - post_response_handler: PostResponseHandler, - model_id: str = None, - workspace_id: str = None, - paragraph_list=None, - manage: PipelineManage = None, - padding_problem_text: str = None, - stream: bool = True, - chat_user_id=None, - chat_user_type=None, - no_references_setting=None, - model_params_setting=None, - model_setting=None, - mcp_tool_ids=None, - mcp_servers="", - mcp_source="referencing", - tool_ids=None, - application_ids=None, - skill_tool_ids=None, - mcp_output_enable=True, - **kwargs, - ): - chat_model = ( - get_model_instance_by_model_workspace_id(model_id, workspace_id, **(model_params_setting or {})) - if model_id is not None - else None - ) - if stream: - return self.execute_stream( - message_list, - chat_id, - problem_text, - post_response_handler, - chat_model, - paragraph_list, - manage, - padding_problem_text, - chat_user_id, - chat_user_type, - no_references_setting, - model_setting, - mcp_tool_ids, - mcp_servers, - mcp_source, - tool_ids, - application_ids, - skill_tool_ids, - workspace_id, - mcp_output_enable, - ) - else: - return self.execute_block( - message_list, - chat_id, - problem_text, - post_response_handler, - chat_model, - paragraph_list, - manage, - padding_problem_text, - chat_user_id, - chat_user_type, - no_references_setting, - model_setting, - mcp_tool_ids, - mcp_servers, - mcp_source, - tool_ids, - application_ids, - skill_tool_ids, - workspace_id, - mcp_output_enable, - ) - - def get_details(self, manage, **kwargs): - # 提取长期记忆 - extract_long_term_memory.apply_async( - args=( - manage.context.get("workspace_id"), - manage.context.get("application_id"), - manage.context.get("chat_user_id"), - ), - countdown=1, - ) - return { - "status": self.status, - "err_message": self.err_message, - "step_type": "chat_step", - "run_time": self.context.get("run_time") or 0, - "model_id": str(manage.context["model_id"]), - "message_list": self.reset_message_list( - self.context["step_args"].get("message_list"), self.context.get("answer_text") - ), - "message_tokens": self.context.get("message_tokens"), - "answer_tokens": self.context.get("answer_tokens"), - "cost": 0, - } - - @staticmethod - def reset_message_list(message_list: List[BaseMessage], answer_text): - result = [ - { - "role": "user" - if isinstance(message, HumanMessage) - else ("system" if isinstance(message, SystemMessage) else "ai"), - "content": message.content, - } - for message in message_list - ] - result.append({"role": "ai", "content": answer_text}) - return result - - def _handle_mcp_request( - self, - mcp_source, - mcp_servers, - mcp_tool_ids, - tool_ids, - application_ids, - skill_tool_ids, - mcp_output_enable, - chat_model, - system_prompt, - message_list, - agent_id, - chat_id, - workspace_id, - runtime_user_id=None, - ): - - mcp_servers_config = {} - - # 迁移过来mcp_source是None - if mcp_source is None: - mcp_source = "custom" - # 兼容老数据 - if not mcp_tool_ids: - mcp_tool_ids = [] - if mcp_source == "custom" and mcp_servers: - mcp_servers_config = json.loads(mcp_servers) - elif mcp_tool_ids: - mcp_tools = QuerySet(Tool).filter(id__in=mcp_tool_ids).values() - for mcp_tool in mcp_tools: - if mcp_tool and mcp_tool["is_active"]: - mcp_servers_config = {**mcp_servers_config, **json.loads(mcp_tool["code"])} - # 校验代码是否包括禁止的关键字 - ToolExecutor().validate_mcp_transport(json.dumps(mcp_servers_config)) - - tool_init_params = {} - tools = get_tools("APPLICATION", agent_id, tool_ids, workspace_id, runtime_user_id) - if tool_ids and len(tool_ids) > 0: # 如果有工具ID,则将其转换为MCP - self.context["tool_ids"] = tool_ids - for tool_id in tool_ids: - tool = QuerySet(Tool).filter(id=tool_id, tool_type=ToolType.CUSTOM).first() - if tool is None or tool.is_active is False: - continue - executor = ToolExecutor() - init_params_default_value = {i["field"]: i.get('default_value') for i in tool.init_field_list} - if tool.init_params is not None: - tool_init_params = init_params_default_value | json.loads(rsa_long_decrypt(tool.init_params)) - else: - tool_init_params = init_params_default_value - tool_config = executor.get_tool_mcp_config(tool, tool_init_params) - - mcp_servers_config[str(tool.id)] = tool_config - - if application_ids and len(application_ids) > 0: - self.context["application_ids"] = application_ids - for application_id in application_ids: - app = QuerySet(Application).filter(id=application_id, is_publish=True).first() - if app is None: - continue - app_key = QuerySet(ApplicationApiKey).filter(application_id=application_id, is_active=True).first() - if app_key is not None: - api_key = app_key.secret_key - application_access_token = ( - QuerySet(ApplicationAccessToken).filter(application_id=app_key.application_id).first() - ) - if application_access_token is not None and application_access_token.authentication: - raise AppApiException( - 500, - _("Agent 【{name}】 access token authentication is not supported for agent tool").format( - name=app.name - ), - ) - else: - raise AppApiException( - 500, _("Agent Key is required for agent tool 【{name}】").format(name=app.name) - ) - executor = ToolExecutor() - app_config = executor.get_app_mcp_config(api_key) - mcp_servers_config[app.name] = app_config - - if skill_tool_ids and len(skill_tool_ids) > 0: - self.context["skill_tool_ids"] = skill_tool_ids - skill_file_items = [] - - for tool_id in skill_tool_ids: - tool = QuerySet(Tool).filter(id=tool_id, is_active=True).first() - if tool is None or tool.is_active is False: - continue - init_params_default_value = {i["field"]: i.get("default_value") for i in tool.init_field_list} - if tool.init_params is not None: - params = init_params_default_value | json.loads(rsa_long_decrypt(tool.init_params)) - else: - params = init_params_default_value - - skill_file_items.append({"tool_id": str(tool.id), "file_id": tool.code, "params": params}) - mcp_servers_config["skills"] = skill_file_items - - if len(mcp_servers_config) > 0 or len(tools) > 0: - source_id = agent_id - source_type = "APPLICATION" - return mcp_response_generator( - chat_model, - system_prompt, - message_list, - mcp_servers_config, - mcp_output_enable, - tool_init_params, - source_id, - source_type, - chat_id, - tools, - ) - - return None - - def get_stream_result( - self, - message_list: List[BaseMessage], - chat_model: BaseChatModel = None, - paragraph_list=None, - no_references_setting=None, - problem_text=None, - mcp_tool_ids=None, - mcp_servers="", - mcp_source="referencing", - tool_ids=None, - application_ids=None, - skill_tool_ids=None, - workspace_id=None, - mcp_output_enable=True, - agent_id=None, - chat_id=None, - chat_user_id=None, - chat_user_type=None, - ): - if paragraph_list is None: - paragraph_list = [] - directly_return_chunk_list = [ - AIMessageChunk(content=paragraph.content) - for paragraph in paragraph_list - if ( - paragraph.hit_handling_method == "directly_return" - and paragraph.similarity >= paragraph.directly_return_similarity - ) - ] - if directly_return_chunk_list is not None and len(directly_return_chunk_list) > 0: - return iter(directly_return_chunk_list), False - elif len(paragraph_list) == 0 and no_references_setting.get("status") == "designated_answer": - return iter( - [AIMessageChunk(content=no_references_setting.get("value").replace("{question}", problem_text))] - ), False - if chat_model is None: - return iter( - [ - AIMessageChunk( - _( - "Sorry, the AI model is not configured. Please go to the application to set up the AI model first." - ) - ) - ] - ), False - else: - user_system_prompt = None - filtered_message_list = [] - long_term_memory = ( - QuerySet(ApplicationLongTermMemory).filter(chat_user_id=chat_user_id, application_id=agent_id).first() - ) - if long_term_memory is not None: - memory = long_term_memory.memory - else: - memory = "" - - # print(chat_user_id, chat_user_type) - for msg in message_list: - if isinstance(msg, SystemMessage): - if isinstance(msg.content, str): - user_system_prompt = msg.content.replace("{memory}", memory) - msg.content = user_system_prompt - elif isinstance(msg.content, list): - user_system_prompt = "".join( - item.get("text", "") if isinstance(item, dict) else str(item) for item in msg.content - ) - else: - user_system_prompt = str(msg.content) - else: - filtered_message_list.append(msg) - runtime_user_id = get_runtime_user_id(chat_user_id=chat_user_id, chat_user_type=chat_user_type) - # 过滤tool_id - all_tool_ids = list(set((mcp_tool_ids or []) + (tool_ids or []) + (skill_tool_ids or []))) - authorized_set = set(filter_authorized_ids("tool", all_tool_ids, workspace_id, user_id=runtime_user_id)) - - mcp_tool_ids = [i for i in (mcp_tool_ids or []) if i in authorized_set] - tool_ids = [i for i in (tool_ids or []) if i in authorized_set] - skill_tool_ids = [i for i in (skill_tool_ids or []) if i in authorized_set] - # 处理 MCP 请求 - mcp_result = self._handle_mcp_request( - mcp_source, - mcp_servers, - mcp_tool_ids, - tool_ids, - application_ids, - skill_tool_ids, - mcp_output_enable, - chat_model, - user_system_prompt, - filtered_message_list, - agent_id, - chat_id, - workspace_id, - runtime_user_id, - ) - if mcp_result: - return mcp_result, True - return chat_model.stream(message_list), True - - def execute_stream( - self, - message_list: List[BaseMessage], - chat_id, - problem_text, - post_response_handler: PostResponseHandler, - chat_model: BaseChatModel = None, - paragraph_list=None, - manage: PipelineManage = None, - padding_problem_text: str = None, - chat_user_id=None, - chat_user_type=None, - no_references_setting=None, - model_setting=None, - mcp_tool_ids=None, - mcp_servers="", - mcp_source="referencing", - tool_ids=None, - application_ids=None, - skill_tool_ids=None, - workspace_id=None, - mcp_output_enable=True, - ): - chat_result, is_ai_chat = self.get_stream_result( - message_list, - chat_model, - paragraph_list, - no_references_setting, - problem_text, - mcp_tool_ids, - mcp_servers, - mcp_source, - tool_ids, - application_ids, - skill_tool_ids, - workspace_id, - mcp_output_enable, - manage.context.get("application_id"), - chat_id, - chat_user_id, - chat_user_type, - ) - chat_record_id = ( - self.context.get("step_args", {}).get("chat_record_id") - if self.context.get("step_args", {}).get("chat_record_id") - else uuid.uuid7() - ) - r = StreamingHttpResponse( - streaming_content=event_content( - chat_result, - chat_id, - chat_record_id, - paragraph_list, - post_response_handler, - manage, - self, - chat_model, - message_list, - problem_text, - padding_problem_text, - chat_user_id, - chat_user_type, - is_ai_chat, - model_setting, - ), - content_type="text/event-stream;charset=utf-8", - ) - - r["Cache-Control"] = "no-cache" - return r - - def get_block_result( - self, - message_list: List[BaseMessage], - chat_model: BaseChatModel = None, - paragraph_list=None, - no_references_setting=None, - problem_text=None, - mcp_tool_ids=None, - mcp_servers="", - mcp_source="referencing", - tool_ids=None, - application_ids=None, - skill_tool_ids=None, - workspace_id=None, - mcp_output_enable=True, - application_id=None, - chat_id=None, - chat_user_id=None, - chat_user_type=None, - ): - if paragraph_list is None: - paragraph_list = [] - directly_return_chunk_list = [ - AIMessageChunk(content=paragraph.content) - for paragraph in paragraph_list - if ( - paragraph.hit_handling_method == "directly_return" - and paragraph.similarity >= paragraph.directly_return_similarity - ) - ] - if directly_return_chunk_list is not None and len(directly_return_chunk_list) > 0: - return directly_return_chunk_list[0], False - elif len(paragraph_list) == 0 and no_references_setting.get("status") == "designated_answer": - return AIMessage(no_references_setting.get("value").replace("{question}", problem_text)), False - if chat_model is None: - return AIMessage( - _("Sorry, the AI model is not configured. Please go to the application to set up the AI model first.") - ), False - else: - runtime_user_id = get_runtime_user_id(chat_user_id=chat_user_id, chat_user_type=chat_user_type) - # 过滤tool_id - all_tool_ids = list(set((mcp_tool_ids or []) + (tool_ids or []) + (skill_tool_ids or []))) - authorized_set = set(filter_authorized_ids("tool", all_tool_ids, workspace_id, user_id=runtime_user_id)) - - mcp_tool_ids = [i for i in (mcp_tool_ids or []) if i in authorized_set] - tool_ids = [i for i in (tool_ids or []) if i in authorized_set] - skill_tool_ids = [i for i in (skill_tool_ids or []) if i in authorized_set] - # 处理 MCP 请求 - mcp_result = self._handle_mcp_request( - mcp_source, - mcp_servers, - mcp_tool_ids, - tool_ids, - application_ids, - skill_tool_ids, - mcp_output_enable, - chat_model, - "", - message_list, - application_id, - chat_id, - workspace_id, - runtime_user_id, - ) - if mcp_result: - return mcp_result, True - return chat_model.invoke(message_list), True - - def execute_block( - self, - message_list: List[BaseMessage], - chat_id, - problem_text, - post_response_handler: PostResponseHandler, - chat_model: BaseChatModel = None, - paragraph_list=None, - manage: PipelineManage = None, - padding_problem_text: str = None, - chat_user_id=None, - chat_user_type=None, - no_references_setting=None, - model_setting=None, - mcp_tool_ids=None, - mcp_servers="", - mcp_source="referencing", - tool_ids=None, - application_ids=None, - skill_tool_ids=None, - workspace_id=None, - mcp_output_enable=True, - ): - reasoning_content_enable = model_setting.get("reasoning_content_enable", False) - reasoning_content_start = model_setting.get("reasoning_content_start", "") - reasoning_content_end = model_setting.get("reasoning_content_end", "") - reasoning = Reasoning(reasoning_content_start, reasoning_content_end) - chat_record_id = uuid.uuid7() - # 调用模型 - try: - chat_result, is_ai_chat = self.get_block_result( - message_list, - chat_model, - paragraph_list, - no_references_setting, - problem_text, - mcp_tool_ids, - mcp_servers, - mcp_source, - tool_ids, - application_ids, - skill_tool_ids, - workspace_id, - mcp_output_enable, - manage.context.get("application_id"), - chat_id, - chat_user_id, - chat_user_type, - ) - if is_ai_chat: - request_token = chat_model.get_num_tokens_from_messages(message_list) - response_token = chat_model.get_num_tokens(chat_result.content) - else: - request_token = 0 - response_token = 0 - write_context(self, manage, request_token, response_token, chat_result.content) - reasoning_result = reasoning.get_reasoning_content(chat_result) - reasoning_result_end = reasoning.get_end_reasoning_content() - content = reasoning_result.get("content") + reasoning_result_end.get("content") - if "reasoning_content" in chat_result.response_metadata: - reasoning_content = chat_result.response_metadata.get("reasoning_content", "") or "" - else: - reasoning_content = (reasoning_result.get("reasoning_content") or "") + ( - reasoning_result_end.get("reasoning_content") or "" - ) - post_response_handler.handler( - chat_id, - chat_record_id, - paragraph_list, - problem_text, - content, - manage, - self, - padding_problem_text, - reasoning_content=reasoning_content, - ) - if not manage.debug: - add_access_num(chat_user_id, chat_user_type, manage.context.get("application_id")) - return manage.get_base_to_response().to_block_response( - str(chat_id), - str(chat_record_id), - content, - True, - request_token, - response_token, - { - "reasoning_content": reasoning_content if reasoning_content_enable else "", - "answer_list": [ - {"content": content, "reasoning_content": reasoning_content if reasoning_content_enable else ""} - ], - }, - ) - except Exception as e: - all_text = "Exception:" + str(e) - write_context(self, manage, 0, 0, all_text) - post_response_handler.handler( - chat_id, - chat_record_id, - paragraph_list, - problem_text, - all_text, - manage, - self, - padding_problem_text, - reasoning_content="", - ) - if not manage.debug: - add_access_num(chat_user_id, chat_user_type, manage.context.get("application_id")) - return manage.get_base_to_response().to_block_response( - str(chat_id), str(chat_record_id), all_text, True, 0, 0, _status=status.HTTP_500_INTERNAL_SERVER_ERROR - ) +# coding=utf-8 +""" +@project: maxkb +@Author:虎 +@file: base_chat_step.py +@date:2024/1/9 18:25 +@desc: 对话step Base实现 +""" + +import json +import time +import traceback +from typing import List + +import uuid_utils.compat as uuid +from application.chat_pipeline.I_base_chat_pipeline import ParagraphPipelineModel +from application.chat_pipeline.pipeline_manage import PipelineManage +from application.chat_pipeline.step.chat_step.i_chat_step import IChatStep, PostResponseHandler +from application.flow.tools import Reasoning, get_tools, mcp_response_generator +from application.long_term_memory import extract_long_term_memory +from application.models import ( + Application, + ApplicationAccessToken, + ApplicationApiKey, + ApplicationChatUserStats, + ApplicationLongTermMemory, + ChatUserType, +) +from common.exception.app_exception import AppApiException +from common.utils.logger import maxkb_logger +from common.utils.rsa_util import rsa_long_decrypt +from common.utils.shared_resource_auth import filter_authorized_ids, get_runtime_user_id +from common.utils.tool_code import ToolExecutor +from django.db.models import QuerySet +from django.http import StreamingHttpResponse +from django.utils.translation import gettext as _ +from langchain.chat_models.base import BaseChatModel +from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, HumanMessage, SystemMessage +from models_provider.tools import get_model_instance_by_model_workspace_id +from rest_framework import status +from tools.models import Tool, ToolType + + +def add_access_num(chat_user_id=None, chat_user_type=None, application_id=None): + if [ChatUserType.ANONYMOUS_USER.value, ChatUserType.CHAT_USER.value].__contains__( + chat_user_type + ) and application_id is not None: + application_public_access_client = ( + QuerySet(ApplicationChatUserStats) + .filter(chat_user_id=chat_user_id, chat_user_type=chat_user_type, application_id=application_id) + .first() + ) + if application_public_access_client is not None: + application_public_access_client.access_num = application_public_access_client.access_num + 1 + application_public_access_client.intraday_access_num = ( + application_public_access_client.intraday_access_num + 1 + ) + application_public_access_client.save() + + +def write_context(step, manage, request_token, response_token, all_text): + step.context["message_tokens"] = request_token + step.context["answer_tokens"] = response_token + current_time = time.time() + step.context["answer_text"] = all_text + step.context["run_time"] = current_time - step.context["start_time"] + manage.context["run_time"] = current_time - manage.context["start_time"] + manage.context["message_tokens"] = manage.context["message_tokens"] + request_token + manage.context["answer_tokens"] = manage.context["answer_tokens"] + response_token + + +def event_content( + response, + chat_id, + chat_record_id, + paragraph_list: List[ParagraphPipelineModel], + post_response_handler: PostResponseHandler, + manage, + step, + chat_model, + message_list: List[BaseMessage], + problem_text: str, + padding_problem_text: str = None, + chat_user_id=None, + chat_user_type=None, + is_ai_chat: bool = None, + model_setting=None, +): + if model_setting is None: + model_setting = {} + reasoning_content_enable = model_setting.get("reasoning_content_enable", False) + reasoning_content_start = model_setting.get("reasoning_content_start", "") + reasoning_content_end = model_setting.get("reasoning_content_end", "") + reasoning = Reasoning(reasoning_content_start, reasoning_content_end) + all_text = "" + reasoning_content = "" + try: + response_reasoning_content = False + for chunk in response: + reasoning_chunk = reasoning.get_reasoning_content(chunk) + content_chunk = reasoning_chunk.get("content") + if "reasoning_content" in chunk.additional_kwargs: + response_reasoning_content = True + reasoning_content_chunk = chunk.additional_kwargs.get("reasoning_content", "") + else: + reasoning_content_chunk = reasoning_chunk.get("reasoning_content") + content_chunk = reasoning._normalize_content(content_chunk) + all_text += content_chunk + if reasoning_content_chunk is None: + reasoning_content_chunk = "" + reasoning_content += reasoning_content_chunk + yield manage.get_base_to_response().to_stream_chunk_response( + chat_id, + str(chat_record_id), + "ai-chat-node", + [], + content_chunk, + False, + 0, + 0, + { + "node_is_end": False, + "view_type": "many_view", + "node_type": "ai-chat-node", + "real_node_id": "ai-chat-node", + "reasoning_content": reasoning_content_chunk if reasoning_content_enable else "", + }, + ) + reasoning_chunk = reasoning.get_end_reasoning_content() + all_text += reasoning_chunk.get("content") + reasoning_content_chunk = "" + if not response_reasoning_content: + reasoning_content_chunk = reasoning_chunk.get("reasoning_content") + yield manage.get_base_to_response().to_stream_chunk_response( + chat_id, + str(chat_record_id), + "ai-chat-node", + [], + reasoning_chunk.get("content"), + False, + 0, + 0, + { + "node_is_end": False, + "view_type": "many_view", + "node_type": "ai-chat-node", + "real_node_id": "ai-chat-node", + "reasoning_content": reasoning_content_chunk if reasoning_content_enable else "", + }, + ) + # 获取token + if is_ai_chat: + try: + request_token = chat_model.get_num_tokens_from_messages(message_list) + response_token = chat_model.get_num_tokens(all_text) + except Exception as e: + request_token = 0 + response_token = 0 + else: + request_token = 0 + response_token = 0 + write_context(step, manage, request_token, response_token, all_text) + post_response_handler.handler( + chat_id, + chat_record_id, + paragraph_list, + problem_text, + all_text, + manage, + step, + padding_problem_text, + reasoning_content=reasoning_content if reasoning_content_enable else "", + ) + yield manage.get_base_to_response().to_stream_chunk_response( + chat_id, + str(chat_record_id), + "ai-chat-node", + [], + "", + True, + request_token, + response_token, + {"node_is_end": True, "view_type": "many_view", "node_type": "ai-chat-node"}, + ) + if not manage.debug: + add_access_num(chat_user_id, chat_user_type, manage.context.get("application_id")) + except BaseException as e: + if isinstance(e, GeneratorExit): + maxkb_logger.error(f"Generator was closed (client disconnected)") + else: + maxkb_logger.error(f"{str(e)}:{traceback.format_exc()}") + all_text = "Exception:" + str(e) + write_context(step, manage, 0, 0, all_text) + post_response_handler.handler( + chat_id, + chat_record_id, + paragraph_list, + problem_text, + all_text, + manage, + step, + padding_problem_text, + reasoning_content=reasoning_content if reasoning_content_enable else "", + ) + if not manage.debug: + add_access_num(chat_user_id, chat_user_type, manage.context.get("application_id")) + yield manage.get_base_to_response().to_stream_chunk_response( + chat_id, + str(chat_record_id), + "ai-chat-node", + [], + all_text, + False, + 0, + 0, + { + "node_is_end": False, + "view_type": "many_view", + "node_type": "ai-chat-node", + "real_node_id": "ai-chat-node", + "reasoning_content": "", + }, + ) + + +class BaseChatStep(IChatStep): + def execute( + self, + message_list: List[BaseMessage], + chat_id, + problem_text, + post_response_handler: PostResponseHandler, + model_id: str = None, + workspace_id: str = None, + paragraph_list=None, + manage: PipelineManage = None, + padding_problem_text: str = None, + stream: bool = True, + chat_user_id=None, + chat_user_type=None, + no_references_setting=None, + model_params_setting=None, + model_setting=None, + mcp_tool_ids=None, + mcp_servers="", + mcp_source="referencing", + tool_ids=None, + application_ids=None, + skill_tool_ids=None, + mcp_output_enable=True, + **kwargs, + ): + chat_model = ( + get_model_instance_by_model_workspace_id(model_id, workspace_id, **(model_params_setting or {})) + if model_id is not None + else None + ) + if stream: + return self.execute_stream( + message_list, + chat_id, + problem_text, + post_response_handler, + chat_model, + paragraph_list, + manage, + padding_problem_text, + chat_user_id, + chat_user_type, + no_references_setting, + model_setting, + mcp_tool_ids, + mcp_servers, + mcp_source, + tool_ids, + application_ids, + skill_tool_ids, + workspace_id, + mcp_output_enable, + ) + else: + return self.execute_block( + message_list, + chat_id, + problem_text, + post_response_handler, + chat_model, + paragraph_list, + manage, + padding_problem_text, + chat_user_id, + chat_user_type, + no_references_setting, + model_setting, + mcp_tool_ids, + mcp_servers, + mcp_source, + tool_ids, + application_ids, + skill_tool_ids, + workspace_id, + mcp_output_enable, + ) + + def get_details(self, manage, **kwargs): + # 提取长期记忆 + extract_long_term_memory.apply_async( + args=( + manage.context.get("workspace_id"), + manage.context.get("application_id"), + manage.context.get("chat_user_id"), + ), + countdown=1, + ) + return { + "status": self.status, + "err_message": self.err_message, + "step_type": "chat_step", + "run_time": self.context.get("run_time") or 0, + "model_id": str(manage.context["model_id"]), + "message_list": self.reset_message_list( + self.context["step_args"].get("message_list"), self.context.get("answer_text") + ), + "message_tokens": self.context.get("message_tokens"), + "answer_tokens": self.context.get("answer_tokens"), + "cost": 0, + } + + @staticmethod + def reset_message_list(message_list: List[BaseMessage], answer_text): + result = [ + { + "role": "user" + if isinstance(message, HumanMessage) + else ("system" if isinstance(message, SystemMessage) else "ai"), + "content": message.content, + } + for message in message_list + ] + result.append({"role": "ai", "content": answer_text}) + return result + + def _handle_mcp_request( + self, + mcp_source, + mcp_servers, + mcp_tool_ids, + tool_ids, + application_ids, + skill_tool_ids, + mcp_output_enable, + chat_model, + system_prompt, + message_list, + agent_id, + chat_id, + workspace_id, + runtime_user_id=None, + ): + + mcp_servers_config = {} + + # 迁移过来mcp_source是None + if mcp_source is None: + mcp_source = "custom" + # 兼容老数据 + if not mcp_tool_ids: + mcp_tool_ids = [] + if mcp_source == "custom" and mcp_servers: + mcp_servers_config = json.loads(mcp_servers) + elif mcp_tool_ids: + mcp_tools = QuerySet(Tool).filter(id__in=mcp_tool_ids).values() + for mcp_tool in mcp_tools: + if mcp_tool and mcp_tool["is_active"]: + mcp_tool_config = json.loads(mcp_tool["code"]) + # 引用多个 MCP 工具时服务名是用户自定义的,可能重名;直接合并会 + # 静默覆盖前面的服务,只有最后一个生效。这里显式报错指出冲突服务名。 + conflict_servers = set(mcp_servers_config) & set(mcp_tool_config) + if conflict_servers: + raise AppApiException( + 500, + _( + "MCP server 【{servers}】 is defined in multiple referenced MCP tools, rename it so every server takes effect" + ).format(servers="、".join(sorted(conflict_servers))), + ) + mcp_servers_config = {**mcp_servers_config, **mcp_tool_config} + # 校验代码是否包括禁止的关键字 + ToolExecutor().validate_mcp_transport(json.dumps(mcp_servers_config)) + + tool_init_params = {} + tools = get_tools("APPLICATION", agent_id, tool_ids, workspace_id, runtime_user_id) + if tool_ids and len(tool_ids) > 0: # 如果有工具ID,则将其转换为MCP + self.context["tool_ids"] = tool_ids + for tool_id in tool_ids: + tool = QuerySet(Tool).filter(id=tool_id, tool_type=ToolType.CUSTOM).first() + if tool is None or tool.is_active is False: + continue + executor = ToolExecutor() + init_params_default_value = {i["field"]: i.get('default_value') for i in tool.init_field_list} + if tool.init_params is not None: + tool_init_params = init_params_default_value | json.loads(rsa_long_decrypt(tool.init_params)) + else: + tool_init_params = init_params_default_value + tool_config = executor.get_tool_mcp_config(tool, tool_init_params) + + mcp_servers_config[str(tool.id)] = tool_config + + if application_ids and len(application_ids) > 0: + self.context["application_ids"] = application_ids + for application_id in application_ids: + app = QuerySet(Application).filter(id=application_id, is_publish=True).first() + if app is None: + continue + app_key = QuerySet(ApplicationApiKey).filter(application_id=application_id, is_active=True).first() + if app_key is not None: + api_key = app_key.secret_key + application_access_token = ( + QuerySet(ApplicationAccessToken).filter(application_id=app_key.application_id).first() + ) + if application_access_token is not None and application_access_token.authentication: + raise AppApiException( + 500, + _("Agent 【{name}】 access token authentication is not supported for agent tool").format( + name=app.name + ), + ) + else: + raise AppApiException( + 500, _("Agent Key is required for agent tool 【{name}】").format(name=app.name) + ) + executor = ToolExecutor() + app_config = executor.get_app_mcp_config(api_key) + mcp_servers_config[app.name] = app_config + + if skill_tool_ids and len(skill_tool_ids) > 0: + self.context["skill_tool_ids"] = skill_tool_ids + skill_file_items = [] + + for tool_id in skill_tool_ids: + tool = QuerySet(Tool).filter(id=tool_id, is_active=True).first() + if tool is None or tool.is_active is False: + continue + init_params_default_value = {i["field"]: i.get("default_value") for i in tool.init_field_list} + if tool.init_params is not None: + params = init_params_default_value | json.loads(rsa_long_decrypt(tool.init_params)) + else: + params = init_params_default_value + + skill_file_items.append({"tool_id": str(tool.id), "file_id": tool.code, "params": params}) + mcp_servers_config["skills"] = skill_file_items + + if len(mcp_servers_config) > 0 or len(tools) > 0: + source_id = agent_id + source_type = "APPLICATION" + return mcp_response_generator( + chat_model, + system_prompt, + message_list, + mcp_servers_config, + mcp_output_enable, + tool_init_params, + source_id, + source_type, + chat_id, + tools, + ) + + return None + + def get_stream_result( + self, + message_list: List[BaseMessage], + chat_model: BaseChatModel = None, + paragraph_list=None, + no_references_setting=None, + problem_text=None, + mcp_tool_ids=None, + mcp_servers="", + mcp_source="referencing", + tool_ids=None, + application_ids=None, + skill_tool_ids=None, + workspace_id=None, + mcp_output_enable=True, + agent_id=None, + chat_id=None, + chat_user_id=None, + chat_user_type=None, + ): + if paragraph_list is None: + paragraph_list = [] + directly_return_chunk_list = [ + AIMessageChunk(content=paragraph.content) + for paragraph in paragraph_list + if ( + paragraph.hit_handling_method == "directly_return" + and paragraph.similarity >= paragraph.directly_return_similarity + ) + ] + if directly_return_chunk_list is not None and len(directly_return_chunk_list) > 0: + return iter(directly_return_chunk_list), False + elif len(paragraph_list) == 0 and no_references_setting.get("status") == "designated_answer": + return iter( + [AIMessageChunk(content=no_references_setting.get("value").replace("{question}", problem_text))] + ), False + if chat_model is None: + return iter( + [ + AIMessageChunk( + _( + "Sorry, the AI model is not configured. Please go to the application to set up the AI model first." + ) + ) + ] + ), False + else: + user_system_prompt = None + filtered_message_list = [] + long_term_memory = ( + QuerySet(ApplicationLongTermMemory).filter(chat_user_id=chat_user_id, application_id=agent_id).first() + ) + if long_term_memory is not None: + memory = long_term_memory.memory + else: + memory = "" + + # print(chat_user_id, chat_user_type) + for msg in message_list: + if isinstance(msg, SystemMessage): + if isinstance(msg.content, str): + user_system_prompt = msg.content.replace("{memory}", memory) + msg.content = user_system_prompt + elif isinstance(msg.content, list): + user_system_prompt = "".join( + item.get("text", "") if isinstance(item, dict) else str(item) for item in msg.content + ) + else: + user_system_prompt = str(msg.content) + else: + filtered_message_list.append(msg) + runtime_user_id = get_runtime_user_id(chat_user_id=chat_user_id, chat_user_type=chat_user_type) + # 过滤tool_id + all_tool_ids = list(set((mcp_tool_ids or []) + (tool_ids or []) + (skill_tool_ids or []))) + authorized_set = set(filter_authorized_ids("tool", all_tool_ids, workspace_id, user_id=runtime_user_id)) + + mcp_tool_ids = [i for i in (mcp_tool_ids or []) if i in authorized_set] + tool_ids = [i for i in (tool_ids or []) if i in authorized_set] + skill_tool_ids = [i for i in (skill_tool_ids or []) if i in authorized_set] + # 处理 MCP 请求 + mcp_result = self._handle_mcp_request( + mcp_source, + mcp_servers, + mcp_tool_ids, + tool_ids, + application_ids, + skill_tool_ids, + mcp_output_enable, + chat_model, + user_system_prompt, + filtered_message_list, + agent_id, + chat_id, + workspace_id, + runtime_user_id, + ) + if mcp_result: + return mcp_result, True + return chat_model.stream(message_list), True + + def execute_stream( + self, + message_list: List[BaseMessage], + chat_id, + problem_text, + post_response_handler: PostResponseHandler, + chat_model: BaseChatModel = None, + paragraph_list=None, + manage: PipelineManage = None, + padding_problem_text: str = None, + chat_user_id=None, + chat_user_type=None, + no_references_setting=None, + model_setting=None, + mcp_tool_ids=None, + mcp_servers="", + mcp_source="referencing", + tool_ids=None, + application_ids=None, + skill_tool_ids=None, + workspace_id=None, + mcp_output_enable=True, + ): + chat_result, is_ai_chat = self.get_stream_result( + message_list, + chat_model, + paragraph_list, + no_references_setting, + problem_text, + mcp_tool_ids, + mcp_servers, + mcp_source, + tool_ids, + application_ids, + skill_tool_ids, + workspace_id, + mcp_output_enable, + manage.context.get("application_id"), + chat_id, + chat_user_id, + chat_user_type, + ) + chat_record_id = ( + self.context.get("step_args", {}).get("chat_record_id") + if self.context.get("step_args", {}).get("chat_record_id") + else uuid.uuid7() + ) + r = StreamingHttpResponse( + streaming_content=event_content( + chat_result, + chat_id, + chat_record_id, + paragraph_list, + post_response_handler, + manage, + self, + chat_model, + message_list, + problem_text, + padding_problem_text, + chat_user_id, + chat_user_type, + is_ai_chat, + model_setting, + ), + content_type="text/event-stream;charset=utf-8", + ) + + r["Cache-Control"] = "no-cache" + return r + + def get_block_result( + self, + message_list: List[BaseMessage], + chat_model: BaseChatModel = None, + paragraph_list=None, + no_references_setting=None, + problem_text=None, + mcp_tool_ids=None, + mcp_servers="", + mcp_source="referencing", + tool_ids=None, + application_ids=None, + skill_tool_ids=None, + workspace_id=None, + mcp_output_enable=True, + application_id=None, + chat_id=None, + chat_user_id=None, + chat_user_type=None, + ): + if paragraph_list is None: + paragraph_list = [] + directly_return_chunk_list = [ + AIMessageChunk(content=paragraph.content) + for paragraph in paragraph_list + if ( + paragraph.hit_handling_method == "directly_return" + and paragraph.similarity >= paragraph.directly_return_similarity + ) + ] + if directly_return_chunk_list is not None and len(directly_return_chunk_list) > 0: + return directly_return_chunk_list[0], False + elif len(paragraph_list) == 0 and no_references_setting.get("status") == "designated_answer": + return AIMessage(no_references_setting.get("value").replace("{question}", problem_text)), False + if chat_model is None: + return AIMessage( + _("Sorry, the AI model is not configured. Please go to the application to set up the AI model first.") + ), False + else: + runtime_user_id = get_runtime_user_id(chat_user_id=chat_user_id, chat_user_type=chat_user_type) + # 过滤tool_id + all_tool_ids = list(set((mcp_tool_ids or []) + (tool_ids or []) + (skill_tool_ids or []))) + authorized_set = set(filter_authorized_ids("tool", all_tool_ids, workspace_id, user_id=runtime_user_id)) + + mcp_tool_ids = [i for i in (mcp_tool_ids or []) if i in authorized_set] + tool_ids = [i for i in (tool_ids or []) if i in authorized_set] + skill_tool_ids = [i for i in (skill_tool_ids or []) if i in authorized_set] + # 处理 MCP 请求 + mcp_result = self._handle_mcp_request( + mcp_source, + mcp_servers, + mcp_tool_ids, + tool_ids, + application_ids, + skill_tool_ids, + mcp_output_enable, + chat_model, + "", + message_list, + application_id, + chat_id, + workspace_id, + runtime_user_id, + ) + if mcp_result: + return mcp_result, True + return chat_model.invoke(message_list), True + + def execute_block( + self, + message_list: List[BaseMessage], + chat_id, + problem_text, + post_response_handler: PostResponseHandler, + chat_model: BaseChatModel = None, + paragraph_list=None, + manage: PipelineManage = None, + padding_problem_text: str = None, + chat_user_id=None, + chat_user_type=None, + no_references_setting=None, + model_setting=None, + mcp_tool_ids=None, + mcp_servers="", + mcp_source="referencing", + tool_ids=None, + application_ids=None, + skill_tool_ids=None, + workspace_id=None, + mcp_output_enable=True, + ): + reasoning_content_enable = model_setting.get("reasoning_content_enable", False) + reasoning_content_start = model_setting.get("reasoning_content_start", "") + reasoning_content_end = model_setting.get("reasoning_content_end", "") + reasoning = Reasoning(reasoning_content_start, reasoning_content_end) + chat_record_id = uuid.uuid7() + # 调用模型 + try: + chat_result, is_ai_chat = self.get_block_result( + message_list, + chat_model, + paragraph_list, + no_references_setting, + problem_text, + mcp_tool_ids, + mcp_servers, + mcp_source, + tool_ids, + application_ids, + skill_tool_ids, + workspace_id, + mcp_output_enable, + manage.context.get("application_id"), + chat_id, + chat_user_id, + chat_user_type, + ) + if is_ai_chat: + request_token = chat_model.get_num_tokens_from_messages(message_list) + response_token = chat_model.get_num_tokens(chat_result.content) + else: + request_token = 0 + response_token = 0 + write_context(self, manage, request_token, response_token, chat_result.content) + reasoning_result = reasoning.get_reasoning_content(chat_result) + reasoning_result_end = reasoning.get_end_reasoning_content() + content = reasoning_result.get("content") + reasoning_result_end.get("content") + if "reasoning_content" in chat_result.response_metadata: + reasoning_content = chat_result.response_metadata.get("reasoning_content", "") or "" + else: + reasoning_content = (reasoning_result.get("reasoning_content") or "") + ( + reasoning_result_end.get("reasoning_content") or "" + ) + post_response_handler.handler( + chat_id, + chat_record_id, + paragraph_list, + problem_text, + content, + manage, + self, + padding_problem_text, + reasoning_content=reasoning_content, + ) + if not manage.debug: + add_access_num(chat_user_id, chat_user_type, manage.context.get("application_id")) + return manage.get_base_to_response().to_block_response( + str(chat_id), + str(chat_record_id), + content, + True, + request_token, + response_token, + { + "reasoning_content": reasoning_content if reasoning_content_enable else "", + "answer_list": [ + {"content": content, "reasoning_content": reasoning_content if reasoning_content_enable else ""} + ], + }, + ) + except Exception as e: + all_text = "Exception:" + str(e) + write_context(self, manage, 0, 0, all_text) + post_response_handler.handler( + chat_id, + chat_record_id, + paragraph_list, + problem_text, + all_text, + manage, + self, + padding_problem_text, + reasoning_content="", + ) + if not manage.debug: + add_access_num(chat_user_id, chat_user_type, manage.context.get("application_id")) + return manage.get_base_to_response().to_block_response( + str(chat_id), str(chat_record_id), all_text, True, 0, 0, _status=status.HTTP_500_INTERNAL_SERVER_ERROR + ) 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 2b2f961352d..033289dfedc 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 @@ -1,640 +1,651 @@ -# coding=utf-8 -""" -@project: maxkb -@Author:虎 -@file: base_question_node.py -@date:2024/6/4 14:30 -@desc: -""" - -import base64 -import json -import re -import time -from functools import reduce -from imghdr import what -from typing import Dict, List - -from common.exception.app_exception import AppApiException -from common.utils.rsa_util import rsa_long_decrypt -from common.utils.shared_resource_auth import filter_authorized_ids, get_runtime_user_id -from common.utils.tool_code import ToolExecutor -from django.db.models import QuerySet -from django.utils.translation import gettext as _ -from knowledge.models import File -from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage -from models_provider.models import Model -from models_provider.tools import get_model_credential, get_model_instance_by_model_workspace_id -from tools.models import Tool, ToolType - -from application.flow.common import WorkflowMode -from application.flow.i_step_node import INode, NodeResult -from application.flow.step_node.ai_chat_step_node.i_chat_node import IChatNode -from application.flow.tools import Reasoning, get_tools, mcp_response_generator -from application.models import Application, ApplicationAccessToken, ApplicationApiKey - - -def _write_context( - node_variable: Dict, workflow_variable: Dict, node: INode, workflow, answer: str, reasoning_content: str -): - chat_model = node_variable.get("chat_model") - message_tokens = chat_model.get_num_tokens_from_messages(node_variable.get("message_list")) - answer_tokens = chat_model.get_num_tokens(answer) - node.context["message_tokens"] = message_tokens - node.context["answer_tokens"] = answer_tokens - node.context["answer"] = answer - node.context["question"] = node_variable["question"] - node.context["run_time"] = time.time() - node.context["start_time"] - node.context["reasoning_content"] = reasoning_content - if workflow.is_result(node, NodeResult(node_variable, workflow_variable)): - node.answer_text = answer - - -def write_context_stream(node_variable: Dict, workflow_variable: Dict, node: INode, workflow): - """ - 写入上下文数据 (流式) - @param node_variable: 节点数据 - @param workflow_variable: 全局数据 - @param node: 节点 - @param workflow: 工作流管理器 - """ - response = node_variable.get("result") - answer = "" - reasoning_content = "" - model_setting = node.context.get( - "model_setting", - {"reasoning_content_enable": False, "reasoning_content_end": "", "reasoning_content_start": ""}, - ) - reasoning = Reasoning( - model_setting.get("reasoning_content_start", ""), model_setting.get("reasoning_content_end", "") - ) - response_reasoning_content = False - - for chunk in response: - if workflow.is_the_task_interrupted(): - break - reasoning_chunk = reasoning.get_reasoning_content(chunk) - content_chunk = reasoning_chunk.get("content") - if "reasoning_content" in chunk.additional_kwargs: - response_reasoning_content = True - reasoning_content_chunk = chunk.additional_kwargs.get("reasoning_content", "") - else: - reasoning_content_chunk = reasoning_chunk.get("reasoning_content") - answer += content_chunk - if reasoning_content_chunk is None: - reasoning_content_chunk = "" - reasoning_content += reasoning_content_chunk - yield { - "content": content_chunk, - "reasoning_content": reasoning_content_chunk - if model_setting.get("reasoning_content_enable", False) - else "", - } - - reasoning_chunk = reasoning.get_end_reasoning_content() - answer += reasoning_chunk.get("content") - reasoning_content_chunk = "" - if not response_reasoning_content: - reasoning_content_chunk = reasoning_chunk.get("reasoning_content") - yield { - "content": reasoning_chunk.get("content"), - "reasoning_content": reasoning_content_chunk if model_setting.get("reasoning_content_enable", False) else "", - } - _write_context(node_variable, workflow_variable, node, workflow, answer, reasoning_content) - - -def write_context(node_variable: Dict, workflow_variable: Dict, node: INode, workflow): - """ - 写入上下文数据 - @param node_variable: 节点数据 - @param workflow_variable: 全局数据 - @param node: 节点实例对象 - @param workflow: 工作流管理器 - """ - response = node_variable.get("result") - model_setting = node.context.get( - "model_setting", - {"reasoning_content_enable": False, "reasoning_content_end": "", "reasoning_content_start": ""}, - ) - reasoning = Reasoning(model_setting.get("reasoning_content_start"), model_setting.get("reasoning_content_end")) - reasoning_result = reasoning.get_reasoning_content(response) - reasoning_result_end = reasoning.get_end_reasoning_content() - content = reasoning_result.get("content") + reasoning_result_end.get("content") - meta = {**response.response_metadata, **response.additional_kwargs} - if "reasoning_content" in meta: - reasoning_content = meta.get("reasoning_content", "") or "" - else: - reasoning_content = (reasoning_result.get("reasoning_content") or "") + ( - reasoning_result_end.get("reasoning_content") or "" - ) - _write_context(node_variable, workflow_variable, node, workflow, content, reasoning_content) - - -CHAT_FILE_LIST_FIELDS = ("image_list", "document_list", "audio_list", "video_list", "other_list") - - -def get_default_model_params_setting(model_id): - model = QuerySet(Model).filter(id=model_id).first() - credential = get_model_credential(model.provider, model.model_type, model.model_name) - model_params_setting = credential.get_model_params_setting_form(model.model_name).get_default_form_data() - return model_params_setting - - -def get_node_message(chat_record, runtime_node_id): - node_details = chat_record.get_node_details_runtime_node_id(runtime_node_id) - if node_details is None: - return [] - return [HumanMessage(node_details.get("question")), AIMessage(node_details.get("answer"))] - - -def get_workflow_message(chat_record): - return [chat_record.get_human_message(), chat_record.get_ai_message()] - - -def get_message(chat_record, dialogue_type, runtime_node_id): - return ( - get_node_message(chat_record, runtime_node_id) if dialogue_type == "NODE" else get_workflow_message(chat_record) - ) - - -class BaseChatNode(IChatNode): - def save_context(self, details, workflow_manage): - self.context["answer"] = details.get("answer") - self.context["question"] = details.get("question") - self.context["reasoning_content"] = details.get("reasoning_content") - self.context["exception_message"] = details.get("err_message") - if self.node_params.get("is_result", False): - self.answer_text = details.get("answer") - - def execute( - self, - model_id, - system, - prompt, - dialogue_number, - history_chat_record, - stream, - chat_id, - chat_record_id, - model_params_setting=None, - model_id_type=None, - model_id_reference=None, - dialogue_type=None, - model_setting=None, - mcp_servers=None, - mcp_tool_id=None, - mcp_tool_ids=None, - mcp_source=None, - tool_ids=None, - application_ids=None, - skill_tool_ids=None, - mcp_output_enable=True, - **kwargs, - ) -> NodeResult: - if dialogue_type is None: - dialogue_type = "WORKFLOW" - - if model_id_type == "reference" and model_id_reference: - reference_data = self.workflow_manage.get_reference_field( - model_id_reference[0], - model_id_reference[1:], - ) - - if reference_data and isinstance(reference_data, dict): - model_id = reference_data.get("model_id", model_id) - model_params_setting = reference_data.get("model_params_setting") - elif model_id_type == "default": - default_setting = self.workflow_manage.get_default_model_setting("LLM") - model_id = default_setting.get("model_id") - model_params_setting = default_setting.get("model_params_setting", model_params_setting) - if model_id is None or model_id == "": - raise Exception(_("Model is not allowed to be empty")) - - if model_params_setting is None and model_id: - model_params_setting = get_default_model_params_setting(model_id) - - if model_setting is None: - model_setting = { - "reasoning_content_enable": False, - "reasoning_content_end": "", - "reasoning_content_start": "", - } - self.context["model_setting"] = model_setting - workspace_id = self.workflow_manage.get_body().get("workspace_id") - chat_model = get_model_instance_by_model_workspace_id(model_id, workspace_id, **(model_params_setting or {})) - history_message = self.get_history_message( - history_chat_record, dialogue_number, dialogue_type, self.runtime_node_id - ) - self.context["history_message"] = [ - {"content": message.content, "role": message.type} - for message in (history_message if history_message is not None else []) - ] - question = self.generate_prompt_question(prompt, chat_model) - self.context["question"] = question.content - system = self.workflow_manage.generate_prompt(system) - self.context["system"] = system - message_list = self.generate_message_list(question, history_message) - self.context["message_list"] = message_list - body = self.workflow_manage.get_body() - runtime_user_id = get_runtime_user_id( - user_id=body.get("user_id"), - chat_user_id=body.get("chat_user_id"), - chat_user_type=body.get("chat_user_type"), - ) - - # 过滤tool_id - all_tool_ids = list( - set( - (mcp_tool_ids or []) - + (tool_ids or []) - + (skill_tool_ids or []) - + ([mcp_tool_id] if mcp_tool_id else []) - ) - ) - authorized_set = set(filter_authorized_ids("tool", all_tool_ids, workspace_id, user_id=runtime_user_id)) - - mcp_tool_ids = [i for i in (mcp_tool_ids or []) if i in authorized_set] - tool_ids = [i for i in (tool_ids or []) if i in authorized_set] - skill_tool_ids = [i for i in (skill_tool_ids or []) if i in authorized_set] - mcp_tool_id = mcp_tool_id if (mcp_tool_id and mcp_tool_id in authorized_set) else None - # 处理 MCP 请求 - mcp_result = self._handle_mcp_request( - mcp_source, - mcp_servers, - mcp_tool_id, - mcp_tool_ids, - tool_ids, - application_ids, - skill_tool_ids, - mcp_output_enable, - chat_model, - SystemMessage(system), - message_list, - history_message, - question, - chat_id, - workspace_id, - runtime_user_id, - ) - if mcp_result: - return mcp_result - message_list = [SystemMessage(system)] + message_list - if stream: - r = chat_model.stream(message_list) - return NodeResult( - {"result": r, "chat_model": chat_model, "message_list": message_list, "question": question.content}, - {}, - _write_context=write_context_stream, - ) - else: - r = chat_model.invoke(message_list) - return NodeResult( - { - "result": r, - "chat_model": chat_model, - "message_list": message_list, - "history_message": [ - {"content": message.content, "role": message.type} - for message in (history_message if history_message is not None else []) - ], - "question": question.content, - }, - {}, - _write_context=write_context, - ) - - def _handle_mcp_request( - self, - mcp_source, - mcp_servers, - mcp_tool_id, - mcp_tool_ids, - tool_ids, - application_ids, - skill_tool_ids, - mcp_output_enable, - chat_model, - system_prompt, - message_list, - history_message, - question, - chat_id, - workspace_id, - runtime_user_id=None, - ): - - mcp_servers_config = {} - - # 迁移过来mcp_source是None - if mcp_source is None: - mcp_source = "custom" - # 兼容老数据 - if not mcp_tool_ids: - mcp_tool_ids = [] - if mcp_tool_id: - mcp_tool_ids = list(set(mcp_tool_ids + [mcp_tool_id])) - if mcp_source == "custom" and mcp_servers: - mcp_servers_config = json.loads(mcp_servers) - mcp_servers_config = self.handle_variables(mcp_servers_config) - elif mcp_tool_ids: - mcp_tools = QuerySet(Tool).filter(id__in=mcp_tool_ids).values() - for mcp_tool in mcp_tools: - if mcp_tool and mcp_tool["is_active"]: - mcp_servers_config = {**mcp_servers_config, **json.loads(mcp_tool["code"])} - mcp_servers_config = self.handle_variables(mcp_servers_config) - # 校验代码是否包括禁止的关键字 - ToolExecutor().validate_mcp_transport(json.dumps(mcp_servers_config)) - - tool_init_params = {} - tools = get_tools( - self.workflow_manage.get_source_type(), - self.workflow_manage.get_source_id(), - tool_ids, - workspace_id, - runtime_user_id, - ) - if tool_ids and len(tool_ids) > 0: # 如果有工具ID,则将其转换为MCP - self.context["tool_ids"] = tool_ids - custom_tools_map = { - str(t.id): t for t in QuerySet(Tool).filter(id__in=tool_ids, tool_type=ToolType.CUSTOM, is_active=True) - } - for tool_id in tool_ids: - tool = custom_tools_map.get(str(tool_id)) - if tool is None: - continue - executor = ToolExecutor() - init_params_default_value = {i["field"]: i.get("default_value") for i in tool.init_field_list} - if tool.init_params is not None: - tool_init_params = init_params_default_value | json.loads(rsa_long_decrypt(tool.init_params)) - else: - tool_init_params = init_params_default_value - - tool_config = executor.get_tool_mcp_config(tool, tool_init_params) - - mcp_servers_config[str(tool.id)] = tool_config - - if application_ids and len(application_ids) > 0: - self.context["application_ids"] = application_ids - apps_map = {str(a.id): a for a in QuerySet(Application).filter(id__in=application_ids, is_publish=True)} - app_keys_map = { - str(ak.application_id): ak - for ak in QuerySet(ApplicationApiKey).filter(application_id__in=application_ids, is_active=True) - } - app_access_tokens_map = { - str(at.application_id): at - for at in QuerySet(ApplicationAccessToken).filter(application_id__in=application_ids) - } - for application_id in application_ids: - app = apps_map.get(str(application_id)) - if app is None: - continue - app_key = app_keys_map.get(str(application_id)) - if app_key is not None: - api_key = app_key.secret_key - application_access_token = app_access_tokens_map.get(str(app_key.application_id)) - if application_access_token is not None and application_access_token.authentication: - raise AppApiException( - 500, - _("Agent 【{name}】 access token authentication is not supported for agent tool").format( - name=app.name - ), - ) - else: - raise AppApiException( - 500, _("Agent Key is required for agent tool 【{name}】").format(name=app.name) - ) - executor = ToolExecutor() - app_config = executor.get_app_mcp_config(api_key, self.get_chat_files(), self.get_form_data()) - mcp_servers_config[app.name] = app_config - - if skill_tool_ids and len(skill_tool_ids) > 0: - self.context["skill_tool_ids"] = skill_tool_ids - skill_file_items = [] - skill_tools_map = {str(t.id): t for t in QuerySet(Tool).filter(id__in=skill_tool_ids, is_active=True)} - for tool_id in skill_tool_ids: - tool = skill_tools_map.get(str(tool_id)) - if tool is None: - continue - init_params_default_value = {i["field"]: i.get("default_value") for i in tool.init_field_list} - if tool.init_params is not None: - params = init_params_default_value | json.loads(rsa_long_decrypt(tool.init_params)) - else: - params = init_params_default_value - - skill_file_items.append({"tool_id": str(tool.id), "file_id": tool.code, "params": params}) - mcp_servers_config["skills"] = skill_file_items - - if len(mcp_servers_config) > 0 or len(tools) > 0: - # 安全获取 application - application_id = None - tool_id = None - knowledge_id = None - if [WorkflowMode.KNOWLEDGE, WorkflowMode.KNOWLEDGE_LOOP].__contains__( - self.workflow_manage.flow.workflow_mode - ): - knowledge_id = self.workflow_params.get("knowledge_id") - elif [WorkflowMode.APPLICATION, WorkflowMode.APPLICATION_LOOP].__contains__( - self.workflow_manage.flow.workflow_mode - ): - application_id = self.workflow_manage.work_flow_post_handler.chat_info.application.id - elif [WorkflowMode.TOOL, WorkflowMode.TOOL_LOOP].__contains__(self.workflow_manage.flow.workflow_mode): - tool_id = self.workflow_params.get("tool_id") - - source_id = application_id or knowledge_id or tool_id - source_type = "APPLICATION" if application_id else "KNOWLEDGE" if knowledge_id else "TOOL" - r = mcp_response_generator( - chat_model, - system_prompt, - message_list, - mcp_servers_config, - mcp_output_enable, - tool_init_params, - source_id, - source_type, - chat_id, - tools, - ) - return NodeResult( - { - "result": r, - "chat_model": chat_model, - "message_list": message_list, - "history_message": [ - {"content": message.content, "role": message.type} - for message in (history_message if history_message is not None else []) - ], - "question": question.content, - }, - {}, - _write_context=write_context_stream, - ) - - return None - - def get_chat_files(self): - """ - 获取本次对话上传的文件, 用于透传给被当作工具调用的应用/MCP - """ - chat_files = {} - for field in CHAT_FILE_LIST_FIELDS: - file_list = getattr(self.workflow_manage, field, None) or [] - items = [ - {key: item.get(key) for key in ("name", "url", "file_id") if item.get(key) is not None} - for item in file_list - if isinstance(item, dict) - ] - if items: - chat_files[field] = items - return chat_files - - def get_form_data(self): - """ - 获取当前会话的用户输入参数,用于透传给作为工具调用的子智能体。 - - 循环工作流会创建独立的 WorkflowManage,并将自身的 form_data 初始化为 - 空字典,因此需要继续从父工作流中查找原始用户输入。 - """ - workflow_manage = self.workflow_manage - visited = set() - while workflow_manage is not None and id(workflow_manage) not in visited: - visited.add(id(workflow_manage)) - form_data = getattr(workflow_manage, "form_data", None) - if isinstance(form_data, dict) and form_data: - return form_data.copy() - workflow_manage = getattr(workflow_manage, "parentWorkflowManage", None) - return {} - - def handle_variables(self, tool_params): - # 处理参数中的变量 - for k, v in tool_params.items(): - if type(v) == str: - tool_params[k] = self.workflow_manage.generate_prompt(tool_params[k]) - elif type(v) == dict: - self.handle_variables(v) - elif (type(v) == list) and len(v) > 0 and (type(v[0]) == str): - tool_params[k] = self.get_reference_content(v) - return tool_params - - def get_reference_content(self, fields: List[str]): - return str(self.workflow_manage.get_reference_field(fields[0], fields[1:])) if fields else "" - - @staticmethod - def get_history_message(history_chat_record, dialogue_number, dialogue_type, runtime_node_id): - start_index = len(history_chat_record) - dialogue_number - history_message = reduce( - lambda x, y: [*x, *y], - [ - get_message(history_chat_record[index], dialogue_type, runtime_node_id) - for index in range(start_index if start_index > 0 else 0, len(history_chat_record)) - ], - [], - ) - for message in history_message: - if isinstance(message.content, str): - message.content = re.sub(r".*?<\/form_rander>", "", message.content, flags=re.DOTALL) - return history_message - - def generate_prompt_question(self, prompt, model): - image = self.get_image() - video = self.get_video() - vision = self.is_vision() - videos = [] - images = [] - if image and vision: - images = self._process_images(image) - if video and vision: - videos = self._process_videos(video, model) - prompt = self.workflow_manage.generate_prompt(prompt) - if images or videos: - return HumanMessage(content=[*videos, *images, {"type": "text", "text": prompt}]) - return HumanMessage(content=prompt) - - def is_vision(self): - if "vision" in self.node_params_serializer.data: - return self.node_params_serializer.data.get("vision") - return False - - def get_image(self): - if "image_list" in self.node_params_serializer.data: - image = self.workflow_manage.get_reference_field( - self.node_params_serializer.data.get("image_list")[0], - self.node_params_serializer.data.get("image_list")[1:], - ) - return image - return None - - def get_video(self): - if "video_list" in self.node_params_serializer.data: - video = self.workflow_manage.get_reference_field( - self.node_params_serializer.data.get("video_list")[0], - self.node_params_serializer.data.get("video_list")[1:], - ) - return video - return None - - def _process_videos(self, image, video_model): - videos = [] - 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: - 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}}) - elif "url" in img and img["url"].startswith("http"): - videos.append({"type": "video_url", "video_url": {"url": img["url"]}}) - return videos - - def _process_images(self, image): - """ - 处理图像数据,转换为模型可识别的格式 - """ - images = [] - 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: - 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}"}} - ) - elif "url" in img and img["url"].startswith("http"): - images.append({"type": "image_url", "image_url": {"url": img["url"]}}) - return images - - def generate_message_list(self, question, history_message): - return [*history_message, question] - - @staticmethod - def reset_message_list(message_list: List[BaseMessage], answer_text): - result = [ - {"role": "user" if isinstance(message, HumanMessage) else "ai", "content": message.content} - for message in message_list - ] - result.append({"role": "ai", "content": answer_text}) - return result - - def get_details(self, index: int, **kwargs): - return { - "name": self.node.properties.get("stepName"), - "index": index, - "run_time": self.context.get("run_time"), - "system": self.context.get("system"), - "history_message": self.context.get("history_message"), - "question": self.context.get("question"), - "answer": self.context.get("answer"), - "reasoning_content": self.context.get("reasoning_content"), - "enableException": self.node.properties.get("enableException"), - "type": self.node.type, - "message_tokens": self.context.get("message_tokens"), - "answer_tokens": self.context.get("answer_tokens"), - "status": self.status, - "err_message": self.err_message, - } +# coding=utf-8 +""" +@project: maxkb +@Author:虎 +@file: base_question_node.py +@date:2024/6/4 14:30 +@desc: +""" + +import base64 +import json +import re +import time +from functools import reduce +from imghdr import what +from typing import Dict, List + +from common.exception.app_exception import AppApiException +from common.utils.rsa_util import rsa_long_decrypt +from common.utils.shared_resource_auth import filter_authorized_ids, get_runtime_user_id +from common.utils.tool_code import ToolExecutor +from django.db.models import QuerySet +from django.utils.translation import gettext as _ +from knowledge.models import File +from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage +from models_provider.models import Model +from models_provider.tools import get_model_credential, get_model_instance_by_model_workspace_id +from tools.models import Tool, ToolType + +from application.flow.common import WorkflowMode +from application.flow.i_step_node import INode, NodeResult +from application.flow.step_node.ai_chat_step_node.i_chat_node import IChatNode +from application.flow.tools import Reasoning, get_tools, mcp_response_generator +from application.models import Application, ApplicationAccessToken, ApplicationApiKey + + +def _write_context( + node_variable: Dict, workflow_variable: Dict, node: INode, workflow, answer: str, reasoning_content: str +): + chat_model = node_variable.get("chat_model") + message_tokens = chat_model.get_num_tokens_from_messages(node_variable.get("message_list")) + answer_tokens = chat_model.get_num_tokens(answer) + node.context["message_tokens"] = message_tokens + node.context["answer_tokens"] = answer_tokens + node.context["answer"] = answer + node.context["question"] = node_variable["question"] + node.context["run_time"] = time.time() - node.context["start_time"] + node.context["reasoning_content"] = reasoning_content + if workflow.is_result(node, NodeResult(node_variable, workflow_variable)): + node.answer_text = answer + + +def write_context_stream(node_variable: Dict, workflow_variable: Dict, node: INode, workflow): + """ + 写入上下文数据 (流式) + @param node_variable: 节点数据 + @param workflow_variable: 全局数据 + @param node: 节点 + @param workflow: 工作流管理器 + """ + response = node_variable.get("result") + answer = "" + reasoning_content = "" + model_setting = node.context.get( + "model_setting", + {"reasoning_content_enable": False, "reasoning_content_end": "", "reasoning_content_start": ""}, + ) + reasoning = Reasoning( + model_setting.get("reasoning_content_start", ""), model_setting.get("reasoning_content_end", "") + ) + response_reasoning_content = False + + for chunk in response: + if workflow.is_the_task_interrupted(): + break + reasoning_chunk = reasoning.get_reasoning_content(chunk) + content_chunk = reasoning_chunk.get("content") + if "reasoning_content" in chunk.additional_kwargs: + response_reasoning_content = True + reasoning_content_chunk = chunk.additional_kwargs.get("reasoning_content", "") + else: + reasoning_content_chunk = reasoning_chunk.get("reasoning_content") + answer += content_chunk + if reasoning_content_chunk is None: + reasoning_content_chunk = "" + reasoning_content += reasoning_content_chunk + yield { + "content": content_chunk, + "reasoning_content": reasoning_content_chunk + if model_setting.get("reasoning_content_enable", False) + else "", + } + + reasoning_chunk = reasoning.get_end_reasoning_content() + answer += reasoning_chunk.get("content") + reasoning_content_chunk = "" + if not response_reasoning_content: + reasoning_content_chunk = reasoning_chunk.get("reasoning_content") + yield { + "content": reasoning_chunk.get("content"), + "reasoning_content": reasoning_content_chunk if model_setting.get("reasoning_content_enable", False) else "", + } + _write_context(node_variable, workflow_variable, node, workflow, answer, reasoning_content) + + +def write_context(node_variable: Dict, workflow_variable: Dict, node: INode, workflow): + """ + 写入上下文数据 + @param node_variable: 节点数据 + @param workflow_variable: 全局数据 + @param node: 节点实例对象 + @param workflow: 工作流管理器 + """ + response = node_variable.get("result") + model_setting = node.context.get( + "model_setting", + {"reasoning_content_enable": False, "reasoning_content_end": "", "reasoning_content_start": ""}, + ) + reasoning = Reasoning(model_setting.get("reasoning_content_start"), model_setting.get("reasoning_content_end")) + reasoning_result = reasoning.get_reasoning_content(response) + reasoning_result_end = reasoning.get_end_reasoning_content() + content = reasoning_result.get("content") + reasoning_result_end.get("content") + meta = {**response.response_metadata, **response.additional_kwargs} + if "reasoning_content" in meta: + reasoning_content = meta.get("reasoning_content", "") or "" + else: + reasoning_content = (reasoning_result.get("reasoning_content") or "") + ( + reasoning_result_end.get("reasoning_content") or "" + ) + _write_context(node_variable, workflow_variable, node, workflow, content, reasoning_content) + + +CHAT_FILE_LIST_FIELDS = ("image_list", "document_list", "audio_list", "video_list", "other_list") + + +def get_default_model_params_setting(model_id): + model = QuerySet(Model).filter(id=model_id).first() + credential = get_model_credential(model.provider, model.model_type, model.model_name) + model_params_setting = credential.get_model_params_setting_form(model.model_name).get_default_form_data() + return model_params_setting + + +def get_node_message(chat_record, runtime_node_id): + node_details = chat_record.get_node_details_runtime_node_id(runtime_node_id) + if node_details is None: + return [] + return [HumanMessage(node_details.get("question")), AIMessage(node_details.get("answer"))] + + +def get_workflow_message(chat_record): + return [chat_record.get_human_message(), chat_record.get_ai_message()] + + +def get_message(chat_record, dialogue_type, runtime_node_id): + return ( + get_node_message(chat_record, runtime_node_id) if dialogue_type == "NODE" else get_workflow_message(chat_record) + ) + + +class BaseChatNode(IChatNode): + def save_context(self, details, workflow_manage): + self.context["answer"] = details.get("answer") + self.context["question"] = details.get("question") + self.context["reasoning_content"] = details.get("reasoning_content") + self.context["exception_message"] = details.get("err_message") + if self.node_params.get("is_result", False): + self.answer_text = details.get("answer") + + def execute( + self, + model_id, + system, + prompt, + dialogue_number, + history_chat_record, + stream, + chat_id, + chat_record_id, + model_params_setting=None, + model_id_type=None, + model_id_reference=None, + dialogue_type=None, + model_setting=None, + mcp_servers=None, + mcp_tool_id=None, + mcp_tool_ids=None, + mcp_source=None, + tool_ids=None, + application_ids=None, + skill_tool_ids=None, + mcp_output_enable=True, + **kwargs, + ) -> NodeResult: + if dialogue_type is None: + dialogue_type = "WORKFLOW" + + if model_id_type == "reference" and model_id_reference: + reference_data = self.workflow_manage.get_reference_field( + model_id_reference[0], + model_id_reference[1:], + ) + + if reference_data and isinstance(reference_data, dict): + model_id = reference_data.get("model_id", model_id) + model_params_setting = reference_data.get("model_params_setting") + elif model_id_type == "default": + default_setting = self.workflow_manage.get_default_model_setting("LLM") + model_id = default_setting.get("model_id") + model_params_setting = default_setting.get("model_params_setting", model_params_setting) + if model_id is None or model_id == "": + raise Exception(_("Model is not allowed to be empty")) + + if model_params_setting is None and model_id: + model_params_setting = get_default_model_params_setting(model_id) + + if model_setting is None: + model_setting = { + "reasoning_content_enable": False, + "reasoning_content_end": "", + "reasoning_content_start": "", + } + self.context["model_setting"] = model_setting + workspace_id = self.workflow_manage.get_body().get("workspace_id") + chat_model = get_model_instance_by_model_workspace_id(model_id, workspace_id, **(model_params_setting or {})) + history_message = self.get_history_message( + history_chat_record, dialogue_number, dialogue_type, self.runtime_node_id + ) + self.context["history_message"] = [ + {"content": message.content, "role": message.type} + for message in (history_message if history_message is not None else []) + ] + question = self.generate_prompt_question(prompt, chat_model) + self.context["question"] = question.content + system = self.workflow_manage.generate_prompt(system) + self.context["system"] = system + message_list = self.generate_message_list(question, history_message) + self.context["message_list"] = message_list + body = self.workflow_manage.get_body() + runtime_user_id = get_runtime_user_id( + user_id=body.get("user_id"), + chat_user_id=body.get("chat_user_id"), + chat_user_type=body.get("chat_user_type"), + ) + + # 过滤tool_id + all_tool_ids = list( + set( + (mcp_tool_ids or []) + + (tool_ids or []) + + (skill_tool_ids or []) + + ([mcp_tool_id] if mcp_tool_id else []) + ) + ) + authorized_set = set(filter_authorized_ids("tool", all_tool_ids, workspace_id, user_id=runtime_user_id)) + + mcp_tool_ids = [i for i in (mcp_tool_ids or []) if i in authorized_set] + tool_ids = [i for i in (tool_ids or []) if i in authorized_set] + skill_tool_ids = [i for i in (skill_tool_ids or []) if i in authorized_set] + mcp_tool_id = mcp_tool_id if (mcp_tool_id and mcp_tool_id in authorized_set) else None + # 处理 MCP 请求 + mcp_result = self._handle_mcp_request( + mcp_source, + mcp_servers, + mcp_tool_id, + mcp_tool_ids, + tool_ids, + application_ids, + skill_tool_ids, + mcp_output_enable, + chat_model, + SystemMessage(system), + message_list, + history_message, + question, + chat_id, + workspace_id, + runtime_user_id, + ) + if mcp_result: + return mcp_result + message_list = [SystemMessage(system)] + message_list + if stream: + r = chat_model.stream(message_list) + return NodeResult( + {"result": r, "chat_model": chat_model, "message_list": message_list, "question": question.content}, + {}, + _write_context=write_context_stream, + ) + else: + r = chat_model.invoke(message_list) + return NodeResult( + { + "result": r, + "chat_model": chat_model, + "message_list": message_list, + "history_message": [ + {"content": message.content, "role": message.type} + for message in (history_message if history_message is not None else []) + ], + "question": question.content, + }, + {}, + _write_context=write_context, + ) + + def _handle_mcp_request( + self, + mcp_source, + mcp_servers, + mcp_tool_id, + mcp_tool_ids, + tool_ids, + application_ids, + skill_tool_ids, + mcp_output_enable, + chat_model, + system_prompt, + message_list, + history_message, + question, + chat_id, + workspace_id, + runtime_user_id=None, + ): + + mcp_servers_config = {} + + # 迁移过来mcp_source是None + if mcp_source is None: + mcp_source = "custom" + # 兼容老数据 + if not mcp_tool_ids: + mcp_tool_ids = [] + if mcp_tool_id: + mcp_tool_ids = list(set(mcp_tool_ids + [mcp_tool_id])) + if mcp_source == "custom" and mcp_servers: + mcp_servers_config = json.loads(mcp_servers) + mcp_servers_config = self.handle_variables(mcp_servers_config) + elif mcp_tool_ids: + mcp_tools = QuerySet(Tool).filter(id__in=mcp_tool_ids).values() + for mcp_tool in mcp_tools: + if mcp_tool and mcp_tool["is_active"]: + mcp_tool_config = json.loads(mcp_tool["code"]) + # 引用多个 MCP 工具时服务名是用户自定义的,可能重名;直接合并会 + # 静默覆盖前面的服务,只有最后一个生效。这里显式报错指出冲突服务名。 + conflict_servers = set(mcp_servers_config) & set(mcp_tool_config) + if conflict_servers: + raise AppApiException( + 500, + _( + "MCP server 【{servers}】 is defined in multiple referenced MCP tools, rename it so every server takes effect" + ).format(servers="、".join(sorted(conflict_servers))), + ) + mcp_servers_config = {**mcp_servers_config, **mcp_tool_config} + mcp_servers_config = self.handle_variables(mcp_servers_config) + # 校验代码是否包括禁止的关键字 + ToolExecutor().validate_mcp_transport(json.dumps(mcp_servers_config)) + + tool_init_params = {} + tools = get_tools( + self.workflow_manage.get_source_type(), + self.workflow_manage.get_source_id(), + tool_ids, + workspace_id, + runtime_user_id, + ) + if tool_ids and len(tool_ids) > 0: # 如果有工具ID,则将其转换为MCP + self.context["tool_ids"] = tool_ids + custom_tools_map = { + str(t.id): t for t in QuerySet(Tool).filter(id__in=tool_ids, tool_type=ToolType.CUSTOM, is_active=True) + } + for tool_id in tool_ids: + tool = custom_tools_map.get(str(tool_id)) + if tool is None: + continue + executor = ToolExecutor() + init_params_default_value = {i["field"]: i.get("default_value") for i in tool.init_field_list} + if tool.init_params is not None: + tool_init_params = init_params_default_value | json.loads(rsa_long_decrypt(tool.init_params)) + else: + tool_init_params = init_params_default_value + + tool_config = executor.get_tool_mcp_config(tool, tool_init_params) + + mcp_servers_config[str(tool.id)] = tool_config + + if application_ids and len(application_ids) > 0: + self.context["application_ids"] = application_ids + apps_map = {str(a.id): a for a in QuerySet(Application).filter(id__in=application_ids, is_publish=True)} + app_keys_map = { + str(ak.application_id): ak + for ak in QuerySet(ApplicationApiKey).filter(application_id__in=application_ids, is_active=True) + } + app_access_tokens_map = { + str(at.application_id): at + for at in QuerySet(ApplicationAccessToken).filter(application_id__in=application_ids) + } + for application_id in application_ids: + app = apps_map.get(str(application_id)) + if app is None: + continue + app_key = app_keys_map.get(str(application_id)) + if app_key is not None: + api_key = app_key.secret_key + application_access_token = app_access_tokens_map.get(str(app_key.application_id)) + if application_access_token is not None and application_access_token.authentication: + raise AppApiException( + 500, + _("Agent 【{name}】 access token authentication is not supported for agent tool").format( + name=app.name + ), + ) + else: + raise AppApiException( + 500, _("Agent Key is required for agent tool 【{name}】").format(name=app.name) + ) + executor = ToolExecutor() + app_config = executor.get_app_mcp_config(api_key, self.get_chat_files(), self.get_form_data()) + mcp_servers_config[app.name] = app_config + + if skill_tool_ids and len(skill_tool_ids) > 0: + self.context["skill_tool_ids"] = skill_tool_ids + skill_file_items = [] + skill_tools_map = {str(t.id): t for t in QuerySet(Tool).filter(id__in=skill_tool_ids, is_active=True)} + for tool_id in skill_tool_ids: + tool = skill_tools_map.get(str(tool_id)) + if tool is None: + continue + init_params_default_value = {i["field"]: i.get("default_value") for i in tool.init_field_list} + if tool.init_params is not None: + params = init_params_default_value | json.loads(rsa_long_decrypt(tool.init_params)) + else: + params = init_params_default_value + + skill_file_items.append({"tool_id": str(tool.id), "file_id": tool.code, "params": params}) + mcp_servers_config["skills"] = skill_file_items + + if len(mcp_servers_config) > 0 or len(tools) > 0: + # 安全获取 application + application_id = None + tool_id = None + knowledge_id = None + if [WorkflowMode.KNOWLEDGE, WorkflowMode.KNOWLEDGE_LOOP].__contains__( + self.workflow_manage.flow.workflow_mode + ): + knowledge_id = self.workflow_params.get("knowledge_id") + elif [WorkflowMode.APPLICATION, WorkflowMode.APPLICATION_LOOP].__contains__( + self.workflow_manage.flow.workflow_mode + ): + application_id = self.workflow_manage.work_flow_post_handler.chat_info.application.id + elif [WorkflowMode.TOOL, WorkflowMode.TOOL_LOOP].__contains__(self.workflow_manage.flow.workflow_mode): + tool_id = self.workflow_params.get("tool_id") + + source_id = application_id or knowledge_id or tool_id + source_type = "APPLICATION" if application_id else "KNOWLEDGE" if knowledge_id else "TOOL" + r = mcp_response_generator( + chat_model, + system_prompt, + message_list, + mcp_servers_config, + mcp_output_enable, + tool_init_params, + source_id, + source_type, + chat_id, + tools, + ) + return NodeResult( + { + "result": r, + "chat_model": chat_model, + "message_list": message_list, + "history_message": [ + {"content": message.content, "role": message.type} + for message in (history_message if history_message is not None else []) + ], + "question": question.content, + }, + {}, + _write_context=write_context_stream, + ) + + return None + + def get_chat_files(self): + """ + 获取本次对话上传的文件, 用于透传给被当作工具调用的应用/MCP + """ + chat_files = {} + for field in CHAT_FILE_LIST_FIELDS: + file_list = getattr(self.workflow_manage, field, None) or [] + items = [ + {key: item.get(key) for key in ("name", "url", "file_id") if item.get(key) is not None} + for item in file_list + if isinstance(item, dict) + ] + if items: + chat_files[field] = items + return chat_files + + def get_form_data(self): + """ + 获取当前会话的用户输入参数,用于透传给作为工具调用的子智能体。 + + 循环工作流会创建独立的 WorkflowManage,并将自身的 form_data 初始化为 + 空字典,因此需要继续从父工作流中查找原始用户输入。 + """ + workflow_manage = self.workflow_manage + visited = set() + while workflow_manage is not None and id(workflow_manage) not in visited: + visited.add(id(workflow_manage)) + form_data = getattr(workflow_manage, "form_data", None) + if isinstance(form_data, dict) and form_data: + return form_data.copy() + workflow_manage = getattr(workflow_manage, "parentWorkflowManage", None) + return {} + + def handle_variables(self, tool_params): + # 处理参数中的变量 + for k, v in tool_params.items(): + if type(v) == str: + tool_params[k] = self.workflow_manage.generate_prompt(tool_params[k]) + elif type(v) == dict: + self.handle_variables(v) + elif (type(v) == list) and len(v) > 0 and (type(v[0]) == str): + tool_params[k] = self.get_reference_content(v) + return tool_params + + def get_reference_content(self, fields: List[str]): + return str(self.workflow_manage.get_reference_field(fields[0], fields[1:])) if fields else "" + + @staticmethod + def get_history_message(history_chat_record, dialogue_number, dialogue_type, runtime_node_id): + start_index = len(history_chat_record) - dialogue_number + history_message = reduce( + lambda x, y: [*x, *y], + [ + get_message(history_chat_record[index], dialogue_type, runtime_node_id) + for index in range(start_index if start_index > 0 else 0, len(history_chat_record)) + ], + [], + ) + for message in history_message: + if isinstance(message.content, str): + message.content = re.sub(r".*?<\/form_rander>", "", message.content, flags=re.DOTALL) + return history_message + + def generate_prompt_question(self, prompt, model): + image = self.get_image() + video = self.get_video() + vision = self.is_vision() + videos = [] + images = [] + if image and vision: + images = self._process_images(image) + if video and vision: + videos = self._process_videos(video, model) + prompt = self.workflow_manage.generate_prompt(prompt) + if images or videos: + return HumanMessage(content=[*videos, *images, {"type": "text", "text": prompt}]) + return HumanMessage(content=prompt) + + def is_vision(self): + if "vision" in self.node_params_serializer.data: + return self.node_params_serializer.data.get("vision") + return False + + def get_image(self): + if "image_list" in self.node_params_serializer.data: + image = self.workflow_manage.get_reference_field( + self.node_params_serializer.data.get("image_list")[0], + self.node_params_serializer.data.get("image_list")[1:], + ) + return image + return None + + def get_video(self): + if "video_list" in self.node_params_serializer.data: + video = self.workflow_manage.get_reference_field( + self.node_params_serializer.data.get("video_list")[0], + self.node_params_serializer.data.get("video_list")[1:], + ) + return video + return None + + def _process_videos(self, image, video_model): + videos = [] + 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: + 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}}) + elif "url" in img and img["url"].startswith("http"): + videos.append({"type": "video_url", "video_url": {"url": img["url"]}}) + return videos + + def _process_images(self, image): + """ + 处理图像数据,转换为模型可识别的格式 + """ + images = [] + 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: + 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}"}} + ) + elif "url" in img and img["url"].startswith("http"): + images.append({"type": "image_url", "image_url": {"url": img["url"]}}) + return images + + def generate_message_list(self, question, history_message): + return [*history_message, question] + + @staticmethod + def reset_message_list(message_list: List[BaseMessage], answer_text): + result = [ + {"role": "user" if isinstance(message, HumanMessage) else "ai", "content": message.content} + for message in message_list + ] + result.append({"role": "ai", "content": answer_text}) + return result + + def get_details(self, index: int, **kwargs): + return { + "name": self.node.properties.get("stepName"), + "index": index, + "run_time": self.context.get("run_time"), + "system": self.context.get("system"), + "history_message": self.context.get("history_message"), + "question": self.context.get("question"), + "answer": self.context.get("answer"), + "reasoning_content": self.context.get("reasoning_content"), + "enableException": self.node.properties.get("enableException"), + "type": self.node.type, + "message_tokens": self.context.get("message_tokens"), + "answer_tokens": self.context.get("answer_tokens"), + "status": self.status, + "err_message": self.err_message, + }