From 09577b9b6589e367ac51ae50cd04d9f1a3f18aa5 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Mon, 2 Mar 2026 21:00:40 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E9=87=8D=E6=9E=84=E7=9F=A5?= =?UTF-8?q?=E8=AF=86=E5=BA=93=E5=B7=A5=E5=85=B7=E6=A8=A1=E5=9D=97=EF=BC=8C?= =?UTF-8?q?=E4=BC=98=E5=8C=96=E5=B7=A5=E5=85=B7=E8=8E=B7=E5=8F=96=E9=80=BB?= =?UTF-8?q?=E8=BE=91=E5=B9=B6=E7=A7=BB=E9=99=A4=E5=86=97=E4=BD=99=E4=BB=A3?= =?UTF-8?q?=E7=A0=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/agents/chatbot/graph.py | 14 +- src/agents/common/toolkits/kbs/__init__.py | 3 + src/agents/common/toolkits/kbs/tools.py | 159 +++++++++++++++++++++ src/agents/common/tools.py | 148 +------------------ 4 files changed, 174 insertions(+), 150 deletions(-) create mode 100644 src/agents/common/toolkits/kbs/__init__.py create mode 100644 src/agents/common/toolkits/kbs/tools.py diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index 633aade3..984d5a64 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -1,7 +1,11 @@ from deepagents.middleware.filesystem import FilesystemMiddleware -from langchain.agents import create_agent -from langchain.agents.middleware import ModelRetryMiddleware +from deepagents.middleware.patch_tool_calls import PatchToolCallsMiddleware +from langchain.agents import create_agent +from langchain.agents.middleware import ( + TodoListMiddleware, + ModelRetryMiddleware, +) from src.agents.common import BaseAgent, load_chat_model from src.agents.common.backends import create_agent_composite_backend from src.agents.common.middlewares import ( @@ -19,7 +23,7 @@ def _create_fs_backend(rt): class ChatbotAgent(BaseAgent): name = "智能体助手" description = "基础的对话机器人,可以回答问题,可在配置中启用需要的工具。" - capabilities = ["file_upload", "files"] # 支持文件上传功能 + capabilities = ["file_upload", "files", "todo"] # 支持文件上传功能 def __init__(self, **kwargs): super().__init__(**kwargs) @@ -34,13 +38,15 @@ class ChatbotAgent(BaseAgent): # 使用 create_agent 创建智能体 # 注意:tools 参数由 RuntimeConfigMiddleware 在 wrap_model_call 中动态设置 graph = create_agent( - model=load_chat_model(context.model), + model=load_chat_model(fully_specified_name=context.model), system_prompt=context.system_prompt, middleware=[ save_attachments_to_fs, # 附件注入提示词 FilesystemMiddleware(backend=_create_fs_backend), # 文件系统后端 RuntimeConfigMiddleware(extra_tools=all_mcp_tools), # 运行时配置应用(模型/工具/知识库/MCP/提示词) ModelRetryMiddleware(), # 模型重试中间件 + TodoListMiddleware(), + PatchToolCallsMiddleware(), ], checkpointer=await self._get_checkpointer(), ) diff --git a/src/agents/common/toolkits/kbs/__init__.py b/src/agents/common/toolkits/kbs/__init__.py new file mode 100644 index 00000000..531dcc91 --- /dev/null +++ b/src/agents/common/toolkits/kbs/__init__.py @@ -0,0 +1,3 @@ +from .tools import get_kb_based_tools + +__all__ = ["get_kb_based_tools"] diff --git a/src/agents/common/toolkits/kbs/tools.py b/src/agents/common/toolkits/kbs/tools.py new file mode 100644 index 00000000..2cf5c0fa --- /dev/null +++ b/src/agents/common/toolkits/kbs/tools.py @@ -0,0 +1,159 @@ +"""知识库工具模块""" +import inspect +import traceback +from typing import Any + +from langchain_core.tools import StructuredTool +from pydantic import BaseModel, Field + +from src import knowledge_base +from src.utils import logger + + +class KnowledgeRetrieverModel(BaseModel): + query_text: str | None = Field( + default=None, + description=( + "查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。" + ) + ) + operation: str = Field( + default="search", + description=( + "操作类型:'search' 表示检索知识库内容,'get_mindmap' 表示获取知识库的思维导图结构。" + "当用户询问知识库的整体结构、文件分类、知识架构时,使用 'get_mindmap'。" + "当用户需要查询具体内容时,使用 'search'。" + ), + ) + + +class CommonKnowledgeRetriever(KnowledgeRetrieverModel): + """Common knowledge retriever model.""" + + file_name: str | None = Field( + default=None, + description=( + "(非必要不启用此参数,留空即可)当操作类型为 'search' 且已经读取思维导图之后,可以指定文件关键词,支持模糊匹配。\n" + "仅当检索结果过多且不相关,需要进一步缩小范围时使用。" + ) + ) + + +def get_kb_based_tools(db_names: list[str] | None = None) -> list: + """获取所有知识库基于的工具""" + # 获取所有知识库 + kb_tools = [] + retrievers = knowledge_base.get_retrievers() + if db_names is None: + db_ids = None + else: + db_ids = [kb_id for kb_id, kb in retrievers.items() if kb["name"] in db_names] + + def _create_retriever_wrapper(db_id: str, retriever_info: dict[str, Any]): + """创建检索器包装函数的工厂函数,避免闭包变量捕获问题""" + + async def async_retriever_wrapper( + query_text: str, operation: str = "search", file_name: str | None = None + ) -> Any: + """异步检索器包装函数,支持检索和获取思维导图""" + + # 获取思维导图 + if operation == "get_mindmap": + try: + logger.debug(f"Getting mindmap for database {db_id}") + + from src.repositories.knowledge_base_repository import KnowledgeBaseRepository + + kb_repo = KnowledgeBaseRepository() + kb = await kb_repo.get_by_id(db_id) + + if kb is None: + return f"知识库 {retriever_info['name']} 不存在" + + mindmap_data = kb.mindmap + + if not mindmap_data: + return f"知识库 {retriever_info['name']} 还没有生成思维导图。" + + # 将思维导图数据转换为文本格式,便于AI理解 + def mindmap_to_text(node, level=0): + """递归将思维导图JSON转换为层级文本""" + indent = " " * level + text = f"{indent}- {node.get('content', '')}\n" + for child in node.get("children", []): + text += mindmap_to_text(child, level + 1) + return text + + mindmap_text = f"知识库 {retriever_info['name']} 的思维导图结构:\n\n" + mindmap_text += mindmap_to_text(mindmap_data) + + logger.debug(f"Successfully retrieved mindmap for {db_id}") + return mindmap_text + + except Exception as e: + logger.error(f"Error getting mindmap for {db_id}: {e}") + return f"获取思维导图失败: {str(e)}" + + # 默认:检索知识库 + retriever = retriever_info["retriever"] + try: + logger.debug(f"Retrieving from database {db_id} with query: {query_text}") + kwargs = {} + if file_name: + kwargs["file_name"] = file_name + + if inspect.iscoroutinefunction(retriever): + result = await retriever(query_text, **kwargs) + else: + result = retriever(query_text, **kwargs) + logger.debug(f"Retrieved {len(result) if isinstance(result, list) else 'N/A'} results from {db_id}") + return result + except Exception as e: + logger.error(f"Error in retriever {db_id}: {e}") + return f"检索失败: {str(e)}" + + return async_retriever_wrapper + + for db_id, retrieve_info in retrievers.items(): + if db_ids is not None and db_id not in db_ids: + continue + + try: + # 构建工具描述 + description = ( + f"使用 {retrieve_info['name']} 知识库的多功能工具。\n" + f"知识库描述:{retrieve_info['description'] or '没有描述。'}\n\n" + f"支持的操作:\n" + f"1. 'search' - 检索知识库内容:根据关键词查询相关文档片段\n" + f"2. 'get_mindmap' - 获取思维导图:查看知识库的整体结构和文件分类\n\n" + f"使用建议:\n" + f"- 需要查询具体内容时,使用 operation='search'\n" + f"- 想了解知识库结构、文件分类时,使用 operation='get_mindmap'" + ) + + # 使用工厂函数创建检索器包装函数,避免闭包问题 + retriever_wrapper = _create_retriever_wrapper(db_id, retrieve_info) + + safename = retrieve_info["name"].replace(" ", "_")[:20] + + args_schema = KnowledgeRetrieverModel + if retrieve_info["metadata"]["kb_type"] in ["milvus"]: + args_schema = CommonKnowledgeRetriever + + # 使用 StructuredTool.from_function 创建异步工具 + tool = StructuredTool.from_function( + coroutine=retriever_wrapper, + name=safename, + description=description, + args_schema=args_schema, + metadata=retrieve_info["metadata"] | {"tag": ["knowledgebase"]}, + ) + + kb_tools.append(tool) + # logger.debug(f"Successfully created tool {tool_id} for database {db_id}") + + except Exception as e: + logger.error(f"Failed to create tool for database {db_id}: {e}, \n{traceback.format_exc()}") + continue + + return kb_tools diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index b891644a..c8568cfc 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -1,4 +1,3 @@ -import asyncio import os import traceback import uuid @@ -6,11 +5,10 @@ from typing import Annotated, Any import requests from langchain.tools import tool -from langchain_core.tools import StructuredTool from langgraph.types import interrupt -from pydantic import BaseModel, Field -from src import config, graph_base, knowledge_base +from src import config, graph_base +from src.agents.common.toolkits.kbs import get_kb_based_tools from src.services.mcp_service import get_enabled_mcp_tools from src.storage.minio import aupload_file_to_minio from src.utils import logger @@ -147,148 +145,6 @@ def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge return f"知识图谱查询失败: {str(e)}" -class KnowledgeRetrieverModel(BaseModel): - query_text: str = Field( - description=( - "查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。" - ) - ) - operation: str = Field( - default="search", - description=( - "操作类型:'search' 表示检索知识库内容,'get_mindmap' 表示获取知识库的思维导图结构。" - "当用户询问知识库的整体结构、文件分类、知识架构时,使用 'get_mindmap'。" - "当用户需要查询具体内容时,使用 'search'。" - ), - ) - - -class CommonKnowledgeRetriever(KnowledgeRetrieverModel): - """Common knowledge retriever model.""" - - file_name: str = Field(description="限定文件名称,当操作类型为 'search' 时,可以指定文件名称,支持模糊匹配") - - -def get_kb_based_tools(db_names: list[str] | None = None) -> list: - """获取所有知识库基于的工具""" - # 获取所有知识库 - kb_tools = [] - retrievers = knowledge_base.get_retrievers() - if db_names is None: - db_ids = None - else: - db_ids = [kb_id for kb_id, kb in retrievers.items() if kb["name"] in db_names] - - def _create_retriever_wrapper(db_id: str, retriever_info: dict[str, Any]): - """创建检索器包装函数的工厂函数,避免闭包变量捕获问题""" - - async def async_retriever_wrapper( - query_text: str, operation: str = "search", file_name: str | None = None - ) -> Any: - """异步检索器包装函数,支持检索和获取思维导图""" - - # 获取思维导图 - if operation == "get_mindmap": - try: - logger.debug(f"Getting mindmap for database {db_id}") - - from src.repositories.knowledge_base_repository import KnowledgeBaseRepository - - kb_repo = KnowledgeBaseRepository() - kb = await kb_repo.get_by_id(db_id) - - if kb is None: - return f"知识库 {retriever_info['name']} 不存在" - - mindmap_data = kb.mindmap - - if not mindmap_data: - return f"知识库 {retriever_info['name']} 还没有生成思维导图。" - - # 将思维导图数据转换为文本格式,便于AI理解 - def mindmap_to_text(node, level=0): - """递归将思维导图JSON转换为层级文本""" - indent = " " * level - text = f"{indent}- {node.get('content', '')}\n" - for child in node.get("children", []): - text += mindmap_to_text(child, level + 1) - return text - - mindmap_text = f"知识库 {retriever_info['name']} 的思维导图结构:\n\n" - mindmap_text += mindmap_to_text(mindmap_data) - - logger.debug(f"Successfully retrieved mindmap for {db_id}") - return mindmap_text - - except Exception as e: - logger.error(f"Error getting mindmap for {db_id}: {e}") - return f"获取思维导图失败: {str(e)}" - - # 默认:检索知识库 - retriever = retriever_info["retriever"] - try: - logger.debug(f"Retrieving from database {db_id} with query: {query_text}") - kwargs = {} - if file_name: - kwargs["file_name"] = file_name - - if asyncio.iscoroutinefunction(retriever): - result = await retriever(query_text, **kwargs) - else: - result = retriever(query_text, **kwargs) - logger.debug(f"Retrieved {len(result) if isinstance(result, list) else 'N/A'} results from {db_id}") - return result - except Exception as e: - logger.error(f"Error in retriever {db_id}: {e}") - return f"检索失败: {str(e)}" - - return async_retriever_wrapper - - for db_id, retrieve_info in retrievers.items(): - if db_ids is not None and db_id not in db_ids: - continue - - try: - # 构建工具描述 - description = ( - f"使用 {retrieve_info['name']} 知识库的多功能工具。\n" - f"知识库描述:{retrieve_info['description'] or '没有描述。'}\n\n" - f"支持的操作:\n" - f"1. 'search' - 检索知识库内容:根据关键词查询相关文档片段\n" - f"2. 'get_mindmap' - 获取思维导图:查看知识库的整体结构和文件分类\n\n" - f"使用建议:\n" - f"- 需要查询具体内容时,使用 operation='search'\n" - f"- 想了解知识库结构、文件分类时,使用 operation='get_mindmap'" - ) - - # 使用工厂函数创建检索器包装函数,避免闭包问题 - retriever_wrapper = _create_retriever_wrapper(db_id, retrieve_info) - - safename = retrieve_info["name"].replace(" ", "_")[:20] - - args_schema = KnowledgeRetrieverModel - if retrieve_info["metadata"]["kb_type"] in ["milvus"]: - args_schema = CommonKnowledgeRetriever - - # 使用 StructuredTool.from_function 创建异步工具 - tool = StructuredTool.from_function( - coroutine=retriever_wrapper, - name=safename, - description=description, - args_schema=args_schema, - metadata=retrieve_info["metadata"] | {"tag": ["knowledgebase"]}, - ) - - kb_tools.append(tool) - # logger.debug(f"Successfully created tool {tool_id} for database {db_id}") - - except Exception as e: - logger.error(f"Failed to create tool for database {db_id}: {e}, \n{traceback.format_exc()}") - continue - - return kb_tools - - def gen_tool_info(tools) -> list[dict[str, Any]]: """获取所有工具的信息(用于前端展示)""" tools_info = []