diff --git a/apps/application/workflow/message/aggregator/impl/tool_aggregator.py b/apps/application/workflow/message/aggregator/impl/tool_aggregator.py index 029ddc0224b..f2b5a5b1e85 100644 --- a/apps/application/workflow/message/aggregator/impl/tool_aggregator.py +++ b/apps/application/workflow/message/aggregator/impl/tool_aggregator.py @@ -1,10 +1,11 @@ # coding=utf-8 """ - @project: MaxKB - @file: tool_aggregator.py - @date:2026/7/22 16:24 - @desc: ToolContent 聚合器 +@project: MaxKB +@file: tool_aggregator.py +@date:2026/7/22 16:24 +@desc: ToolContent 聚合器 """ + from application.workflow.message.aggregator.content_aggregator import ContentAggregator from application.workflow.message.struct.tool_content import ToolContent @@ -18,7 +19,7 @@ class ToolAggregator(ContentAggregator[ToolContent]): def aggregate(self, prev: ToolContent, chunk: ToolContent) -> ToolContent: """ 聚合工具内容 - + @param prev: 之前的内容 @param chunk: 新的内容块 @return: 合并后的内容 @@ -26,20 +27,20 @@ def aggregate(self, prev: ToolContent, chunk: ToolContent) -> ToolContent: if prev is None: return chunk - # 合并 content (tool_name) - prev_content = prev.content if prev.content else "" - chunk_content = chunk.content if chunk.content else "" - merged_content = chunk_content if chunk_content else prev_content + # 合并 name (tool_name):取新回退旧 + prev_name = prev.name if prev.name else "" + chunk_name = chunk.name if chunk.name else "" + merged_name = chunk_name if chunk_name else prev_name - # 合并 arguments + # 合并 arguments:拼接 prev_arguments = prev.arguments if prev.arguments else "" chunk_arguments = chunk.arguments if chunk.arguments else "" merged_arguments = prev_arguments + chunk_arguments - # 合并 result - prev_result = prev.result if prev.result else "" - chunk_result = chunk.result if chunk.result else "" - merged_result = prev_result + chunk_result + # 合并 content(即 result 结果):拼接 + prev_content = prev.content if prev.content else "" + chunk_content = chunk.content if chunk.content else "" + merged_content = prev_content + chunk_content # 合并基础字段 merged_id = chunk.id if chunk.id else prev.id @@ -47,7 +48,9 @@ def aggregate(self, prev: ToolContent, chunk: ToolContent) -> ToolContent: merged_node_info = chunk.node_info if chunk.node_info else prev.node_info merged_position = chunk.position if chunk.position else prev.position - result = ToolContent(merged_id, merged_content, merged_arguments, merged_result, - merged_status, merged_node_info, merged_position) + # ToolContent(_id, tool_name, arguments, result, status, node_info, position) + result = ToolContent( + merged_id, merged_name, merged_arguments, merged_content, merged_status, merged_node_info, merged_position + ) return result diff --git a/apps/application/workflow/message/struct/tool_content.py b/apps/application/workflow/message/struct/tool_content.py index 3f89ea6d04b..e632df00406 100644 --- a/apps/application/workflow/message/struct/tool_content.py +++ b/apps/application/workflow/message/struct/tool_content.py @@ -1,27 +1,37 @@ # coding=utf-8 """ - @project: MaxKB - @Author:虎虎 - @file: tool_content.py - @date:2026/6/30 16:17 - @desc: +@project: MaxKB +@Author:虎虎 +@file: tool_content.py +@date:2026/6/30 16:17 +@desc: """ + from application.workflow.content_type import ContentType from application.workflow.message.struct.content import Content, NodeInfo, Position from application.workflow.status import Status class ToolContent(Content): - def __init__(self, _id, tool_name: str, arguments: str, result: str, status: Status, node_info: NodeInfo, - position: Position, **kwargs): - self.content = tool_name + def __init__( + self, + _id, + tool_name: str, + arguments: str, + result: str, + status: Status, + node_info: NodeInfo, + position: Position, + **kwargs, + ): + self.name = tool_name self.arguments = arguments - self.result = result + self.content = result super().__init__(_id, status, ContentType.TOOL, node_info, position, **kwargs) def to_dict(self): result = super().to_dict() - result['content'] = self.content - result['arguments'] = self.arguments - result['result'] = self.result + result["content"] = self.content + result["arguments"] = self.arguments + result["name"] = self.name return result diff --git a/apps/application/workflow/nodes/ai_chat_node/agent.py b/apps/application/workflow/nodes/ai_chat_node/agent.py index ea322db9d5f..9e8a561bd83 100644 --- a/apps/application/workflow/nodes/ai_chat_node/agent.py +++ b/apps/application/workflow/nodes/ai_chat_node/agent.py @@ -11,36 +11,22 @@ """ import asyncio -import io import json import os -import queue import re import shutil -import threading -import time -import zipfile import langchain_core.messages.ai as _lc_ai_module import uuid_utils.compat as uuid -from asgiref.sync import sync_to_async from deepagents import create_deep_agent -from django.db.models import OuterRef, QuerySet, Subquery -from langchain_core.messages import AIMessageChunk, ToolMessage -from langchain_core.tools import StructuredTool from langchain_core.utils._merge import merge_lists as _original_merge_lists from langchain_mcp_adapters.client import MultiServerMCPClient from langgraph.checkpoint.memory import MemorySaver -from pydantic import Field, create_model from application.workflow.backend.sandbox_shell import SandboxShellBackend -from application.workflow.message.aggregator import AggregationManager -from application.workflow.status import Status -from common.utils.logger import maxkb_logger -from knowledge.models import File -from knowledge.models.knowledge_action import State +from application.workflow.i_node import CancelledException +from application.workflow.nodes.ai_chat_node.tools.skill import init_skills from maxkb.const import CONFIG -from tools.models import Tool, ToolRecord, ToolType, ToolWorkflowVersion # --------------------------------------------------------------------------- @@ -81,49 +67,7 @@ def _norm(lst): _lc_ai_module.merge_lists = _merge_lists_normalize_empty_tool_chunk_ids -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""" +def _get_tool_call_id(raw_id): if not raw_id: return None if not isinstance(raw_id, str): @@ -147,68 +91,77 @@ def _extract_tool_id(raw_id): return tool_id or raw_id -async def _initialize_skills(mcp_servers, temp_dir): - skills_dir = os.path.join(temp_dir, "skills") - mcp_config = json.loads(mcp_servers) - 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) +class ToolCallStreamManagement: + def __init__(self): + self.index_id_map = {} + self.id_name_map = {} + self.tool_uuid_map = {} + self.use_tool_id_list = set() - # 获取技能解压后的顶级目录名 - top_level_dirs = set() - for member in members: - parts = member.split("/") - if parts[0]: - top_level_dirs.add(parts[0]) + @staticmethod + def get_fallback_tool_calls(msg): + source = msg.tool_calls or msg.invalid_tool_calls + if source: + return [(tc.get("index"), tc.get("id"), tc.get("name"), tc.get("args", "")) for tc in source] + result = [] + for tc in msg.additional_kwargs.get("tool_calls", []): + func = tc.get("function") + if isinstance(func, dict): + result.append((tc.get("index"), tc.get("id"), func.get("name"), func.get("arguments", ""))) + else: + result.append((tc.get("index"), tc.get("id"), tc.get("name"), tc.get("arguments", ""))) + return result - # 将 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) + def get_tool_id(self, index, raw_id): + if raw_id and str(raw_id).strip(): + tool_id = _get_tool_call_id(str(raw_id).strip()) + if index is not None: + self.index_id_map[index] = tool_id + return tool_id + if index is not None: + return self.index_id_map.get(index) + return None + + def get_tool_name(self, tool_id, default=None): + return self.id_name_map.get(tool_id, default) - os.system("chmod -R g+rx " + temp_dir) # 确保技能目录可访问 + def add_tool_id(self, tool_id): + self.use_tool_id_list.add(tool_id) - client = MultiServerMCPClient(mcp_config) + def get_tool_uuid(self, tool_id): + if tool_id not in self.tool_uuid_map: + self.tool_uuid_map[tool_id] = str(uuid.uuid7()) + return self.tool_uuid_map.get(tool_id) - return client + def tool_id_is_used(self, tool_id): + return tool_id in self.use_tool_id_list + def set_tool_id_name(self, tool_id, name): + self.id_name_map[tool_id] = name -async def _yield_mcp_response( + +def create_agent( chat_model, system_prompt, message_list, mcp_servers, - mcp_output_enable=True, - tool_init_params={}, - source_id=None, - source_type=None, - temp_dir=None, + call_back, chat_id=None, + skill_tool_ids=None, extra_tools=None, ): - try: + # 创建临时文件夹 + 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) + + async def _run(): checkpointer = MemorySaver() - client = await _initialize_skills(mcp_servers, temp_dir) + await init_skills(skill_tool_ids, temp_dir) + client = MultiServerMCPClient(json.loads(mcp_servers)) tools = await client.get_tools() for tool in tools: tool.handle_tool_error = True @@ -232,512 +185,25 @@ async def _yield_mcp_response( 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: - 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 - + msg = chunk[0] + call_back.on_next(msg) + + def _classify_error(e): + # 取消:原样保留(保持节点取消语义);MCP TaskGroup 的 ExceptionGroup:展开取真实异常并包成 RuntimeError + if isinstance(e, CancelledException): + return e + if isinstance(e, ExceptionGroup): + while isinstance(e, ExceptionGroup): + e = e.exceptions[0] + return RuntimeError(f"{type(e).__name__}: {str(e)}") + + error = None + try: + asyncio.run(_run()) 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): - 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) - - 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 - - -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 _save_workflow_tool_record( - tool_record_id, tool_id, workspace_id, source_type, source_id, wf_manage, aggregation, parameters, start_time, error -): - """ - 工具工作流执行结束后落库执行记录(替代旧引擎 ToolWorkflowPostHandler.handler)。 - 实实行(非调试)直接插入 ToolRecord,字段与工具记录查询端点保持一致。 - """ - workflow = wf_manage.workflow - base_node = workflow.get_node("tool-base-node") - input_field_list = base_node.properties.get("user_input_field_list", []) if base_node else [] - output_field_list = base_node.properties.get("user_output_field_list", []) if base_node else [] - input_data = {f.get("field"): parameters.get(f.get("field")) for f in input_field_list} - # 新引擎工具输出统一收口于全局 output 上下文(tool-start-node 初始化、变量赋值节点写入) - output = wf_manage.context.get("output", {}) - details = wf_manage.get_details() - if error: - state = State.FAILURE - else: - has_fail = any((d or {}).get("status") == Status.FAIL.value for d in (details or [])) - state = State.FAILURE if has_fail else State.SUCCESS - ToolRecord( - id=tool_record_id, - tool_id=tool_id, - workspace_id=workspace_id, - source_type=source_type, - source_id=source_id, - state=state, - run_time=time.time() - start_time, - meta={ - "input_field_list": input_field_list, - "output_field_list": output_field_list, - "input": input_data, - "output": output, - "details": details, - "answer_text_list": aggregation.get_contents(), - }, - ).save() - - -def get_workflow_func(source_type, source_id, tool, qv, workspace_id): - tool_id = tool.id - - def inner(**kwargs): - # 使用新工作流引擎执行工具工作流,方式与 tool_workflow_lib_node 保持一致。 - from application.workflow.common import WorkflowType, new_instance - from application.workflow.nodes import get_node_class - from application.workflow.workflow_manage import CallBack, WorkflowManage - - tool_record_id = str(uuid.uuid7()) - sub_workflow = new_instance(qv.work_flow, WorkflowType.TOOL) - start_time = time.time() - sub_parameters = { - "chat_record_id": tool_record_id, - "tool_id": str(tool_id), - "stream": True, - "workspace_id": workspace_id, - "default_model_setting": qv.default_model_setting or {}, - **kwargs, - } - - # WorkflowManage.run() 在后台线程异步执行节点,完成时机由 on_complete 回调驱动, - # 而 inner 作为 LangChain 同步工具函数必须阻塞到子工作流结束再返回其输出。 - aggregation = AggregationManager() - done_event = threading.Event() - result_holder = {"output": {}, "error": None} - - def on_next(wf_manage, content): - # 逐块聚合,用于执行记录的 answer_text_list(不直接转发给上游) - aggregation.aggregate(content) - - def on_complete(wf_manage, error): - try: - # 工具工作流输出统一写入 context['output'] - result_holder["output"] = dict(wf_manage.context.get("output", {}) or {}) - # 执行结束落库工具执行记录 - _save_workflow_tool_record( - tool_record_id, - tool_id, - workspace_id, - source_type, - source_id, - wf_manage, - aggregation, - sub_parameters, - start_time, - error, - ) - finally: - result_holder["error"] = error - done_event.set() - - call_back = CallBack(on_next, on_complete) - - def get_start_node_fn(wf, wm): - start_node = wf.get_node("tool-start-node") - node_class = get_node_class("tool-start-node", WorkflowType.TOOL) - return node_class(start_node, wm, lambda n: n.properties.get("node_data", {})) - - sub_manage = WorkflowManage( - workflow=sub_workflow, - parameters=sub_parameters, - workflow_type=WorkflowType.TOOL, - call_back=call_back, - get_start_node=get_start_node_fn, - ) - sub_manage.start_node.workflow_manage = sub_manage - sub_manage.run() - done_event.wait() - if result_holder["error"]: - raise result_holder["error"] - return result_holder["output"] - - return inner - - -def get_workflow_tools(source_type, source_id, tool_workflow_ids, workspace_id): - 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) - 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 + error = _classify_error(e) + finally: + # 清理临时文件夹 + shutil.rmtree(temp_dir, ignore_errors=True) + call_back.on_complete(error) diff --git a/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py b/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py index e468eefd5ae..86589127f3c 100644 --- a/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py +++ b/apps/application/workflow/nodes/ai_chat_node/ai_chat_node.py @@ -11,34 +11,46 @@ import json import re from functools import reduce +from typing import Callable, Optional import uuid_utils.compat as uuid from django.db.models import QuerySet from django.utils.translation import gettext_lazy as _ -from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage +from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage, AIMessageChunk from rest_framework import serializers -from application.workflow.message.aggregator import AggregationManager -from application.workflow.nodes.ai_chat_node.agent import get_workflow_tools, mcp_response_generator -from application.models import Application, ApplicationAccessToken, ApplicationApiKey from application.workflow.common import WorkflowType from application.workflow.i_node import INode +from application.workflow.message.aggregator import AggregationManager from application.workflow.message.struct.content import NodeInfo, Position, Content from application.workflow.message.struct.reasoning_content import ReasoningContent from application.workflow.message.struct.text_content import TextContent from application.workflow.message.struct.tool_content import ToolContent +from application.workflow.nodes.ai_chat_node.agent import create_agent, ToolCallStreamManagement, _get_tool_call_id +from application.workflow.nodes.ai_chat_node.tools import ( + get_application_tools, + get_mcp_servers, + get_tool_tools, +) from application.workflow.status import Status from application.workflow.tools import Reasoning -from common.exception.app_exception import AppApiException from common.utils.common import guess_image_format from common.utils.messages_util import to_ai_message_list, to_human_message_list -from common.utils.rsa_util import rsa_long_decrypt from common.utils.shared_resource_auth import filter_authorized_ids from common.utils.tool_code import ToolExecutor from knowledge.models import File 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 + + +class AgentCallBack: + def __init__( + self, + on_next: Callable[[any], None], + on_complete: Callable[[Optional[Exception]], None], + ): + self.on_next = on_next + self.on_complete = on_complete class ChatNodeSerializer(serializers.Serializer): @@ -313,13 +325,10 @@ def execute(self): chat_model, SystemMessage(system), message_list, - history_message, question, chat_id, workspace_id, workflow_type, - reasoning_content_id, - text_content_id, is_result, ) if not mcp_handled: @@ -467,37 +476,14 @@ def _handle_mcp( chat_model, system_prompt, message_list, - history_message, question, chat_id, workspace_id, workflow_type, - reasoning_content_id, text_content_id, is_result=False, ): - mcp_servers_config = {} - - 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 = {} + # 工具记录来源(source_type / source_id) if workflow_type == WorkflowType.KNOWLEDGE: source_id = self.get_workflow_parameters().get("knowledge_id") source_type = "KNOWLEDGE" @@ -508,118 +494,129 @@ def _handle_mcp( source_id = self.get_workflow_parameters().get("application_id") source_type = "APPLICATION" - tools = get_workflow_tools(source_type, chat_id, tool_ids, workspace_id) - if tool_ids and len(tool_ids) > 0: - 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: - 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 - ), + # 工具(workflow/custom) + 智能体(子应用) → LangChain tools; + # MCP(自定义/库内) → mcp_servers 配置;技能 → 交给引擎侧 init_skills 初始化 + tools = get_tool_tools(source_type, source_id, tool_ids, workspace_id) + get_application_tools( + source_type, source_id, application_ids, workspace_id, self.get_workflow_parameters() + ) + mcp_servers_config = get_mcp_servers(mcp_source, mcp_servers, mcp_tool_id, mcp_tool_ids, self._handle_variables) + ToolExecutor().validate_mcp_transport(json.dumps(mcp_servers_config)) + + if tools or mcp_servers_config or skill_tool_ids: + node_info = NodeInfo(self.get_node_id(), self.get_node_name(), Status.RUNNING) + # 使用可变状态在回调间共享(answer 累积、当前文本 content id、工具 content id 映射) + state = {"answer": "", "text_id": text_content_id, "tool_id_map": {}} + tool_stream = ToolCallStreamManagement() + node_info = NodeInfo(self.get_node_id(), self.get_node_name(), Status.RUNNING) + + def on_next(chunk): + self._check_cancelled() + if mcp_output_enable and isinstance(chunk, AIMessageChunk): + if chunk.tool_call_chunks: + for tc in chunk.tool_call_chunks: + tool_id = tool_stream.get_tool_id(tc.get("index"), tc.get("id")) + if not tool_id: + continue + tool_stream.add_tool_id(tool_id) + if tc.get("name") or tc.get("args"): + tool_stream.set_tool_id_name(tool_id, tc.get("name")) + self.write( + ToolContent( + tool_stream.get_tool_uuid(tool_id), + tc.get("name"), + tc.get("args"), + "", + Status.RUNNING, + node_info, + Position(self.get_node_id()), + ) + ) + else: + for index, raw_id, name, args in tool_stream.get_fallback_tool_calls(chunk): + tool_id = tool_stream.get_tool_id(index, raw_id) + if not tool_id or not tool_stream.add_tool_id(tool_id): + continue + self.write( + ToolContent( + tool_stream.get_tool_uuid(tool_id), + name, + args, + "", + Status.RUNNING, + node_info, + Position(self.get_node_id()), + ) + ) + + if mcp_output_enable and isinstance(chunk, ToolMessage): + tool_id = _get_tool_call_id(chunk.tool_call_id) or chunk.tool_call_id + chunk.name = tool_stream.get_tool_name(tool_id, chunk.name) + try: + if isinstance(chunk.content, str): + tool_result = json.loads(chunk.content) + elif isinstance(chunk.content, dict): + tool_result = chunk.content + elif isinstance(chunk.content, list): + tool_result = chunk.content[0] if len(chunk.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 + tool_result = ( + text_result if isinstance(text_result, str) else json.dumps(text_result, ensure_ascii=False) + ) + except Exception: + tool_result = chunk.content + result = ( + tool_result if isinstance(tool_result, str) else json.dumps(tool_result, ensure_ascii=False) + ) + self.write( + ToolContent( + tool_stream.get_tool_uuid(tool_id), + "", + "", + result, + Status.SUCCESS, + NodeInfo(self.get_node_id(), self.get_node_name(), Status.SUCCESS), + Position(self.get_node_id()), ) - 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: - 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 is_result and chunk.content: + self.write( + TextContent( + tool_stream.get_tool_uuid(chunk.id), + chunk.content, + Status.RUNNING, + node_info, + Position(self.get_node_id()), + ) + ) - if len(mcp_servers_config) > 0 or len(tools) > 0: - node_info = NodeInfo(self.get_node_id(), self.get_node_name(), Status.RUNNING) - tool_content_id = str(uuid.uuid7()) - r = mcp_response_generator( + def on_complete(error): + if error: + raise error + self._write_final_context(chat_model, message_list, question.content, state["answer"], "") + self.write( + TextContent( + tool_stream.get_tool_uuid("text"), + "", + Status.SUCCESS, + NodeInfo(self.get_node_id(), self.get_node_name(), Status.SUCCESS), + Position(self.get_node_id()), + ) + ) + + create_agent( chat_model, system_prompt, message_list, json.dumps(mcp_servers_config), - mcp_output_enable, - tool_init_params, - source_id, - source_type, + AgentCallBack(on_next, on_complete), chat_id, + skill_tool_ids, tools, ) - answer = "" - tool_calls_map = {} - for chunk in r: - self._check_cancelled() - if isinstance(chunk, ToolMessage): - tool_call = tool_calls_map.get(chunk.tool_call_id, {}) - self.write( - ToolContent( - tool_content_id, - tool_call.get("name", getattr(chunk, "name", "")), - json.dumps(tool_call.get("args", {}), ensure_ascii=False), - chunk.content, - Status.RUNNING, - node_info, - Position(self.get_node_id()), - ) - ) - continue - - if hasattr(chunk, "tool_calls") and chunk.tool_calls: - for tool_call in chunk.tool_calls: - tool_calls_map[tool_call.get("id", "")] = tool_call - - answer += chunk.content if hasattr(chunk, "content") else str(chunk) - if chunk.content: - self.write( - TextContent( - text_content_id, chunk.content, Status.RUNNING, node_info, Position(self.get_node_id()) - ) - ) - self._write_final_context(chat_model, message_list, question.content, answer, "") return True return False diff --git a/apps/application/workflow/nodes/ai_chat_node/tools/__init__.py b/apps/application/workflow/nodes/ai_chat_node/tools/__init__.py new file mode 100644 index 00000000000..03482beb2cd --- /dev/null +++ b/apps/application/workflow/nodes/ai_chat_node/tools/__init__.py @@ -0,0 +1,15 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author: 虎虎虎 +@file: __init__.py +@date: 2026/9/15 16:58 +@desc: +""" + +from .application import get_application_tools +from .mcp import get_mcp_servers +from .skill import init_skills +from .tool import get_tool_tools + +__all__ = ["get_tool_tools", "get_application_tools", "get_mcp_servers", "init_skills"] diff --git a/apps/application/workflow/nodes/ai_chat_node/tools/application.py b/apps/application/workflow/nodes/ai_chat_node/tools/application.py new file mode 100644 index 00000000000..272b1f70220 --- /dev/null +++ b/apps/application/workflow/nodes/ai_chat_node/tools/application.py @@ -0,0 +1,191 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author: 虎虎虎 +@file: application.py +@date: 2026/9/15 16:10 +@desc: +""" + +import re +import threading + +import uuid_utils.compat as uuid +from django.db.models import QuerySet +from langchain_core.tools import StructuredTool +from pydantic import Field + +from .base import build_schema + + +def _application_string_to_uuid(input_str): + return str(uuid.uuid5(uuid.NAMESPACE_DNS, input_str)) + + +def get_application_args(): + """ + 应用(Agent)工具对模型暴露的入参:固定单个必填 message。 + + 与 chat.mcp.tools.MCPToolHandler.list_tools 的 inputSchema 保持一致, + 这样从远端 MCP 代理切换为进程内直调时,模型侧契约不变。 + """ + return build_schema( + { + "message": (str, Field(..., required=True, description="The message to send to the AI.")), + } + ) + + +def get_application_func(source_type, source_id, application, workflow_params, workspace_id): + """ + 构建应用(Agent)工具的执行函数。 + + 工具调用是同步的,不支持子应用表单中断,命中表单时返回已累积文本。 + """ + application_id = str(application.id) + + def inner(message: str = ""): + from application.models import Application, ApplicationVersion, Chat, ChatRecord, ChatSourceChoices + from application.workflow.common import WorkflowType, new_instance + from application.workflow.content_type import ContentType + from application.workflow.nodes import get_start_node + from application.workflow.workflow_manage import CallBack, WorkflowManage + from chat.serializers.chat import get_work_flow + from chat.serializers.chat_history import ChatHistory + + question = str(message or "") + chat_id = workflow_params.get("chat_id") + chat_user_id = workflow_params.get("chat_user_id") + chat_user_type = workflow_params.get("chat_user_type") + ip_address = workflow_params.get("ip_address") or "-" + source = workflow_params.get("source") or {"type": ChatSourceChoices.ONLINE.value} + debug = workflow_params.get("debug", False) + + # 自引用守卫:子应用不能是当前应用本身 + if application_id == str(workflow_params.get("application_id") or ""): + raise Exception("The sub application cannot use the current agent") + + # 派生子应用聊天 id(父对话 + 子应用稳定映射),与 application_node 一致 + current_chat_id = _application_string_to_uuid(str(chat_id) + application_id) + asker = workflow_params.get("chat_user") + Chat.objects.get_or_create( + id=current_chat_id, + defaults={ + "application_id": application_id, + "abstract": question[0:1024], + "chat_user_id": chat_user_id, + "chat_user_type": chat_user_type, + "ip_address": ip_address, + "source": source, + "asker": asker, + }, + ) + + # 解析子应用工作流(debug 取本体,否则取最新发布版本) + if debug: + sub_application = QuerySet(Application).filter(id=application_id).first() + else: + sub_application = ( + QuerySet(ApplicationVersion).filter(application_id=application_id).order_by("-create_time")[0:1].first() + ) + if sub_application is None: + raise Exception("The application has not been published. Please use it after publishing.") + + sub_workflow = new_instance(get_work_flow(sub_application), WorkflowType.APPLICATION) + + # 生成子应用记录 id 并建 ChatRecord(子对话可追溯) + sub_chat_record_id = str(uuid.uuid7()) + QuerySet(ChatRecord).create( + id=sub_chat_record_id, + chat_id=current_chat_id, + problem_text=question[0:1024], + answer_text="", + details={}, + message_tokens=0, + answer_tokens=0, + answer_text_list=[[]], + index=0, + ip_address=ip_address or "", + source=source, + workflow_context={}, + question={"content": question}, + messages=[], + ) + + # 组装子应用参数(复制父工作流参数并覆盖子应用相关字段) + sub_parameters = dict(workflow_params) + sub_parameters.update( + { + "chat_id": current_chat_id, + "chat_record_id": sub_chat_record_id, + "application_id": application_id, + "question": question, + "stream": True, + "form_data": {}, + "position": None, + "history_chat_record": ChatHistory(current_chat_id).load(exclude_record_id=sub_chat_record_id), + "image_list": [], + "document_list": [], + "audio_list": [], + "video_list": [], + "default_model_setting": sub_application.default_model_setting or {}, + } + ) + + done_event = threading.Event() + result_holder = {"answer": "", "error": None} + + def on_next(wf_manage, content): + # 逐块聚合子应用文本回答(不直接转发给上游,作为工具结果一次性返回) + if content.type == ContentType.TEXT: + result_holder["answer"] += content.content or "" + + def on_complete(wf_manage, error): + try: + # 持久化子应用上下文,供后续追溯 + QuerySet(ChatRecord).filter(id=sub_chat_record_id).update(workflow_context=wf_manage.context) + finally: + result_holder["error"] = error + done_event.set() + + call_back = CallBack(on_next, on_complete) + + def get_start_node_fn(wf, wm): + return get_start_node(wf, wm, WorkflowType.APPLICATION, None) + + sub_manage = WorkflowManage( + sub_workflow, sub_parameters, WorkflowType.APPLICATION, call_back, get_start_node_fn + ) + sub_manage.start_node.workflow_manage = sub_manage + sub_manage.run() + done_event.wait() + if result_holder["error"]: + raise result_holder["error"] + + answer = result_holder["answer"] + # 去除 标签(与 MCPToolHandler.call_tool 一致) + answer = re.sub(r".*?", "", answer, flags=re.DOTALL) + return answer + + return inner + + +def get_application_tools(source_type, source_id, application_ids, workspace_id, workflow_params): + if not application_ids: + return [] + from application.models import Application + + applications = QuerySet(Application).filter(id__in=application_ids, is_publish=True) + results = [] + for application in applications: + func = get_application_func(source_type, source_id, application, workflow_params, workspace_id) + args = get_application_args() + structured_tool = StructuredTool.from_function( + func=func, + name=application.name, + description=f"{application.name} {application.desc or ''}".strip(), + args_schema=args, + ) + results.append(structured_tool) + + return results diff --git a/apps/application/workflow/nodes/ai_chat_node/tools/base.py b/apps/application/workflow/nodes/ai_chat_node/tools/base.py new file mode 100644 index 00000000000..ac13933f716 --- /dev/null +++ b/apps/application/workflow/nodes/ai_chat_node/tools/base.py @@ -0,0 +1,30 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author: 虎虎虎 +@file: base.py +@date: 2026/9/15 +@desc: +""" + +from pydantic import create_model + + +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 diff --git a/apps/application/workflow/nodes/ai_chat_node/tools/mcp.py b/apps/application/workflow/nodes/ai_chat_node/tools/mcp.py new file mode 100644 index 00000000000..9f3c44e4708 --- /dev/null +++ b/apps/application/workflow/nodes/ai_chat_node/tools/mcp.py @@ -0,0 +1,37 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author: 虎虎虎 +@file: mcp.py +@date: 2026/9/15 +@desc: +""" + +import json + +from django.db.models import QuerySet + +from tools.models import Tool + + +def get_mcp_servers(mcp_source, mcp_servers, mcp_tool_id, mcp_tool_ids, handle_variables): + """ + tool-mcp-custom:mcp_source == "custom" 时用节点传入的自定义 MCP JSON。 + tool-mcp:否则用库内 MCP 工具(Tool.code 存 MCP server 配置)。 + """ + 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])) + + mcp_servers_config = {} + if mcp_source == "custom" and mcp_servers: + mcp_servers_config = handle_variables(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 = handle_variables({**mcp_servers_config, **json.loads(mcp_tool["code"])}) + return mcp_servers_config diff --git a/apps/application/workflow/nodes/ai_chat_node/tools/skill.py b/apps/application/workflow/nodes/ai_chat_node/tools/skill.py new file mode 100644 index 00000000000..42d362ed48d --- /dev/null +++ b/apps/application/workflow/nodes/ai_chat_node/tools/skill.py @@ -0,0 +1,66 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author: 虎虎虎 +@file: skill.py +@date: 2026/9/15 +@desc: +""" + +import io +import json +import os +import zipfile + +from asgiref.sync import sync_to_async +from django.db.models import QuerySet + +from common.utils.rsa_util import rsa_long_decrypt +from knowledge.models import File +from tools.models import Tool + + +async def init_skills(skill_tool_ids, temp_dir): + if not skill_tool_ids: + return + skills_dir = os.path.join(temp_dir, "skills") + tools = await sync_to_async(lambda: list(QuerySet(Tool).filter(id__in=skill_tool_ids, is_active=True)))() + if not tools: + return + + for tool in tools: + init_params_default_value = {i["field"]: i.get("default_value") for i in (tool.init_field_list or [])} + 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 + + file = await sync_to_async(lambda t=tool: QuerySet(File).filter(id=t.code).first())() + if not file: + continue + file_bytes = await sync_to_async(file.get_bytes)() + + 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 = [f"{key}={value}" for key, value in params.items()] + 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) # 确保技能目录可访问 diff --git a/apps/application/workflow/nodes/ai_chat_node/tools/tool/__init__.py b/apps/application/workflow/nodes/ai_chat_node/tools/tool/__init__.py new file mode 100644 index 00000000000..864f81d4b90 --- /dev/null +++ b/apps/application/workflow/nodes/ai_chat_node/tools/tool/__init__.py @@ -0,0 +1,26 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author: 虎虎虎 +@file: __init__.py +@date: 2026/9/15 +@desc: +""" + +from .custom import get_custom_tools +from .workflow import get_workflow_tools + +__all__ = ["get_tool_tools", "get_workflow_tools", "get_custom_tools"] + + +def get_tool_tools(source_type, source_id, tool_ids, workspace_id): + """ + 构建工具(Tool)类工具:内部按 tool_type 拆分 workflow / custom,合并返回 LangChain tools。 + + 节点只需传入混合的 tool_ids,各构建器各自按 tool_type 过滤。 + """ + if not tool_ids: + return [] + return get_workflow_tools(source_type, source_id, tool_ids, workspace_id) + get_custom_tools( + source_type, source_id, tool_ids, workspace_id + ) diff --git a/apps/application/workflow/nodes/ai_chat_node/tools/tool/custom.py b/apps/application/workflow/nodes/ai_chat_node/tools/tool/custom.py new file mode 100644 index 00000000000..a665f4b132a --- /dev/null +++ b/apps/application/workflow/nodes/ai_chat_node/tools/tool/custom.py @@ -0,0 +1,117 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author: 虎虎虎 +@file: custom.py +@date: 2026/9/15 +@desc: +""" + +import json +import time + +import uuid_utils.compat as uuid +from django.db.models import QuerySet +from langchain_core.tools import StructuredTool +from pydantic import Field + +from knowledge.models.knowledge_action import State +from tools.models import Tool, ToolRecord, ToolType + +from ..base import build_schema, get_type + + +def get_custom_args(tool): + """ + 从 CUSTOM 工具的 input_field_list 显式构建给模型的 args_schema。 + + input_field_list 项结构:{name, is_required, type(string|int|dict|array|float), source} + """ + input_field_list = tool.input_field_list or [] + return build_schema( + { + field.get("name"): ( + 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 + } + ) + + +def _save_custom_tool_record(tool_id, workspace_id, source_type, source_id, input_params, output, start_time, error): + """ + CUSTOM 工具执行结束后落库执行记录(替代原 MCP 路径的 save_tool_record)。 + + input 仅记录业务入参(不含 init 参数),避免密文/密钥进入记录。 + """ + state = State.FAILURE if error else State.SUCCESS + ToolRecord( + id=uuid.uuid7(), + tool_id=tool_id, + workspace_id=workspace_id, + source_type=source_type, + source_id=source_id, + state=state, + run_time=time.time() - start_time, + meta={ + "input": input_params, + "output": str(error) if error else output, + }, + ).save() + + +def get_custom_func(source_type, source_id, tool, workspace_id): + tool_id = tool.id + code = tool.code + init_field_list = tool.init_field_list or [] + init_params_ciphertext = tool.init_params + + def inner(**kwargs): + # 在进程内直接跑沙箱代码(无 MCP 子进程),方式与工具调试执行 ToolExecutor.exec_code 一致。 + from common.utils.rsa_util import rsa_long_decrypt + from common.utils.tool_code import ToolExecutor + + start_time = time.time() + # 合并初始化参数(默认值 → 已保存的启动参数),服务端注入,模型不可见 + init_params_default_value = {i["field"]: i.get("default_value") for i in init_field_list} + if init_params_ciphertext is not None: + init_params = init_params_default_value | json.loads(rsa_long_decrypt(init_params_ciphertext)) + else: + init_params = init_params_default_value + all_params = init_params | kwargs + + error = None + result = None + try: + result = ToolExecutor().exec_code(code, all_params) + except Exception as e: + error = e + finally: + _save_custom_tool_record(tool_id, workspace_id, source_type, source_id, kwargs, result, start_time, error) + if error: + raise error + return result + + return inner + + +def get_custom_tools(source_type, source_id, tool_ids, workspace_id): + if not tool_ids: + return [] + tools = QuerySet(Tool).filter(id__in=tool_ids, is_active=True, tool_type=ToolType.CUSTOM) + results = [] + for tool in tools: + func = get_custom_func(source_type, source_id, tool, workspace_id) + args = get_custom_args(tool) + structured_tool = StructuredTool.from_function( + func=func, + name=tool.name, + description=tool.desc, + args_schema=args, + ) + results.append(structured_tool) + + return results diff --git a/apps/application/workflow/nodes/ai_chat_node/tools/tool/workflow.py b/apps/application/workflow/nodes/ai_chat_node/tools/tool/workflow.py new file mode 100644 index 00000000000..0a9d785e6dc --- /dev/null +++ b/apps/application/workflow/nodes/ai_chat_node/tools/tool/workflow.py @@ -0,0 +1,183 @@ +# coding=utf-8 +""" +@project: MaxKB +@Author: 虎虎虎 +@file: workflow.py +@date: 2026/9/15 +@desc: +""" + +import threading +import time + +import uuid_utils.compat as uuid +from django.db.models import OuterRef, QuerySet, Subquery +from langchain_core.tools import StructuredTool +from pydantic import Field + +from application.workflow.message.aggregator import AggregationManager +from application.workflow.status import Status +from knowledge.models.knowledge_action import State +from tools.models import Tool, ToolRecord, ToolType, ToolWorkflowVersion + +from ..base import build_schema, get_type + + +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 _save_workflow_tool_record( + tool_record_id, tool_id, workspace_id, source_type, source_id, wf_manage, aggregation, parameters, start_time, error +): + """ + 工具工作流执行结束后落库执行记录(替代旧引擎 ToolWorkflowPostHandler.handler)。 + 实实行(非调试)直接插入 ToolRecord,字段与工具记录查询端点保持一致。 + """ + workflow = wf_manage.workflow + base_node = workflow.get_node("tool-base-node") + input_field_list = base_node.properties.get("user_input_field_list", []) if base_node else [] + output_field_list = base_node.properties.get("user_output_field_list", []) if base_node else [] + input_data = {f.get("field"): parameters.get(f.get("field")) for f in input_field_list} + # 新引擎工具输出统一收口于全局 output 上下文(tool-start-node 初始化、变量赋值节点写入) + output = wf_manage.context.get("output", {}) + details = wf_manage.get_details() + if error: + state = State.FAILURE + else: + has_fail = any((d or {}).get("status") == Status.FAIL.value for d in (details or [])) + state = State.FAILURE if has_fail else State.SUCCESS + ToolRecord( + id=tool_record_id, + tool_id=tool_id, + workspace_id=workspace_id, + source_type=source_type, + source_id=source_id, + state=state, + run_time=time.time() - start_time, + meta={ + "input_field_list": input_field_list, + "output_field_list": output_field_list, + "input": input_data, + "output": output, + "details": details, + "answer_text_list": aggregation.get_contents(), + }, + ).save() + + +def get_workflow_func(source_type, source_id, tool, qv, workspace_id): + tool_id = tool.id + + def inner(**kwargs): + # 使用新工作流引擎执行工具工作流,方式与 tool_workflow_lib_node 保持一致。 + from application.workflow.common import WorkflowType, new_instance + from application.workflow.nodes import get_node_class + from application.workflow.workflow_manage import CallBack, WorkflowManage + + tool_record_id = str(uuid.uuid7()) + sub_workflow = new_instance(qv.work_flow, WorkflowType.TOOL) + start_time = time.time() + sub_parameters = { + "chat_record_id": tool_record_id, + "tool_id": str(tool_id), + "stream": True, + "workspace_id": workspace_id, + "default_model_setting": qv.default_model_setting or {}, + **kwargs, + } + + # WorkflowManage.run() 在后台线程异步执行节点,完成时机由 on_complete 回调驱动, + # 而 inner 作为 LangChain 同步工具函数必须阻塞到子工作流结束再返回其输出。 + aggregation = AggregationManager() + done_event = threading.Event() + result_holder = {"output": {}, "error": None} + + def on_next(wf_manage, content): + # 逐块聚合,用于执行记录的 answer_text_list(不直接转发给上游) + aggregation.aggregate(content) + + def on_complete(wf_manage, error): + try: + # 工具工作流输出统一写入 context['output'] + result_holder["output"] = dict(wf_manage.context.get("output", {}) or {}) + # 执行结束落库工具执行记录 + _save_workflow_tool_record( + tool_record_id, + tool_id, + workspace_id, + source_type, + source_id, + wf_manage, + aggregation, + sub_parameters, + start_time, + error, + ) + finally: + result_holder["error"] = error + done_event.set() + + call_back = CallBack(on_next, on_complete) + + def get_start_node_fn(wf, wm): + start_node = wf.get_node("tool-start-node") + node_class = get_node_class("tool-start-node", WorkflowType.TOOL) + return node_class(start_node, wm, lambda n: n.properties.get("node_data", {})) + + sub_manage = WorkflowManage( + workflow=sub_workflow, + parameters=sub_parameters, + workflow_type=WorkflowType.TOOL, + call_back=call_back, + get_start_node=get_start_node_fn, + ) + sub_manage.start_node.workflow_manage = sub_manage + sub_manage.run() + done_event.wait() + if result_holder["error"]: + raise result_holder["error"] + return result_holder["output"] + + return inner + + +def get_workflow_tools(source_type, source_id, tool_workflow_ids, workspace_id): + 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) + 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 diff --git a/ui/src/components/conversation/content/items/tool.vue b/ui/src/components/conversation/content/items/tool.vue index d34225d62a3..b27e9eef1f2 100644 --- a/ui/src/components/conversation/content/items/tool.vue +++ b/ui/src/components/conversation/content/items/tool.vue @@ -1,28 +1,79 @@ diff --git a/ui/src/components/conversation/index.ts b/ui/src/components/conversation/index.ts index c500cf2edd2..b9e112894c4 100644 --- a/ui/src/components/conversation/index.ts +++ b/ui/src/components/conversation/index.ts @@ -79,8 +79,8 @@ const TOOL = (prev: any, chunk: any) => { return { type: 'TOOL', id: chunk.id ?? prev.id, - toolName: chunk.toolName, - functionArguments: (prev.functionArguments || '') + (chunk.functionArguments || ''), + name: chunk.name, + arguments: (prev.arguments || '') + (chunk.arguments || ''), content: (prev.content || '') + (chunk.content || ''), status: chunk.status ?? prev.status, workflowRunId: chunk.workflowRunId ?? prev.workflowRunId,