From 58637ed9ef72bbee35f80efb741aa2a2f3f769e9 Mon Sep 17 00:00:00 2001 From: wangzifei Date: Tue, 22 Sep 2026 19:07:05 +0800 Subject: [PATCH] fix: skip degraded tool record when a detailed record exists Workflow and function tools already persist a detailed ToolRecord through their post-handlers when execution finishes; the agent stream then wrote a second degraded record (JSON-string input, no execution details) for the same execution from the ToolMessage branch, producing duplicate entries in the tool execution list (issue #7053). Skip the degraded insert when a detailed record already exists for this tool/source pair. Fixes #7053 Signed-off-by: wangzifei --- apps/application/flow/tools.py | 2309 ++++++++++++++++---------------- 1 file changed, 1162 insertions(+), 1147 deletions(-) diff --git a/apps/application/flow/tools.py b/apps/application/flow/tools.py index f1c678a5f8c..2c2499557f2 100644 --- a/apps/application/flow/tools.py +++ b/apps/application/flow/tools.py @@ -1,1147 +1,1162 @@ -# coding=utf-8 -""" -@project: maxkb -@Author:虎 -@file: utils.py -@date:2024/6/6 15:15 -@desc: -""" - -import asyncio -import io -import json -import os -import queue -import re -import shutil -import threading -import zipfile -from functools import reduce -from typing import Iterator - -# --------------------------------------------------------------------------- -# Fix: qwen's OpenAI-compatible streaming sends id='' (empty string) for -# intermediate tool_call_chunks while only the first chunk carries the real -# id ('call_xxx...'). langchain-core's merge_lists treats '' != 'call_xxx' as -# an ID conflict and _appends_ instead of merging → the accumulated AIMessage -# ends up with two separate tool_calls (one with empty args, one with empty -# id) instead of one correct entry. This causes the Qwen API to reject the -# next request with "function.arguments must be in JSON format". -# -# Patch: normalise id='' → None for items that have an 'index' key -# (i.e. tool_call_chunk dicts). merge_lists treats None as "no id" and will -# merge with any existing entry, keeping the real id from the first chunk. -# --------------------------------------------------------------------------- -import langchain_core.messages.ai as _lc_ai_module -import uuid_utils.compat as uuid -from asgiref.sync import sync_to_async -from common.result import result -from common.utils.logger import maxkb_logger -from deepagents import create_deep_agent -from django.db.models import OuterRef, QuerySet, Subquery -from django.http import StreamingHttpResponse -from knowledge.models import File -from knowledge.models.knowledge_action import State -from langchain_core.messages import AIMessageChunk, BaseMessage, BaseMessageChunk, ToolMessage -from langchain_core.tools import StructuredTool -from langchain_core.utils._merge import merge_lists as _original_merge_lists -from langgraph.checkpoint.memory import MemorySaver -from maxkb.const import CONFIG -from pydantic import Field, create_model -from tools.models import Tool, ToolRecord, ToolScope, ToolType, ToolWorkflowVersion - -from application.flow.backend.sandbox_mcp import SandboxMCPBackend -from application.flow.backend.sandbox_shell import SandboxShellBackend -from application.flow.common import Workflow, WorkflowMode -from application.flow.i_step_node import ToolWorkflowPostHandler, WorkFlowPostHandler -from application.serializers.common import ToolExecute - - -def _merge_lists_normalize_empty_tool_chunk_ids(left, *others): - """Wrapper around merge_lists that normalises empty-string IDs to None in - tool_call_chunk items (those with an 'index' key) so that qwen streaming - chunks with id='' are merged correctly by index.""" - - def _norm(lst): - if lst is None: - return lst - result = [] - for item in lst: - if isinstance(item, dict) and "index" in item and item.get("id") == "": - item = {**item, "id": None} - result.append(item) - return result - - return _original_merge_lists( - _norm(left), - *[_norm(o) for o in others], - ) - - -# Replace the module-level reference used by add_ai_message_chunks in ai.py -_lc_ai_module.merge_lists = _merge_lists_normalize_empty_tool_chunk_ids - - -class Reasoning: - def __init__(self, reasoning_content_start, reasoning_content_end): - self.content = "" - self.reasoning_content = "" - self.all_content = "" - self.reasoning_content_start_tag = reasoning_content_start - self.reasoning_content_end_tag = reasoning_content_end - self.reasoning_content_start_tag_len = ( - len(reasoning_content_start) if reasoning_content_start is not None else 0 - ) - self.reasoning_content_end_tag_len = len(reasoning_content_end) if reasoning_content_end is not None else 0 - self.reasoning_content_end_tag_prefix = ( - reasoning_content_end[0] if self.reasoning_content_end_tag_len > 0 else "" - ) - self.reasoning_content_is_start = False - self.reasoning_content_is_end = False - self.reasoning_content_chunk = "" - - def get_end_reasoning_content(self): - if not self.reasoning_content_is_start and not self.reasoning_content_is_end: - r = {"content": self.all_content, "reasoning_content": ""} - self.reasoning_content_chunk = "" - return r - if self.reasoning_content_is_start and not self.reasoning_content_is_end: - r = {"content": "", "reasoning_content": self.reasoning_content_chunk} - self.reasoning_content_chunk = "" - return r - return {"content": "", "reasoning_content": ""} - - def _normalize_content(self, content): - """将不同类型的内容统一转换为字符串""" - if isinstance(content, str): - return content - elif isinstance(content, list): - # 处理包含多种内容类型的列表 - normalized_parts = [] - for item in content: - if isinstance(item, dict): - if item.get("type") == "text": - normalized_parts.append(item.get("text", "")) - return "".join(normalized_parts) - else: - return str(content) - - def get_reasoning_content(self, chunk): - # 如果没有开始思考过程标签那么就全是结果 - if self.reasoning_content_start_tag is None or len(self.reasoning_content_start_tag) == 0: - self.content += chunk.content - return {"content": chunk.content, "reasoning_content": ""} - # 如果没有结束思考过程标签那么就全部是思考过程 - if self.reasoning_content_end_tag is None or len(self.reasoning_content_end_tag) == 0: - return {"content": "", "reasoning_content": chunk.content} - chunk.content = self._normalize_content(chunk.content) - self.all_content += chunk.content - if not self.reasoning_content_is_start and len(self.all_content) >= self.reasoning_content_start_tag_len: - if self.all_content.startswith(self.reasoning_content_start_tag): - self.reasoning_content_is_start = True - self.reasoning_content_chunk = self.all_content[self.reasoning_content_start_tag_len :] - else: - if not self.reasoning_content_is_end: - self.reasoning_content_is_end = True - self.content += self.all_content - return { - "content": self.all_content, - "reasoning_content": chunk.additional_kwargs.get("reasoning_content", "") - if chunk.additional_kwargs - else "", - } - else: - if self.reasoning_content_is_start: - self.reasoning_content_chunk += chunk.content - reasoning_content_end_tag_prefix_index = self.reasoning_content_chunk.find( - self.reasoning_content_end_tag_prefix - ) - if self.reasoning_content_is_end: - self.content += chunk.content - return { - "content": chunk.content, - "reasoning_content": chunk.additional_kwargs.get("reasoning_content", "") - if chunk.additional_kwargs - else "", - } - # 是否包含结束 - if reasoning_content_end_tag_prefix_index > -1: - if ( - len(self.reasoning_content_chunk) - reasoning_content_end_tag_prefix_index - >= self.reasoning_content_end_tag_len - ): - reasoning_content_end_tag_index = self.reasoning_content_chunk.find(self.reasoning_content_end_tag) - if reasoning_content_end_tag_index > -1: - reasoning_content_chunk = self.reasoning_content_chunk[0:reasoning_content_end_tag_index] - content_chunk = self.reasoning_content_chunk[ - reasoning_content_end_tag_index + self.reasoning_content_end_tag_len : - ] - self.reasoning_content += reasoning_content_chunk - self.content += content_chunk - self.reasoning_content_chunk = "" - self.reasoning_content_is_end = True - return {"content": content_chunk, "reasoning_content": reasoning_content_chunk} - else: - reasoning_content_chunk = self.reasoning_content_chunk[ - 0 : reasoning_content_end_tag_prefix_index + 1 - ] - self.reasoning_content_chunk = self.reasoning_content_chunk.replace(reasoning_content_chunk, "") - self.reasoning_content += reasoning_content_chunk - return {"content": "", "reasoning_content": reasoning_content_chunk} - else: - return {"content": "", "reasoning_content": ""} - - else: - if self.reasoning_content_is_end: - self.content += chunk.content - return { - "content": chunk.content, - "reasoning_content": chunk.additional_kwargs.get("reasoning_content", "") - if chunk.additional_kwargs - else "", - } - else: - # aaa - result = {"content": "", "reasoning_content": self.reasoning_content_chunk} - self.reasoning_content += self.reasoning_content_chunk - self.reasoning_content_chunk = "" - return result - - -def event_content(chat_id, chat_record_id, response, workflow, write_context, post_handler: WorkFlowPostHandler): - """ - 用于处理流式输出 - @param chat_id: 会话id - @param chat_record_id: 对话记录id - @param response: 响应数据 - @param workflow: 工作流管理器 - @param write_context 写入节点上下文 - @param post_handler: 后置处理器 - """ - answer = "" - try: - for chunk in response: - answer += chunk.content - yield ( - "data: " - + json.dumps( - { - "chat_id": str(chat_id), - "id": str(chat_record_id), - "operate": True, - "content": chunk.content, - "is_end": False, - }, - ensure_ascii=False, - ) - + "\n\n" - ) - write_context(answer, 200) - post_handler.handler(chat_id, chat_record_id, answer, workflow) - yield ( - "data: " - + json.dumps( - {"chat_id": str(chat_id), "id": str(chat_record_id), "operate": True, "content": "", "is_end": True}, - ensure_ascii=False, - ) - + "\n\n" - ) - except Exception as e: - answer = str(e) - write_context(answer, 500) - post_handler.handler(chat_id, chat_record_id, answer, workflow) - yield ( - "data: " - + json.dumps( - { - "chat_id": str(chat_id), - "id": str(chat_record_id), - "operate": True, - "content": answer, - "is_end": True, - }, - ensure_ascii=False, - ) - + "\n\n" - ) - - -def to_stream_response( - chat_id, chat_record_id, response: Iterator[BaseMessageChunk], workflow, write_context, post_handler -): - """ - 将结果转换为服务流输出 - @param chat_id: 会话id - @param chat_record_id: 对话记录id - @param response: 响应数据 - @param workflow: 工作流管理器 - @param write_context 写入节点上下文 - @param post_handler: 后置处理器 - @return: 响应 - """ - r = StreamingHttpResponse( - streaming_content=event_content(chat_id, chat_record_id, response, workflow, write_context, post_handler), - content_type="text/event-stream;charset=utf-8", - charset="utf-8", - ) - - r["Cache-Control"] = "no-cache" - return r - - -def to_response( - chat_id, chat_record_id, response: BaseMessage, workflow, write_context, post_handler: WorkFlowPostHandler -): - """ - 将结果转换为服务输出 - - @param chat_id: 会话id - @param chat_record_id: 对话记录id - @param response: 响应数据 - @param workflow: 工作流管理器 - @param write_context 写入节点上下文 - @param post_handler: 后置处理器 - @return: 响应 - """ - answer = response.content - write_context(answer) - post_handler.handler(chat_id, chat_record_id, answer, workflow) - return result.success( - {"chat_id": str(chat_id), "id": str(chat_record_id), "operate": True, "content": answer, "is_end": True} - ) - - -def to_response_simple(chat_id, chat_record_id, response: BaseMessage, workflow, post_handler: WorkFlowPostHandler): - answer = response.content - post_handler.handler(chat_id, chat_record_id, answer, workflow) - return result.success( - {"chat_id": str(chat_id), "id": str(chat_record_id), "operate": True, "content": answer, "is_end": True} - ) - - -def to_stream_response_simple(stream_event): - r = StreamingHttpResponse( - streaming_content=stream_event, content_type="text/event-stream;charset=utf-8", charset="utf-8" - ) - - r["Cache-Control"] = "no-cache" - return r - - -def generate_tool_message_complete(icon, name, input_content, output_content): - """生成包含输入和输出的工具消息模版""" - # 确保输入内容是字符串,如果不是则尝试转换为 JSON 字符串 - if not isinstance(input_content, str): - input_content = json.dumps(input_content, ensure_ascii=False) - # 格式化输出 - if not isinstance(output_content, str): - output_content = json.dumps(output_content, ensure_ascii=False) - content = { - "icon": icon, - "title": name, - "type": "simple-tool-calls", - "content": {"input": input_content, "output": output_content}, - } - return f"{json.dumps(content, ensure_ascii=False)}" - - -# 全局单例事件循环 -_global_loop = None -_loop_thread = None -_loop_lock = threading.Lock() - - -def get_global_loop(): - """获取全局共享的事件循环""" - global _global_loop, _loop_thread - - with _loop_lock: - if _global_loop is None: - _global_loop = asyncio.new_event_loop() - - def run_forever(): - asyncio.set_event_loop(_global_loop) - _global_loop.run_forever() - - _loop_thread = threading.Thread(target=run_forever, daemon=True, name="GlobalAsyncLoop") - _loop_thread.start() - - return _global_loop - - -def _extract_tool_id(raw_id): - """从 raw_id 中提取最后一个符合 call_... 模式的 id,若无匹配则返回原值或 None""" - if not raw_id: - return None - if not isinstance(raw_id, str): - raw_id = str(raw_id) - - s = raw_id - prefix = "call_" - positions = [m.start() for m in re.finditer(re.escape(prefix), s)] - if not positions: - return raw_id - - # 取最后一个前缀位置,截到下一个前缀或结尾 - start = positions[-1] - end = len(s) - for pos in positions: - if pos > start: - end = pos - break - - tool_id = s[start:end] - return tool_id or raw_id - - -async def _initialize_skills(mcp_servers, temp_dir) -> SandboxMCPBackend: - skills_dir = os.path.join(temp_dir, "skills") - mcp_config = dict(mcp_servers) # Preserve server-generated InternalMCPConfig objects. - if "skills" in mcp_config: - skill_file_items = mcp_config.pop("skills") - for skill_file in skill_file_items: - # 使用 sync_to_async 包装 ORM 查询 - file = await sync_to_async(lambda: QuerySet(File).filter(id=skill_file["file_id"]).first())() - if not file: - continue - # get_bytes 可能也涉及 IO,也用 sync_to_async 包装 - file_bytes = await sync_to_async(file.get_bytes)() - params = skill_file.get("params", {}) - with zipfile.ZipFile(io.BytesIO(file_bytes), "r") as zip_ref: - members = [m for m in zip_ref.namelist() if not m.startswith("__MACOSX/") and "__MACOSX" not in m] - for member in members: - if ".." in member or member.startswith("/"): - raise ValueError(f"非法路径: {member}") - zip_ref.extractall(skills_dir, members=members) - - # 获取技能解压后的顶级目录名 - top_level_dirs = set() - for member in members: - parts = member.split("/") - if parts[0]: - top_level_dirs.add(parts[0]) - - # 将 params 写入每个顶级目录下的 .env 文件 - if params: - env_lines = [] - for key, value in params.items(): - # 对含空格或特殊字符的值加引号 - env_lines.append(f"{key}={value}") - env_content = "\n".join(env_lines) + "\n" - for top_dir in top_level_dirs: - env_path = os.path.join(skills_dir, top_dir, ".env") - with open(env_path, "w", encoding="utf-8") as f: - f.write(env_content) - - os.system("chmod -R g+rx " + temp_dir) # 确保技能目录可访问 - - return SandboxMCPBackend(mcp_config) - - -async def _yield_mcp_response( - chat_model, - system_prompt, - message_list, - mcp_servers, - mcp_output_enable=True, - tool_init_params={}, - source_id=None, - source_type=None, - temp_dir=None, - chat_id=None, - extra_tools=None, -): - try: - checkpointer = MemorySaver() - mcp_backend = await _initialize_skills(mcp_servers, temp_dir) - tools = await mcp_backend.get_tools() - for tool in tools: - tool.handle_tool_error = True - if extra_tools: - for tool in extra_tools: - tools.append(tool) - - agent = create_deep_agent( - model=chat_model, - backend=SandboxShellBackend(root_dir=temp_dir, virtual_mode=True), - skills=["/skills"], - tools=tools, - system_prompt=system_prompt, - interrupt_on={"write_file": False, "read_file": False, "edit_file": False}, - checkpointer=checkpointer, - ) - recursion_limit = int(CONFIG.get("LANGCHAIN_GRAPH_RECURSION_LIMIT", "100")) - response = agent.astream( - {"messages": message_list}, - config={"recursion_limit": recursion_limit, "configurable": {"thread_id": chat_id}}, - stream_mode="messages", - ) - - tool_calls_info = {} # tool_id -> {'name': ..., 'input': ...} - # key(index/id) -> {'id': ..., 'name': ..., 'arguments': ...} - _tool_fragments = {} - - def _merge_arguments(entry, part_args): - if not isinstance(part_args, str): - try: - part_args = json.dumps(part_args, ensure_ascii=False) - except Exception: - part_args = str(part_args) if part_args else "" - if not part_args: - return - - # Some providers first emit placeholder args like "{}" and then - # stream the real JSON fragments via later chunks. Prefer fragments. - if entry["arguments"] in ("{}", "[]") and part_args.startswith("{"): - entry["arguments"] = part_args - return - - if entry["arguments"]: - try: - existing_obj = json.loads(entry["arguments"]) - new_obj = json.loads(part_args) - if isinstance(existing_obj, dict) and isinstance(new_obj, dict): - merged = {**existing_obj, **new_obj} - entry["arguments"] = json.dumps(merged, ensure_ascii=False) - else: - entry["arguments"] += part_args - except (json.JSONDecodeError, ValueError): - entry["arguments"] += part_args - else: - entry["arguments"] = part_args - - def _get_fragment_key(idx, raw_id): - if idx is not None: - return f"idx:{idx}" - if raw_id and str(raw_id).strip(): - return f"id:{_extract_tool_id(str(raw_id).strip())}" - return None - - def _upsert_fragment(key, raw_id, func_name, part_args): - if key is None: - return - entry = _tool_fragments.setdefault(key, {"id": "", "name": "", "arguments": ""}) - - if raw_id and str(raw_id).strip(): - new_id = str(raw_id).strip() - if entry.get("completed") and entry.get("id") and entry["id"] != new_id: - maxkb_logger.debug(f"Resetting completed fragment {key}: old ID {entry['id']} -> new ID {new_id}") - entry.clear() - entry.update({"id": "", "name": "", "arguments": ""}) - entry["id"] = new_id - - if func_name: - entry["name"] = func_name - - _merge_arguments(entry, part_args) - - async for chunk in response: - # print(chunk) - if isinstance(chunk[0], AIMessageChunk): - # ---------------------------------------------------------------- - # 1. 从 tool_call_chunks 中聚合工具调用片段 - # (qwen/OpenAI streaming 通过 tool_call_chunks 传递, - # additional_kwargs['tool_calls'] 在流式时通常为空) - # ---------------------------------------------------------------- - for tc_chunk in chunk[0].tool_call_chunks or []: - raw_id = tc_chunk.get("id") - key = _get_fragment_key(tc_chunk.get("index"), raw_id) - _upsert_fragment(key, raw_id, tc_chunk.get("name"), tc_chunk.get("args", "")) - - # ---------------------------------------------------------------- - # 1.1 兼容部分模型将工具调用放在 chunk.tool_calls,且 tool_call_chunks - # 的 index 为空(例如 ollama/qwen) - # ---------------------------------------------------------------- - has_tool_call_chunks = bool(chunk[0].tool_call_chunks) - for tool_call in chunk[0].tool_calls or []: - raw_id = tool_call.get("id") - part_args = tool_call.get("args", "") - # qwen-plus often emits {} here as a placeholder while - # the real args are split in tool_call_chunks/invalid_tool_calls. - if has_tool_call_chunks and (part_args == "" or part_args == {} or part_args == []): - part_args = "" - key = _get_fragment_key(tool_call.get("index"), raw_id) - _upsert_fragment(key, raw_id, tool_call.get("name"), part_args) - - # ---------------------------------------------------------------- - # 1.2 兼容 invalid_tool_calls 分片(部分模型会把中间 JSON 片段放这里) - # ---------------------------------------------------------------- - for invalid_tool_call in chunk[0].invalid_tool_calls or []: - raw_id = invalid_tool_call.get("id") - key = _get_fragment_key(invalid_tool_call.get("index"), raw_id) - _upsert_fragment(key, raw_id, invalid_tool_call.get("name"), invalid_tool_call.get("args", "")) - - # ---------------------------------------------------------------- - # 2. 兼容 additional_kwargs['tool_calls'] 方式(旧格式/非流式情况) - # ---------------------------------------------------------------- - legacy_tool_calls = chunk[0].additional_kwargs.get("tool_calls", []) - for tool_call in legacy_tool_calls: - raw_id = tool_call.get("id") - func = tool_call.get("function", {}) - if isinstance(func, dict): - func_name = func.get("name") - part_args = func.get("arguments", "") - else: - func_name = tool_call.get("name") - part_args = tool_call.get("arguments", "") - key = _get_fragment_key(tool_call.get("index"), raw_id) - _upsert_fragment(key, raw_id, func_name, part_args) - - # ---------------------------------------------------------------- - # 3. 检测工具调用结束,更新 tool_calls_info - # ---------------------------------------------------------------- - is_finish_chunk = ( - chunk[0].response_metadata.get("finish_reason") == "tool_calls" or chunk[0].chunk_position == "last" - ) - - if is_finish_chunk: - # 在 finish chunk 时,将所有未完成的 fragment 标记完成并更新 tool_calls_info - maxkb_logger.debug(f"Processing finish chunk. Tool fragments: {_tool_fragments}") - for idx, entry in _tool_fragments.items(): - if entry.get("completed"): - maxkb_logger.debug(f"Skipping fragment {idx}: already completed") - continue - if not entry.get("id"): - maxkb_logger.debug(f"Skipping fragment {idx}: missing id. Fragment: {entry}") - continue - if not entry.get("arguments"): - maxkb_logger.debug(f"Skipping fragment {idx}: missing arguments. Fragment: {entry}") - continue - - if not entry.get("completed") and entry.get("id") and entry.get("arguments"): - try: - parsed_args = json.loads(entry["arguments"]) - filtered_args = ( - {k: v for k, v in parsed_args.items() if k not in tool_init_params} - if tool_init_params - else parsed_args - ) - normalized_id = _extract_tool_id(entry["id"]) - info = {"name": entry["name"], "input": json.dumps(filtered_args, ensure_ascii=False)} - tool_calls_info[entry["id"]] = info - if normalized_id and normalized_id != entry["id"]: - tool_calls_info[normalized_id] = info - entry["completed"] = True - maxkb_logger.debug(f"Added tool call {entry['id']} to tool_calls_info") - except (json.JSONDecodeError, ValueError) as e: - # JSON parsing failed, but still add to tool_calls_info with raw arguments - # to prevent "Tool ID not found" errors when ToolMessage arrives - maxkb_logger.warning( - f"Failed to parse tool arguments at finish for tool {entry.get('id', 'unknown')}: " - f"{entry['arguments']}, error: {e}. Using raw arguments." - ) - normalized_id = _extract_tool_id(entry["id"]) - info = { - "name": entry["name"], - # Use raw arguments - "input": entry["arguments"], - } - tool_calls_info[entry["id"]] = info - if normalized_id and normalized_id != entry["id"]: - tool_calls_info[normalized_id] = info - entry["completed"] = True - - # ---------------------------------------------------------------- - # 4. 修复 tool_call_chunks 中的空 id(回填已知 id) - # ---------------------------------------------------------------- - if chunk[0].tool_call_chunks: - for tc_chunk in chunk[0].tool_call_chunks: - key = _get_fragment_key(tc_chunk.get("index"), tc_chunk.get("id")) - if key is not None: - frag = _tool_fragments.get(key) - if frag and frag.get("id") and not tc_chunk.get("id"): - tc_chunk["id"] = frag["id"] - - # ---------------------------------------------------------------- - # 5. 修复 additional_kwargs['tool_calls'](兼容旧格式) - # 仅在 finish chunk 时写入完整参数,避免污染中间 chunk 的 - # additional_kwargs(中间 chunk 会被 ainvoke 累积,如果写入 - # 不完整 JSON 会导致下一轮 API 调用出现 arguments 非 JSON 格式错误) - # ---------------------------------------------------------------- - if legacy_tool_calls and is_finish_chunk: - fixed_tool_calls = [] - for tool_call in legacy_tool_calls: - key = _get_fragment_key(tool_call.get("index"), tool_call.get("id")) - frag = _tool_fragments.get(key) if key is not None else None - tc = dict(tool_call) - if frag and frag.get("id") and not tc.get("id"): - tc["id"] = frag["id"] - if frag and isinstance(tc.get("function"), dict): - tc["function"] = dict(tc["function"]) - if frag.get("completed"): - tc["function"]["arguments"] = frag["arguments"] - fixed_tool_calls.append(tc) - chunk[0].additional_kwargs["tool_calls"] = fixed_tool_calls - - yield chunk[0] - - if mcp_output_enable and isinstance(chunk[0], ToolMessage): - tool_id = chunk[0].tool_call_id - normalized_tool_id = _extract_tool_id(tool_id) - tool_info = tool_calls_info.get(tool_id) or tool_calls_info.get(normalized_tool_id) - - if tool_info: - try: - if isinstance(chunk[0].content, str): - tool_result = json.loads(chunk[0].content) - elif isinstance(chunk[0].content, dict): - tool_result = chunk[0].content - elif isinstance(chunk[0].content, list): - tool_result = chunk[0].content[0] if len(chunk[0].content) > 0 else {} - else: - tool_result = {} - text = tool_result.get("text") if "text" in tool_result else None - text_result = json.loads(text) if text else tool_result - if text: - tool_lib_id = text_result.pop("tool_id") if "tool_id" in text_result else None - else: - tool_lib_id = tool_result.pop("tool_id") if "tool_id" in tool_result else None - if tool_lib_id: - await save_tool_record(tool_lib_id, tool_info, tool_result, source_id, source_type) - tool_result = json.dumps(text_result, ensure_ascii=False) - except Exception as e: - tool_result = chunk[0].content - content = generate_tool_message_complete( - tool_info.get("icon", ""), tool_info["name"], tool_info["input"], tool_result - ) - chunk[0].content = content - else: - maxkb_logger.warning( - f"Tool ID {tool_id} not found in tool_calls_info. " - f"Normalized Tool ID: {normalized_tool_id}. " - f"Available IDs: {list(tool_calls_info.keys())}. " - f"Tool fragments at this point: {_tool_fragments}" - ) - - yield chunk[0] - - except ExceptionGroup as eg: - - def get_real_error(exc): - if isinstance(exc, ExceptionGroup): - return get_real_error(exc.exceptions[0]) - return exc - - real_error = get_real_error(eg) - error_msg = f"{type(real_error).__name__}: {str(real_error)}" - raise RuntimeError(error_msg) from None - - except Exception as e: - error_msg = f"{type(e).__name__}: {str(e)}" - raise RuntimeError(error_msg) from None - - -async def save_tool_record(tool_id, tool_info, tool_result, source_id, source_type): - from django.db import close_old_connections - await sync_to_async(close_old_connections)() - tool = await sync_to_async(lambda: QuerySet(Tool).filter(id=tool_id).first())() - tool_info["icon"] = tool.icon - tool_record = ToolRecord( - id=uuid.uuid7(), - workspace_id=tool.workspace_id, - tool_id=tool_id, - source_type=source_type, - source_id=source_id, - meta={"input": tool_info["input"], "output": tool_result}, - state=State.SUCCESS, - ) - await sync_to_async(tool_record.save)() - - -def mcp_response_generator( - chat_model, - system_prompt, - message_list, - mcp_servers, - mcp_output_enable=True, - tool_init_params={}, - source_id=None, - source_type=None, - chat_id=None, - extra_tools=None, -): - """使用全局事件循环,不创建新实例""" - result_queue = queue.Queue() - loop = get_global_loop() # 使用共享循环 - # 创建临时文件夹 - if chat_id: - temp_dir = os.path.join("/tmp", chat_id) - else: - temp_dir = os.path.join("/tmp", str(uuid.uuid7())) - skills_dir = os.path.join(temp_dir, "skills") - os.makedirs(skills_dir, exist_ok=True) - - # print(f"Initializing skills in temporary directory: {skills_dir}") - - async def _run(): - try: - async_gen = _yield_mcp_response( - chat_model, - system_prompt, - message_list, - mcp_servers, - mcp_output_enable, - tool_init_params, - source_id, - source_type, - temp_dir, - chat_id, - extra_tools, - ) - async for chunk in async_gen: - result_queue.put(("data", chunk)) - except Exception as e: - maxkb_logger.error(f"Exception: {e}", exc_info=True) - result_queue.put(("error", e)) - finally: - result_queue.put(("done", None)) - - # 在全局循环中调度任务 - asyncio.run_coroutine_threadsafe(_run(), loop) - - while True: - msg_type, data = result_queue.get() - if msg_type == "done": - # 清理临时文件夹 - shutil.rmtree(temp_dir, ignore_errors=True) - break - if msg_type == "error": - # 清理临时文件夹 - shutil.rmtree(temp_dir, ignore_errors=True) - raise data - yield data - - -async def anext_async(agen): - return await agen.__anext__() - - -def _get_node_model_id(node, model_field, mode_field): - """节点为 default/reference 模式时不返回节点内 model_id(运行时才解析,避免脏映射)。""" - node_data = (node.get("properties") or {}).get("node_data") or {} - if node_data.get(mode_field) in ("default", "reference"): - return None - return node_data.get(model_field) - - -# base-node 三类模型:mode 判定与 validate_workflow_default_models/get_base_node_model 保持一致 -# (stt/长期记忆用 'default'/'reference',tts 用大写 'DEFAULT'/'BROWSER') -_base_node_model_specs = ( - ("stt_model_id_type", ("default", "reference"), "stt_model_enable", "stt_model_id"), - ("tts_type", ("DEFAULT", "BROWSER"), "tts_model_enable", "tts_model_id"), - ("long_term_model_id_type", ("default", "reference"), "long_term_enable", "long_term_model_id"), -) - - -def _get_base_node_model_ids(node): - """返回 base-node node_data 中实际自定义的 STT/TTS/长期记忆 model_id(default/BROWSER 时运行时解析,不映射)。""" - node_data = (node.get("properties") or {}).get("node_data") or {} - model_ids = [] - for mode_field, skip_modes, enable_field, model_field in _base_node_model_specs: - if node_data.get(enable_field) and node_data.get(mode_field) not in skip_modes: - if node_data.get(model_field): - model_ids.append(node_data.get(model_field)) - return model_ids - - -target_source_node_mapping = { - "TOOL": { - "tool-lib-node": lambda n: [n.get("properties").get("node_data").get("tool_lib_id")], - "ai-chat-node": lambda n: [ - *(n.get("properties").get("node_data").get("mcp_tool_ids") or []), - *(n.get("properties").get("node_data").get("tool_ids") or []), - *(n.get("properties").get("node_data").get("skill_tool_ids") or []), - ], - "mcp-node": lambda n: [n.get("properties").get("node_data").get("mcp_tool_id")], - "tool-workflow-lib-node": lambda n: [n.get("properties").get("node_data").get("tool_lib_id")], - }, - "MODEL": { - "ai-chat-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], - "question-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], - "speech-to-text-node": lambda n: [v for v in [_get_node_model_id(n, 'stt_model_id', 'stt_model_id_type')] if v], - "text-to-speech-node": lambda n: [v for v in [_get_node_model_id(n, 'tts_model_id', 'tts_model_id_type')] if v], - "image-to-video-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], - "image-generate-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], - "intent-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], - "image-understand-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], - "parameter-extraction-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], - "video-understand-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], - "reranker-node": lambda n: [v for v in [_get_node_model_id(n, 'reranker_model_id', 'reranker_model_id_type')] if v], - "base-node": _get_base_node_model_ids, - }, - "KNOWLEDGE": { - "search-knowledge-node": lambda n: n.get("properties").get("node_data").get("knowledge_id_list"), - "search-document-node": lambda n: n.get("properties").get("node_data").get("knowledge_id_list"), - }, - "APPLICATION": { - "application-node": lambda n: [n.get("properties").get("node_data").get("application_id")], - "ai-chat-node": lambda n: [*(n.get("properties").get("node_data").get("application_ids") or [])], - }, -} - - -def get_node_handle_callback(source_type, source_id): - def node_handle_callback(node): - from system_manage.models.resource_mapping import ResourceMapping - - response = [] - for key, value in target_source_node_mapping.items(): - if node.get("type") in value: - call = value.get(node.get("type")) - target_source_id_list = call(node) - for target_source_id in target_source_id_list: - if target_source_id: - response.append( - ResourceMapping( - source_type=source_type, - target_type=key, - source_id=source_id, - target_id=target_source_id, - ) - ) - return response - - return node_handle_callback - - -def get_workflow_resource(workflow, node_handle): - response = [] - if "nodes" in workflow: - for node in workflow.get("nodes"): - rs = node_handle(node) - if rs: - for r in rs: - response.append(r) - if node.get("type") == "loop-node": - r = get_workflow_resource(node.get("properties", {}).get("node_data", {}).get("loop_body"), node_handle) - for rn in r: - response.append(rn) - return list({(str(item.target_type) + str(item.target_id)): item for item in response}.values()) - return [] - - -application_instance_field_call_dict = { - "TOOL": [ - lambda instance: instance.mcp_tool_ids or [], - lambda instance: instance.skill_tool_ids or [], - lambda instance: instance.tool_ids or [], - ], - "APPLICATION": [ - lambda instance: instance.application_ids or [], - ], - "MODEL": [ - lambda instance: [instance.model_id] if instance.model_id else [], - lambda instance: [instance.long_term_model_id] if instance.long_term_model_id else [], - lambda instance: [instance.tts_model_id] if instance.tts_model_id else [], - lambda instance: [instance.stt_model_id] if instance.stt_model_id else [], - lambda instance: [v.get('model_id') for v in (instance.default_model_setting or {}).values() if (v or {}).get('model_id')], - ], -} -knowledge_instance_field_call_dict = { - "MODEL": [lambda instance: [instance.embedding_model_id] if instance.embedding_model_id else []], -} - - -def get_instance_resource(instance, source_type, source_id, instance_field_call_dict): - response = [] - from system_manage.models.resource_mapping import ResourceMapping - - for target_type, call_list in instance_field_call_dict.items(): - target_id_list = reduce(lambda x, y: [*x, *y], [call(instance) for call in call_list], []) - if target_id_list: - for target_id in target_id_list: - response.append( - ResourceMapping( - source_type=source_type, target_type=target_type, source_id=source_id, target_id=target_id - ) - ) - return response - - -def append_default_model_mapping(instance_mapping, default_model_setting, source_type, source_id): - """把 default_model_setting 各类别 model_id 追加为 MODEL 资源映射(方案A),返回追加后的列表。""" - from system_manage.models.resource_mapping import ResourceMapping, ResourceType - - for value in (default_model_setting or {}).values(): - model_id = (value or {}).get('model_id') - if model_id: - instance_mapping.append( - ResourceMapping( - source_type=source_type, target_type=ResourceType.MODEL, - source_id=str(source_id), target_id=model_id, - ) - ) - return instance_mapping - - -def save_workflow_mapping(workflow, source_type, source_id, other_resource_mapping=None): - if not other_resource_mapping: - other_resource_mapping = [] - from django.db.models import QuerySet - from system_manage.models.resource_mapping import ResourceMapping - - QuerySet(ResourceMapping).filter(source_type=source_type, source_id=source_id).delete() - resource_mapping_list = get_workflow_resource(workflow, get_node_handle_callback(source_type, source_id)) - resource_mapping_list += other_resource_mapping - if resource_mapping_list: - QuerySet(ResourceMapping).bulk_create( - {(str(item.target_type) + str(item.target_id)): item for item in resource_mapping_list}.values() - ) - - -def get_tool_id_list(workflow, with_deep=False): - from tools.models import ToolType, ToolWorkflow - - _result = [] - for node in workflow.get("nodes", []): - if node.get("type") == "tool-lib-node": - tool_id = node.get("properties", {}).get("node_data", {}).get("tool_lib_id") - if tool_id: - _result.append(tool_id) - elif node.get("type") == "loop-node": - r = get_tool_id_list(node.get("properties", {}).get("node_data", {}).get("loop_body", {})) - for item in r: - _result.append(item) - elif node.get("type") == "tool-workflow-lib-node": - tool_id = node.get("properties", {}).get("node_data", {}).get("tool_lib_id") - if tool_id: - _result.append(tool_id) - elif node.get("type") == "ai-chat-node": - node_data = node.get("properties", {}).get("node_data", {}) - mcp_tool_ids = node_data.get("mcp_tool_ids") or [] - skill_tool_ids = node_data.get("skill_tool_ids") or [] - tool_ids = node_data.get("tool_ids") or [] - for _id in mcp_tool_ids + tool_ids + skill_tool_ids: - _result.append(_id) - elif node.get("type") == "mcp-node": - mcp_tool_id = node.get("properties", {}).get("node_data", {}).get("mcp_tool_id") - if mcp_tool_id: - _result.append(mcp_tool_id) - if with_deep: - workflow_list = QuerySet(Tool).filter(id__in=_result, tool_type=ToolType.WORKFLOW) - tool_work_flow_list = QuerySet(ToolWorkflow).filter(tool_id__in=[wl.id for wl in workflow_list]) - for tool_work_flow in tool_work_flow_list: - child_tool_id_list = get_child_tool_id_list(tool_work_flow.work_flow, []) - for c in child_tool_id_list: - _result.append(c) - return _result - - -def get_child_tool_id_list(work_flow, response): - from tools.models import ToolType, ToolWorkflow - - tool_id_list = get_tool_id_list(work_flow, False) - tool_id_list = [tool_id for tool_id in tool_id_list if len([r for r in response if r == tool_id]) == 0] - tool_list = [] - if len(tool_id_list) > 0: - tool_list = QuerySet(Tool).filter(id__in=tool_id_list).exclude(scope=ToolScope.SHARED) - work_flow_tools = [tool for tool in tool_list if tool.tool_type == ToolType.WORKFLOW] - if len(work_flow_tools) > 0: - work_flow_tool_dict = { - tw.tool_id: tw for tw in QuerySet(ToolWorkflow).filter(tool_id__in=[t.id for t in work_flow_tools]) - } - for tool in tool_list: - response.append(str(tool.id)) - if tool.tool_type == ToolType.WORKFLOW: - get_child_tool_id_list(work_flow_tool_dict.get(tool.id).work_flow, response) - else: - for tool in tool_list: - response.append(str(tool.id)) - return response - - -def build_schema(fields: dict): - return create_model("dynamicSchema", **fields) - - -def get_type(_type: str): - if _type == "float": - return float - if _type == "string": - return str - if _type == "int": - return int - if _type == "dict": - return dict - if _type == "array": - return list - if _type == "boolean": - return bool - return object - - -def get_workflow_args(tool, qv): - for node in qv.work_flow.get("nodes"): - if node.get("type") == "tool-base-node": - input_field_list = node.get("properties").get("user_input_field_list") - return build_schema( - { - field.get("field"): ( - get_type(field.get("type")), - Field(..., required=True, description=field.get("desc")) if field.get("is_required") else Field(default=None, required=False, description=field.get("desc")) - ) - for field in input_field_list - } - ) - - return build_schema({}) - - -def get_workflow_func(source_type, source_id, tool, qv, workspace_id, user_id=None): - tool_id = tool.id - tool_record_id = str(uuid.uuid7()) - took_execute = ToolExecute(tool_id, tool_record_id, workspace_id, source_type, source_id, False) - - def inner(**kwargs): - from application.flow.tool_workflow_manage import ToolWorkflowManage - - work_flow_manage = ToolWorkflowManage( - Workflow.new_instance(qv.work_flow, WorkflowMode.TOOL), - { - "chat_record_id": tool_record_id, - "tool_id": tool_id, - "stream": True, - "workspace_id": workspace_id, - "user_id": user_id, - **kwargs, - "default_model_setting": qv.default_model_setting, - }, - ToolWorkflowPostHandler(took_execute, tool_id), - is_the_task_interrupted=lambda: False, - child_node=None, - start_node_id=None, - start_node_data=None, - chat_record=None, - ) - res = work_flow_manage.run() - for r in res: - pass - return work_flow_manage.out_context - - return inner - - -def get_tools(source_type, source_id, tool_workflow_ids, workspace_id, user_id=None): - tools = QuerySet(Tool).filter( - id__in=tool_workflow_ids, is_active=True, tool_type=ToolType.WORKFLOW, workspace_id=workspace_id - ) - latest_subquery = ToolWorkflowVersion.objects.filter(tool_id=OuterRef("tool_id")).order_by("-create_time") - - qs = ToolWorkflowVersion.objects.filter( - tool_id__in=[t.id for t in tools], id=Subquery(latest_subquery.values("id")[:1]) - ) - qd = {q.tool_id: q for q in qs} - results = [] - for tool in tools: - qv = qd.get(tool.id) - func = get_workflow_func(source_type, source_id, tool, qv, workspace_id, user_id=user_id) - args = get_workflow_args(tool, qv) - tool = StructuredTool.from_function( - func=func, - name=tool.name, - description=tool.desc, - args_schema=args, - ) - results.append(tool) - - return results +# coding=utf-8 +""" +@project: maxkb +@Author:虎 +@file: utils.py +@date:2024/6/6 15:15 +@desc: +""" + +import asyncio +import io +import json +import os +import queue +import re +import shutil +import threading +import zipfile +from functools import reduce +from typing import Iterator + +# --------------------------------------------------------------------------- +# Fix: qwen's OpenAI-compatible streaming sends id='' (empty string) for +# intermediate tool_call_chunks while only the first chunk carries the real +# id ('call_xxx...'). langchain-core's merge_lists treats '' != 'call_xxx' as +# an ID conflict and _appends_ instead of merging → the accumulated AIMessage +# ends up with two separate tool_calls (one with empty args, one with empty +# id) instead of one correct entry. This causes the Qwen API to reject the +# next request with "function.arguments must be in JSON format". +# +# Patch: normalise id='' → None for items that have an 'index' key +# (i.e. tool_call_chunk dicts). merge_lists treats None as "no id" and will +# merge with any existing entry, keeping the real id from the first chunk. +# --------------------------------------------------------------------------- +import langchain_core.messages.ai as _lc_ai_module +import uuid_utils.compat as uuid +from asgiref.sync import sync_to_async +from common.result import result +from common.utils.logger import maxkb_logger +from deepagents import create_deep_agent +from django.db.models import OuterRef, QuerySet, Subquery +from django.http import StreamingHttpResponse +from knowledge.models import File +from knowledge.models.knowledge_action import State +from langchain_core.messages import AIMessageChunk, BaseMessage, BaseMessageChunk, ToolMessage +from langchain_core.tools import StructuredTool +from langchain_core.utils._merge import merge_lists as _original_merge_lists +from langgraph.checkpoint.memory import MemorySaver +from maxkb.const import CONFIG +from pydantic import Field, create_model +from tools.models import Tool, ToolRecord, ToolScope, ToolType, ToolWorkflowVersion + +from application.flow.backend.sandbox_mcp import SandboxMCPBackend +from application.flow.backend.sandbox_shell import SandboxShellBackend +from application.flow.common import Workflow, WorkflowMode +from application.flow.i_step_node import ToolWorkflowPostHandler, WorkFlowPostHandler +from application.serializers.common import ToolExecute + + +def _merge_lists_normalize_empty_tool_chunk_ids(left, *others): + """Wrapper around merge_lists that normalises empty-string IDs to None in + tool_call_chunk items (those with an 'index' key) so that qwen streaming + chunks with id='' are merged correctly by index.""" + + def _norm(lst): + if lst is None: + return lst + result = [] + for item in lst: + if isinstance(item, dict) and "index" in item and item.get("id") == "": + item = {**item, "id": None} + result.append(item) + return result + + return _original_merge_lists( + _norm(left), + *[_norm(o) for o in others], + ) + + +# Replace the module-level reference used by add_ai_message_chunks in ai.py +_lc_ai_module.merge_lists = _merge_lists_normalize_empty_tool_chunk_ids + + +class Reasoning: + def __init__(self, reasoning_content_start, reasoning_content_end): + self.content = "" + self.reasoning_content = "" + self.all_content = "" + self.reasoning_content_start_tag = reasoning_content_start + self.reasoning_content_end_tag = reasoning_content_end + self.reasoning_content_start_tag_len = ( + len(reasoning_content_start) if reasoning_content_start is not None else 0 + ) + self.reasoning_content_end_tag_len = len(reasoning_content_end) if reasoning_content_end is not None else 0 + self.reasoning_content_end_tag_prefix = ( + reasoning_content_end[0] if self.reasoning_content_end_tag_len > 0 else "" + ) + self.reasoning_content_is_start = False + self.reasoning_content_is_end = False + self.reasoning_content_chunk = "" + + def get_end_reasoning_content(self): + if not self.reasoning_content_is_start and not self.reasoning_content_is_end: + r = {"content": self.all_content, "reasoning_content": ""} + self.reasoning_content_chunk = "" + return r + if self.reasoning_content_is_start and not self.reasoning_content_is_end: + r = {"content": "", "reasoning_content": self.reasoning_content_chunk} + self.reasoning_content_chunk = "" + return r + return {"content": "", "reasoning_content": ""} + + def _normalize_content(self, content): + """将不同类型的内容统一转换为字符串""" + if isinstance(content, str): + return content + elif isinstance(content, list): + # 处理包含多种内容类型的列表 + normalized_parts = [] + for item in content: + if isinstance(item, dict): + if item.get("type") == "text": + normalized_parts.append(item.get("text", "")) + return "".join(normalized_parts) + else: + return str(content) + + def get_reasoning_content(self, chunk): + # 如果没有开始思考过程标签那么就全是结果 + if self.reasoning_content_start_tag is None or len(self.reasoning_content_start_tag) == 0: + self.content += chunk.content + return {"content": chunk.content, "reasoning_content": ""} + # 如果没有结束思考过程标签那么就全部是思考过程 + if self.reasoning_content_end_tag is None or len(self.reasoning_content_end_tag) == 0: + return {"content": "", "reasoning_content": chunk.content} + chunk.content = self._normalize_content(chunk.content) + self.all_content += chunk.content + if not self.reasoning_content_is_start and len(self.all_content) >= self.reasoning_content_start_tag_len: + if self.all_content.startswith(self.reasoning_content_start_tag): + self.reasoning_content_is_start = True + self.reasoning_content_chunk = self.all_content[self.reasoning_content_start_tag_len :] + else: + if not self.reasoning_content_is_end: + self.reasoning_content_is_end = True + self.content += self.all_content + return { + "content": self.all_content, + "reasoning_content": chunk.additional_kwargs.get("reasoning_content", "") + if chunk.additional_kwargs + else "", + } + else: + if self.reasoning_content_is_start: + self.reasoning_content_chunk += chunk.content + reasoning_content_end_tag_prefix_index = self.reasoning_content_chunk.find( + self.reasoning_content_end_tag_prefix + ) + if self.reasoning_content_is_end: + self.content += chunk.content + return { + "content": chunk.content, + "reasoning_content": chunk.additional_kwargs.get("reasoning_content", "") + if chunk.additional_kwargs + else "", + } + # 是否包含结束 + if reasoning_content_end_tag_prefix_index > -1: + if ( + len(self.reasoning_content_chunk) - reasoning_content_end_tag_prefix_index + >= self.reasoning_content_end_tag_len + ): + reasoning_content_end_tag_index = self.reasoning_content_chunk.find(self.reasoning_content_end_tag) + if reasoning_content_end_tag_index > -1: + reasoning_content_chunk = self.reasoning_content_chunk[0:reasoning_content_end_tag_index] + content_chunk = self.reasoning_content_chunk[ + reasoning_content_end_tag_index + self.reasoning_content_end_tag_len : + ] + self.reasoning_content += reasoning_content_chunk + self.content += content_chunk + self.reasoning_content_chunk = "" + self.reasoning_content_is_end = True + return {"content": content_chunk, "reasoning_content": reasoning_content_chunk} + else: + reasoning_content_chunk = self.reasoning_content_chunk[ + 0 : reasoning_content_end_tag_prefix_index + 1 + ] + self.reasoning_content_chunk = self.reasoning_content_chunk.replace(reasoning_content_chunk, "") + self.reasoning_content += reasoning_content_chunk + return {"content": "", "reasoning_content": reasoning_content_chunk} + else: + return {"content": "", "reasoning_content": ""} + + else: + if self.reasoning_content_is_end: + self.content += chunk.content + return { + "content": chunk.content, + "reasoning_content": chunk.additional_kwargs.get("reasoning_content", "") + if chunk.additional_kwargs + else "", + } + else: + # aaa + result = {"content": "", "reasoning_content": self.reasoning_content_chunk} + self.reasoning_content += self.reasoning_content_chunk + self.reasoning_content_chunk = "" + return result + + +def event_content(chat_id, chat_record_id, response, workflow, write_context, post_handler: WorkFlowPostHandler): + """ + 用于处理流式输出 + @param chat_id: 会话id + @param chat_record_id: 对话记录id + @param response: 响应数据 + @param workflow: 工作流管理器 + @param write_context 写入节点上下文 + @param post_handler: 后置处理器 + """ + answer = "" + try: + for chunk in response: + answer += chunk.content + yield ( + "data: " + + json.dumps( + { + "chat_id": str(chat_id), + "id": str(chat_record_id), + "operate": True, + "content": chunk.content, + "is_end": False, + }, + ensure_ascii=False, + ) + + "\n\n" + ) + write_context(answer, 200) + post_handler.handler(chat_id, chat_record_id, answer, workflow) + yield ( + "data: " + + json.dumps( + {"chat_id": str(chat_id), "id": str(chat_record_id), "operate": True, "content": "", "is_end": True}, + ensure_ascii=False, + ) + + "\n\n" + ) + except Exception as e: + answer = str(e) + write_context(answer, 500) + post_handler.handler(chat_id, chat_record_id, answer, workflow) + yield ( + "data: " + + json.dumps( + { + "chat_id": str(chat_id), + "id": str(chat_record_id), + "operate": True, + "content": answer, + "is_end": True, + }, + ensure_ascii=False, + ) + + "\n\n" + ) + + +def to_stream_response( + chat_id, chat_record_id, response: Iterator[BaseMessageChunk], workflow, write_context, post_handler +): + """ + 将结果转换为服务流输出 + @param chat_id: 会话id + @param chat_record_id: 对话记录id + @param response: 响应数据 + @param workflow: 工作流管理器 + @param write_context 写入节点上下文 + @param post_handler: 后置处理器 + @return: 响应 + """ + r = StreamingHttpResponse( + streaming_content=event_content(chat_id, chat_record_id, response, workflow, write_context, post_handler), + content_type="text/event-stream;charset=utf-8", + charset="utf-8", + ) + + r["Cache-Control"] = "no-cache" + return r + + +def to_response( + chat_id, chat_record_id, response: BaseMessage, workflow, write_context, post_handler: WorkFlowPostHandler +): + """ + 将结果转换为服务输出 + + @param chat_id: 会话id + @param chat_record_id: 对话记录id + @param response: 响应数据 + @param workflow: 工作流管理器 + @param write_context 写入节点上下文 + @param post_handler: 后置处理器 + @return: 响应 + """ + answer = response.content + write_context(answer) + post_handler.handler(chat_id, chat_record_id, answer, workflow) + return result.success( + {"chat_id": str(chat_id), "id": str(chat_record_id), "operate": True, "content": answer, "is_end": True} + ) + + +def to_response_simple(chat_id, chat_record_id, response: BaseMessage, workflow, post_handler: WorkFlowPostHandler): + answer = response.content + post_handler.handler(chat_id, chat_record_id, answer, workflow) + return result.success( + {"chat_id": str(chat_id), "id": str(chat_record_id), "operate": True, "content": answer, "is_end": True} + ) + + +def to_stream_response_simple(stream_event): + r = StreamingHttpResponse( + streaming_content=stream_event, content_type="text/event-stream;charset=utf-8", charset="utf-8" + ) + + r["Cache-Control"] = "no-cache" + return r + + +def generate_tool_message_complete(icon, name, input_content, output_content): + """生成包含输入和输出的工具消息模版""" + # 确保输入内容是字符串,如果不是则尝试转换为 JSON 字符串 + if not isinstance(input_content, str): + input_content = json.dumps(input_content, ensure_ascii=False) + # 格式化输出 + if not isinstance(output_content, str): + output_content = json.dumps(output_content, ensure_ascii=False) + content = { + "icon": icon, + "title": name, + "type": "simple-tool-calls", + "content": {"input": input_content, "output": output_content}, + } + return f"{json.dumps(content, ensure_ascii=False)}" + + +# 全局单例事件循环 +_global_loop = None +_loop_thread = None +_loop_lock = threading.Lock() + + +def get_global_loop(): + """获取全局共享的事件循环""" + global _global_loop, _loop_thread + + with _loop_lock: + if _global_loop is None: + _global_loop = asyncio.new_event_loop() + + def run_forever(): + asyncio.set_event_loop(_global_loop) + _global_loop.run_forever() + + _loop_thread = threading.Thread(target=run_forever, daemon=True, name="GlobalAsyncLoop") + _loop_thread.start() + + return _global_loop + + +def _extract_tool_id(raw_id): + """从 raw_id 中提取最后一个符合 call_... 模式的 id,若无匹配则返回原值或 None""" + if not raw_id: + return None + if not isinstance(raw_id, str): + raw_id = str(raw_id) + + s = raw_id + prefix = "call_" + positions = [m.start() for m in re.finditer(re.escape(prefix), s)] + if not positions: + return raw_id + + # 取最后一个前缀位置,截到下一个前缀或结尾 + start = positions[-1] + end = len(s) + for pos in positions: + if pos > start: + end = pos + break + + tool_id = s[start:end] + return tool_id or raw_id + + +async def _initialize_skills(mcp_servers, temp_dir) -> SandboxMCPBackend: + skills_dir = os.path.join(temp_dir, "skills") + mcp_config = dict(mcp_servers) # Preserve server-generated InternalMCPConfig objects. + if "skills" in mcp_config: + skill_file_items = mcp_config.pop("skills") + for skill_file in skill_file_items: + # 使用 sync_to_async 包装 ORM 查询 + file = await sync_to_async(lambda: QuerySet(File).filter(id=skill_file["file_id"]).first())() + if not file: + continue + # get_bytes 可能也涉及 IO,也用 sync_to_async 包装 + file_bytes = await sync_to_async(file.get_bytes)() + params = skill_file.get("params", {}) + with zipfile.ZipFile(io.BytesIO(file_bytes), "r") as zip_ref: + members = [m for m in zip_ref.namelist() if not m.startswith("__MACOSX/") and "__MACOSX" not in m] + for member in members: + if ".." in member or member.startswith("/"): + raise ValueError(f"非法路径: {member}") + zip_ref.extractall(skills_dir, members=members) + + # 获取技能解压后的顶级目录名 + top_level_dirs = set() + for member in members: + parts = member.split("/") + if parts[0]: + top_level_dirs.add(parts[0]) + + # 将 params 写入每个顶级目录下的 .env 文件 + if params: + env_lines = [] + for key, value in params.items(): + # 对含空格或特殊字符的值加引号 + env_lines.append(f"{key}={value}") + env_content = "\n".join(env_lines) + "\n" + for top_dir in top_level_dirs: + env_path = os.path.join(skills_dir, top_dir, ".env") + with open(env_path, "w", encoding="utf-8") as f: + f.write(env_content) + + os.system("chmod -R g+rx " + temp_dir) # 确保技能目录可访问 + + return SandboxMCPBackend(mcp_config) + + +async def _yield_mcp_response( + chat_model, + system_prompt, + message_list, + mcp_servers, + mcp_output_enable=True, + tool_init_params={}, + source_id=None, + source_type=None, + temp_dir=None, + chat_id=None, + extra_tools=None, +): + try: + checkpointer = MemorySaver() + mcp_backend = await _initialize_skills(mcp_servers, temp_dir) + tools = await mcp_backend.get_tools() + for tool in tools: + tool.handle_tool_error = True + if extra_tools: + for tool in extra_tools: + tools.append(tool) + + agent = create_deep_agent( + model=chat_model, + backend=SandboxShellBackend(root_dir=temp_dir, virtual_mode=True), + skills=["/skills"], + tools=tools, + system_prompt=system_prompt, + interrupt_on={"write_file": False, "read_file": False, "edit_file": False}, + checkpointer=checkpointer, + ) + recursion_limit = int(CONFIG.get("LANGCHAIN_GRAPH_RECURSION_LIMIT", "100")) + response = agent.astream( + {"messages": message_list}, + config={"recursion_limit": recursion_limit, "configurable": {"thread_id": chat_id}}, + stream_mode="messages", + ) + + tool_calls_info = {} # tool_id -> {'name': ..., 'input': ...} + # key(index/id) -> {'id': ..., 'name': ..., 'arguments': ...} + _tool_fragments = {} + + def _merge_arguments(entry, part_args): + if not isinstance(part_args, str): + try: + part_args = json.dumps(part_args, ensure_ascii=False) + except Exception: + part_args = str(part_args) if part_args else "" + if not part_args: + return + + # Some providers first emit placeholder args like "{}" and then + # stream the real JSON fragments via later chunks. Prefer fragments. + if entry["arguments"] in ("{}", "[]") and part_args.startswith("{"): + entry["arguments"] = part_args + return + + if entry["arguments"]: + try: + existing_obj = json.loads(entry["arguments"]) + new_obj = json.loads(part_args) + if isinstance(existing_obj, dict) and isinstance(new_obj, dict): + merged = {**existing_obj, **new_obj} + entry["arguments"] = json.dumps(merged, ensure_ascii=False) + else: + entry["arguments"] += part_args + except (json.JSONDecodeError, ValueError): + entry["arguments"] += part_args + else: + entry["arguments"] = part_args + + def _get_fragment_key(idx, raw_id): + if idx is not None: + return f"idx:{idx}" + if raw_id and str(raw_id).strip(): + return f"id:{_extract_tool_id(str(raw_id).strip())}" + return None + + def _upsert_fragment(key, raw_id, func_name, part_args): + if key is None: + return + entry = _tool_fragments.setdefault(key, {"id": "", "name": "", "arguments": ""}) + + if raw_id and str(raw_id).strip(): + new_id = str(raw_id).strip() + if entry.get("completed") and entry.get("id") and entry["id"] != new_id: + maxkb_logger.debug(f"Resetting completed fragment {key}: old ID {entry['id']} -> new ID {new_id}") + entry.clear() + entry.update({"id": "", "name": "", "arguments": ""}) + entry["id"] = new_id + + if func_name: + entry["name"] = func_name + + _merge_arguments(entry, part_args) + + async for chunk in response: + # print(chunk) + if isinstance(chunk[0], AIMessageChunk): + # ---------------------------------------------------------------- + # 1. 从 tool_call_chunks 中聚合工具调用片段 + # (qwen/OpenAI streaming 通过 tool_call_chunks 传递, + # additional_kwargs['tool_calls'] 在流式时通常为空) + # ---------------------------------------------------------------- + for tc_chunk in chunk[0].tool_call_chunks or []: + raw_id = tc_chunk.get("id") + key = _get_fragment_key(tc_chunk.get("index"), raw_id) + _upsert_fragment(key, raw_id, tc_chunk.get("name"), tc_chunk.get("args", "")) + + # ---------------------------------------------------------------- + # 1.1 兼容部分模型将工具调用放在 chunk.tool_calls,且 tool_call_chunks + # 的 index 为空(例如 ollama/qwen) + # ---------------------------------------------------------------- + has_tool_call_chunks = bool(chunk[0].tool_call_chunks) + for tool_call in chunk[0].tool_calls or []: + raw_id = tool_call.get("id") + part_args = tool_call.get("args", "") + # qwen-plus often emits {} here as a placeholder while + # the real args are split in tool_call_chunks/invalid_tool_calls. + if has_tool_call_chunks and (part_args == "" or part_args == {} or part_args == []): + part_args = "" + key = _get_fragment_key(tool_call.get("index"), raw_id) + _upsert_fragment(key, raw_id, tool_call.get("name"), part_args) + + # ---------------------------------------------------------------- + # 1.2 兼容 invalid_tool_calls 分片(部分模型会把中间 JSON 片段放这里) + # ---------------------------------------------------------------- + for invalid_tool_call in chunk[0].invalid_tool_calls or []: + raw_id = invalid_tool_call.get("id") + key = _get_fragment_key(invalid_tool_call.get("index"), raw_id) + _upsert_fragment(key, raw_id, invalid_tool_call.get("name"), invalid_tool_call.get("args", "")) + + # ---------------------------------------------------------------- + # 2. 兼容 additional_kwargs['tool_calls'] 方式(旧格式/非流式情况) + # ---------------------------------------------------------------- + legacy_tool_calls = chunk[0].additional_kwargs.get("tool_calls", []) + for tool_call in legacy_tool_calls: + raw_id = tool_call.get("id") + func = tool_call.get("function", {}) + if isinstance(func, dict): + func_name = func.get("name") + part_args = func.get("arguments", "") + else: + func_name = tool_call.get("name") + part_args = tool_call.get("arguments", "") + key = _get_fragment_key(tool_call.get("index"), raw_id) + _upsert_fragment(key, raw_id, func_name, part_args) + + # ---------------------------------------------------------------- + # 3. 检测工具调用结束,更新 tool_calls_info + # ---------------------------------------------------------------- + is_finish_chunk = ( + chunk[0].response_metadata.get("finish_reason") == "tool_calls" or chunk[0].chunk_position == "last" + ) + + if is_finish_chunk: + # 在 finish chunk 时,将所有未完成的 fragment 标记完成并更新 tool_calls_info + maxkb_logger.debug(f"Processing finish chunk. Tool fragments: {_tool_fragments}") + for idx, entry in _tool_fragments.items(): + if entry.get("completed"): + maxkb_logger.debug(f"Skipping fragment {idx}: already completed") + continue + if not entry.get("id"): + maxkb_logger.debug(f"Skipping fragment {idx}: missing id. Fragment: {entry}") + continue + if not entry.get("arguments"): + maxkb_logger.debug(f"Skipping fragment {idx}: missing arguments. Fragment: {entry}") + continue + + if not entry.get("completed") and entry.get("id") and entry.get("arguments"): + try: + parsed_args = json.loads(entry["arguments"]) + filtered_args = ( + {k: v for k, v in parsed_args.items() if k not in tool_init_params} + if tool_init_params + else parsed_args + ) + normalized_id = _extract_tool_id(entry["id"]) + info = {"name": entry["name"], "input": json.dumps(filtered_args, ensure_ascii=False)} + tool_calls_info[entry["id"]] = info + if normalized_id and normalized_id != entry["id"]: + tool_calls_info[normalized_id] = info + entry["completed"] = True + maxkb_logger.debug(f"Added tool call {entry['id']} to tool_calls_info") + except (json.JSONDecodeError, ValueError) as e: + # JSON parsing failed, but still add to tool_calls_info with raw arguments + # to prevent "Tool ID not found" errors when ToolMessage arrives + maxkb_logger.warning( + f"Failed to parse tool arguments at finish for tool {entry.get('id', 'unknown')}: " + f"{entry['arguments']}, error: {e}. Using raw arguments." + ) + normalized_id = _extract_tool_id(entry["id"]) + info = { + "name": entry["name"], + # Use raw arguments + "input": entry["arguments"], + } + tool_calls_info[entry["id"]] = info + if normalized_id and normalized_id != entry["id"]: + tool_calls_info[normalized_id] = info + entry["completed"] = True + + # ---------------------------------------------------------------- + # 4. 修复 tool_call_chunks 中的空 id(回填已知 id) + # ---------------------------------------------------------------- + if chunk[0].tool_call_chunks: + for tc_chunk in chunk[0].tool_call_chunks: + key = _get_fragment_key(tc_chunk.get("index"), tc_chunk.get("id")) + if key is not None: + frag = _tool_fragments.get(key) + if frag and frag.get("id") and not tc_chunk.get("id"): + tc_chunk["id"] = frag["id"] + + # ---------------------------------------------------------------- + # 5. 修复 additional_kwargs['tool_calls'](兼容旧格式) + # 仅在 finish chunk 时写入完整参数,避免污染中间 chunk 的 + # additional_kwargs(中间 chunk 会被 ainvoke 累积,如果写入 + # 不完整 JSON 会导致下一轮 API 调用出现 arguments 非 JSON 格式错误) + # ---------------------------------------------------------------- + if legacy_tool_calls and is_finish_chunk: + fixed_tool_calls = [] + for tool_call in legacy_tool_calls: + key = _get_fragment_key(tool_call.get("index"), tool_call.get("id")) + frag = _tool_fragments.get(key) if key is not None else None + tc = dict(tool_call) + if frag and frag.get("id") and not tc.get("id"): + tc["id"] = frag["id"] + if frag and isinstance(tc.get("function"), dict): + tc["function"] = dict(tc["function"]) + if frag.get("completed"): + tc["function"]["arguments"] = frag["arguments"] + fixed_tool_calls.append(tc) + chunk[0].additional_kwargs["tool_calls"] = fixed_tool_calls + + yield chunk[0] + + if mcp_output_enable and isinstance(chunk[0], ToolMessage): + tool_id = chunk[0].tool_call_id + normalized_tool_id = _extract_tool_id(tool_id) + tool_info = tool_calls_info.get(tool_id) or tool_calls_info.get(normalized_tool_id) + + if tool_info: + try: + if isinstance(chunk[0].content, str): + tool_result = json.loads(chunk[0].content) + elif isinstance(chunk[0].content, dict): + tool_result = chunk[0].content + elif isinstance(chunk[0].content, list): + tool_result = chunk[0].content[0] if len(chunk[0].content) > 0 else {} + else: + tool_result = {} + text = tool_result.get("text") if "text" in tool_result else None + text_result = json.loads(text) if text else tool_result + if text: + tool_lib_id = text_result.pop("tool_id") if "tool_id" in text_result else None + else: + tool_lib_id = tool_result.pop("tool_id") if "tool_id" in tool_result else None + if tool_lib_id: + await save_tool_record(tool_lib_id, tool_info, tool_result, source_id, source_type) + tool_result = json.dumps(text_result, ensure_ascii=False) + except Exception as e: + tool_result = chunk[0].content + content = generate_tool_message_complete( + tool_info.get("icon", ""), tool_info["name"], tool_info["input"], tool_result + ) + chunk[0].content = content + else: + maxkb_logger.warning( + f"Tool ID {tool_id} not found in tool_calls_info. " + f"Normalized Tool ID: {normalized_tool_id}. " + f"Available IDs: {list(tool_calls_info.keys())}. " + f"Tool fragments at this point: {_tool_fragments}" + ) + + yield chunk[0] + + except ExceptionGroup as eg: + + def get_real_error(exc): + if isinstance(exc, ExceptionGroup): + return get_real_error(exc.exceptions[0]) + return exc + + real_error = get_real_error(eg) + error_msg = f"{type(real_error).__name__}: {str(real_error)}" + raise RuntimeError(error_msg) from None + + except Exception as e: + error_msg = f"{type(e).__name__}: {str(e)}" + raise RuntimeError(error_msg) from None + + +async def save_tool_record(tool_id, tool_info, tool_result, source_id, source_type): + from django.db import close_old_connections + await sync_to_async(close_old_connections)() + tool = await sync_to_async(lambda: QuerySet(Tool).filter(id=tool_id).first())() + tool_info["icon"] = tool.icon + + # 工作流/函数库工具执行完会由各自的后置处理器先写入一条带 details 的完整记录 + # (ToolWorkflowPostHandler、函数库节点自记录)。同一次执行再走这里会多出一条 + # 只有 input/output 的降级记录,导致工具执行记录重复,因此已有完整记录时跳过。 + def _has_detailed_record(): + return ( + QuerySet(ToolRecord) + .filter(tool_id=tool_id, source_id=source_id, source_type=source_type) + .filter(meta__has_key="details") + .first() + ) is not None + + if await sync_to_async(_has_detailed_record)(): + return + + tool_record = ToolRecord( + id=uuid.uuid7(), + workspace_id=tool.workspace_id, + tool_id=tool_id, + source_type=source_type, + source_id=source_id, + meta={"input": tool_info["input"], "output": tool_result}, + state=State.SUCCESS, + ) + await sync_to_async(tool_record.save)() + + +def mcp_response_generator( + chat_model, + system_prompt, + message_list, + mcp_servers, + mcp_output_enable=True, + tool_init_params={}, + source_id=None, + source_type=None, + chat_id=None, + extra_tools=None, +): + """使用全局事件循环,不创建新实例""" + result_queue = queue.Queue() + loop = get_global_loop() # 使用共享循环 + # 创建临时文件夹 + if chat_id: + temp_dir = os.path.join("/tmp", chat_id) + else: + temp_dir = os.path.join("/tmp", str(uuid.uuid7())) + skills_dir = os.path.join(temp_dir, "skills") + os.makedirs(skills_dir, exist_ok=True) + + # print(f"Initializing skills in temporary directory: {skills_dir}") + + async def _run(): + try: + async_gen = _yield_mcp_response( + chat_model, + system_prompt, + message_list, + mcp_servers, + mcp_output_enable, + tool_init_params, + source_id, + source_type, + temp_dir, + chat_id, + extra_tools, + ) + async for chunk in async_gen: + result_queue.put(("data", chunk)) + except Exception as e: + maxkb_logger.error(f"Exception: {e}", exc_info=True) + result_queue.put(("error", e)) + finally: + result_queue.put(("done", None)) + + # 在全局循环中调度任务 + asyncio.run_coroutine_threadsafe(_run(), loop) + + while True: + msg_type, data = result_queue.get() + if msg_type == "done": + # 清理临时文件夹 + shutil.rmtree(temp_dir, ignore_errors=True) + break + if msg_type == "error": + # 清理临时文件夹 + shutil.rmtree(temp_dir, ignore_errors=True) + raise data + yield data + + +async def anext_async(agen): + return await agen.__anext__() + + +def _get_node_model_id(node, model_field, mode_field): + """节点为 default/reference 模式时不返回节点内 model_id(运行时才解析,避免脏映射)。""" + node_data = (node.get("properties") or {}).get("node_data") or {} + if node_data.get(mode_field) in ("default", "reference"): + return None + return node_data.get(model_field) + + +# base-node 三类模型:mode 判定与 validate_workflow_default_models/get_base_node_model 保持一致 +# (stt/长期记忆用 'default'/'reference',tts 用大写 'DEFAULT'/'BROWSER') +_base_node_model_specs = ( + ("stt_model_id_type", ("default", "reference"), "stt_model_enable", "stt_model_id"), + ("tts_type", ("DEFAULT", "BROWSER"), "tts_model_enable", "tts_model_id"), + ("long_term_model_id_type", ("default", "reference"), "long_term_enable", "long_term_model_id"), +) + + +def _get_base_node_model_ids(node): + """返回 base-node node_data 中实际自定义的 STT/TTS/长期记忆 model_id(default/BROWSER 时运行时解析,不映射)。""" + node_data = (node.get("properties") or {}).get("node_data") or {} + model_ids = [] + for mode_field, skip_modes, enable_field, model_field in _base_node_model_specs: + if node_data.get(enable_field) and node_data.get(mode_field) not in skip_modes: + if node_data.get(model_field): + model_ids.append(node_data.get(model_field)) + return model_ids + + +target_source_node_mapping = { + "TOOL": { + "tool-lib-node": lambda n: [n.get("properties").get("node_data").get("tool_lib_id")], + "ai-chat-node": lambda n: [ + *(n.get("properties").get("node_data").get("mcp_tool_ids") or []), + *(n.get("properties").get("node_data").get("tool_ids") or []), + *(n.get("properties").get("node_data").get("skill_tool_ids") or []), + ], + "mcp-node": lambda n: [n.get("properties").get("node_data").get("mcp_tool_id")], + "tool-workflow-lib-node": lambda n: [n.get("properties").get("node_data").get("tool_lib_id")], + }, + "MODEL": { + "ai-chat-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], + "question-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], + "speech-to-text-node": lambda n: [v for v in [_get_node_model_id(n, 'stt_model_id', 'stt_model_id_type')] if v], + "text-to-speech-node": lambda n: [v for v in [_get_node_model_id(n, 'tts_model_id', 'tts_model_id_type')] if v], + "image-to-video-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], + "image-generate-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], + "intent-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], + "image-understand-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], + "parameter-extraction-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], + "video-understand-node": lambda n: [v for v in [_get_node_model_id(n, 'model_id', 'model_id_type')] if v], + "reranker-node": lambda n: [v for v in [_get_node_model_id(n, 'reranker_model_id', 'reranker_model_id_type')] if v], + "base-node": _get_base_node_model_ids, + }, + "KNOWLEDGE": { + "search-knowledge-node": lambda n: n.get("properties").get("node_data").get("knowledge_id_list"), + "search-document-node": lambda n: n.get("properties").get("node_data").get("knowledge_id_list"), + }, + "APPLICATION": { + "application-node": lambda n: [n.get("properties").get("node_data").get("application_id")], + "ai-chat-node": lambda n: [*(n.get("properties").get("node_data").get("application_ids") or [])], + }, +} + + +def get_node_handle_callback(source_type, source_id): + def node_handle_callback(node): + from system_manage.models.resource_mapping import ResourceMapping + + response = [] + for key, value in target_source_node_mapping.items(): + if node.get("type") in value: + call = value.get(node.get("type")) + target_source_id_list = call(node) + for target_source_id in target_source_id_list: + if target_source_id: + response.append( + ResourceMapping( + source_type=source_type, + target_type=key, + source_id=source_id, + target_id=target_source_id, + ) + ) + return response + + return node_handle_callback + + +def get_workflow_resource(workflow, node_handle): + response = [] + if "nodes" in workflow: + for node in workflow.get("nodes"): + rs = node_handle(node) + if rs: + for r in rs: + response.append(r) + if node.get("type") == "loop-node": + r = get_workflow_resource(node.get("properties", {}).get("node_data", {}).get("loop_body"), node_handle) + for rn in r: + response.append(rn) + return list({(str(item.target_type) + str(item.target_id)): item for item in response}.values()) + return [] + + +application_instance_field_call_dict = { + "TOOL": [ + lambda instance: instance.mcp_tool_ids or [], + lambda instance: instance.skill_tool_ids or [], + lambda instance: instance.tool_ids or [], + ], + "APPLICATION": [ + lambda instance: instance.application_ids or [], + ], + "MODEL": [ + lambda instance: [instance.model_id] if instance.model_id else [], + lambda instance: [instance.long_term_model_id] if instance.long_term_model_id else [], + lambda instance: [instance.tts_model_id] if instance.tts_model_id else [], + lambda instance: [instance.stt_model_id] if instance.stt_model_id else [], + lambda instance: [v.get('model_id') for v in (instance.default_model_setting or {}).values() if (v or {}).get('model_id')], + ], +} +knowledge_instance_field_call_dict = { + "MODEL": [lambda instance: [instance.embedding_model_id] if instance.embedding_model_id else []], +} + + +def get_instance_resource(instance, source_type, source_id, instance_field_call_dict): + response = [] + from system_manage.models.resource_mapping import ResourceMapping + + for target_type, call_list in instance_field_call_dict.items(): + target_id_list = reduce(lambda x, y: [*x, *y], [call(instance) for call in call_list], []) + if target_id_list: + for target_id in target_id_list: + response.append( + ResourceMapping( + source_type=source_type, target_type=target_type, source_id=source_id, target_id=target_id + ) + ) + return response + + +def append_default_model_mapping(instance_mapping, default_model_setting, source_type, source_id): + """把 default_model_setting 各类别 model_id 追加为 MODEL 资源映射(方案A),返回追加后的列表。""" + from system_manage.models.resource_mapping import ResourceMapping, ResourceType + + for value in (default_model_setting or {}).values(): + model_id = (value or {}).get('model_id') + if model_id: + instance_mapping.append( + ResourceMapping( + source_type=source_type, target_type=ResourceType.MODEL, + source_id=str(source_id), target_id=model_id, + ) + ) + return instance_mapping + + +def save_workflow_mapping(workflow, source_type, source_id, other_resource_mapping=None): + if not other_resource_mapping: + other_resource_mapping = [] + from django.db.models import QuerySet + from system_manage.models.resource_mapping import ResourceMapping + + QuerySet(ResourceMapping).filter(source_type=source_type, source_id=source_id).delete() + resource_mapping_list = get_workflow_resource(workflow, get_node_handle_callback(source_type, source_id)) + resource_mapping_list += other_resource_mapping + if resource_mapping_list: + QuerySet(ResourceMapping).bulk_create( + {(str(item.target_type) + str(item.target_id)): item for item in resource_mapping_list}.values() + ) + + +def get_tool_id_list(workflow, with_deep=False): + from tools.models import ToolType, ToolWorkflow + + _result = [] + for node in workflow.get("nodes", []): + if node.get("type") == "tool-lib-node": + tool_id = node.get("properties", {}).get("node_data", {}).get("tool_lib_id") + if tool_id: + _result.append(tool_id) + elif node.get("type") == "loop-node": + r = get_tool_id_list(node.get("properties", {}).get("node_data", {}).get("loop_body", {})) + for item in r: + _result.append(item) + elif node.get("type") == "tool-workflow-lib-node": + tool_id = node.get("properties", {}).get("node_data", {}).get("tool_lib_id") + if tool_id: + _result.append(tool_id) + elif node.get("type") == "ai-chat-node": + node_data = node.get("properties", {}).get("node_data", {}) + mcp_tool_ids = node_data.get("mcp_tool_ids") or [] + skill_tool_ids = node_data.get("skill_tool_ids") or [] + tool_ids = node_data.get("tool_ids") or [] + for _id in mcp_tool_ids + tool_ids + skill_tool_ids: + _result.append(_id) + elif node.get("type") == "mcp-node": + mcp_tool_id = node.get("properties", {}).get("node_data", {}).get("mcp_tool_id") + if mcp_tool_id: + _result.append(mcp_tool_id) + if with_deep: + workflow_list = QuerySet(Tool).filter(id__in=_result, tool_type=ToolType.WORKFLOW) + tool_work_flow_list = QuerySet(ToolWorkflow).filter(tool_id__in=[wl.id for wl in workflow_list]) + for tool_work_flow in tool_work_flow_list: + child_tool_id_list = get_child_tool_id_list(tool_work_flow.work_flow, []) + for c in child_tool_id_list: + _result.append(c) + return _result + + +def get_child_tool_id_list(work_flow, response): + from tools.models import ToolType, ToolWorkflow + + tool_id_list = get_tool_id_list(work_flow, False) + tool_id_list = [tool_id for tool_id in tool_id_list if len([r for r in response if r == tool_id]) == 0] + tool_list = [] + if len(tool_id_list) > 0: + tool_list = QuerySet(Tool).filter(id__in=tool_id_list).exclude(scope=ToolScope.SHARED) + work_flow_tools = [tool for tool in tool_list if tool.tool_type == ToolType.WORKFLOW] + if len(work_flow_tools) > 0: + work_flow_tool_dict = { + tw.tool_id: tw for tw in QuerySet(ToolWorkflow).filter(tool_id__in=[t.id for t in work_flow_tools]) + } + for tool in tool_list: + response.append(str(tool.id)) + if tool.tool_type == ToolType.WORKFLOW: + get_child_tool_id_list(work_flow_tool_dict.get(tool.id).work_flow, response) + else: + for tool in tool_list: + response.append(str(tool.id)) + return response + + +def build_schema(fields: dict): + return create_model("dynamicSchema", **fields) + + +def get_type(_type: str): + if _type == "float": + return float + if _type == "string": + return str + if _type == "int": + return int + if _type == "dict": + return dict + if _type == "array": + return list + if _type == "boolean": + return bool + return object + + +def get_workflow_args(tool, qv): + for node in qv.work_flow.get("nodes"): + if node.get("type") == "tool-base-node": + input_field_list = node.get("properties").get("user_input_field_list") + return build_schema( + { + field.get("field"): ( + get_type(field.get("type")), + Field(..., required=True, description=field.get("desc")) if field.get("is_required") else Field(default=None, required=False, description=field.get("desc")) + ) + for field in input_field_list + } + ) + + return build_schema({}) + + +def get_workflow_func(source_type, source_id, tool, qv, workspace_id, user_id=None): + tool_id = tool.id + tool_record_id = str(uuid.uuid7()) + took_execute = ToolExecute(tool_id, tool_record_id, workspace_id, source_type, source_id, False) + + def inner(**kwargs): + from application.flow.tool_workflow_manage import ToolWorkflowManage + + work_flow_manage = ToolWorkflowManage( + Workflow.new_instance(qv.work_flow, WorkflowMode.TOOL), + { + "chat_record_id": tool_record_id, + "tool_id": tool_id, + "stream": True, + "workspace_id": workspace_id, + "user_id": user_id, + **kwargs, + "default_model_setting": qv.default_model_setting, + }, + ToolWorkflowPostHandler(took_execute, tool_id), + is_the_task_interrupted=lambda: False, + child_node=None, + start_node_id=None, + start_node_data=None, + chat_record=None, + ) + res = work_flow_manage.run() + for r in res: + pass + return work_flow_manage.out_context + + return inner + + +def get_tools(source_type, source_id, tool_workflow_ids, workspace_id, user_id=None): + tools = QuerySet(Tool).filter( + id__in=tool_workflow_ids, is_active=True, tool_type=ToolType.WORKFLOW, workspace_id=workspace_id + ) + latest_subquery = ToolWorkflowVersion.objects.filter(tool_id=OuterRef("tool_id")).order_by("-create_time") + + qs = ToolWorkflowVersion.objects.filter( + tool_id__in=[t.id for t in tools], id=Subquery(latest_subquery.values("id")[:1]) + ) + qd = {q.tool_id: q for q in qs} + results = [] + for tool in tools: + qv = qd.get(tool.id) + func = get_workflow_func(source_type, source_id, tool, qv, workspace_id, user_id=user_id) + args = get_workflow_args(tool, qv) + tool = StructuredTool.from_function( + func=func, + name=tool.name, + description=tool.desc, + args_schema=args, + ) + results.append(tool) + + return results