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,
+ }