From 339cd0f3b21f4438e2d0937bb0804f2674c93029 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Thu, 15 Jan 2026 16:04:19 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E9=87=8D=E6=9E=84=E5=B7=A5?= =?UTF-8?q?=E5=85=B7=E8=8E=B7=E5=8F=96=E9=80=BB=E8=BE=91=EF=BC=8C=E7=A7=BB?= =?UTF-8?q?=E9=99=A4=E8=BF=87=E6=97=B6=E7=9A=84=E5=B7=A5=E5=85=B7=E6=8E=A5?= =?UTF-8?q?=E5=8F=A3=EF=BC=8C=E4=BC=98=E5=8C=96=E4=B8=8A=E4=B8=8B=E6=96=87?= =?UTF-8?q?=E5=B7=A5=E5=85=B7=E7=AE=A1=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/latest/advanced/agents-config.md | 62 ++++++++++- server/routers/chat_router.py | 21 ---- src/agents/__init__.py | 2 +- src/agents/chatbot/context.py | 42 -------- src/agents/chatbot/graph.py | 36 +------ src/agents/chatbot/tools.py | 56 ---------- src/agents/common/context.py | 35 +++++++ src/agents/common/subagents/calc_agent.py | 2 +- src/agents/common/tools.py | 122 ++++++++++++++++------ src/agents/mini_agent/graph.py | 11 +- src/agents/reporter/graph.py | 34 ++++-- src/services/evaluation_service.py | 4 +- 12 files changed, 225 insertions(+), 202 deletions(-) delete mode 100644 src/agents/chatbot/context.py delete mode 100644 src/agents/chatbot/tools.py diff --git a/docs/latest/advanced/agents-config.md b/docs/latest/advanced/agents-config.md index 8c50d3ea..85156132 100644 --- a/docs/latest/advanced/agents-config.md +++ b/docs/latest/advanced/agents-config.md @@ -36,7 +36,67 @@ <<< @/../src/agents/reporter/graph.py +### 工具系统 +系统提供统一的工具获取函数 `get_tools_from_context(context)`,自动从上下文配置中组装工具列表: + +```python +from src.agents.common.tools import get_tools_from_context + +async def get_graph(self, **kwargs): + context = self.get_context() + tools = await get_tools_from_context(context) + # tools 已包含:基础工具、知识库工具、MCP 工具 +``` + +该函数会自动处理三类工具的组装: +1. **基础工具**: 从 `context.tools` 筛选的内置工具 +2. **知识库工具**: 根据 `context.knowledges` 自动生成检索工具 +3. **MCP 工具**: 根据 `context.mcps` 加载并过滤的 MCP 服务器工具 + +### BaseContext 配置字段 + +`BaseContext` 已内置以下常用配置字段,所有智能体可直接复用: + +| 字段 | 类型 | 说明 | +|------|------|------| +| `model` | str | 使用的 LLM 模型 | +| `system_prompt` | str | 系统提示词 | +| `tools` | list[str] | 启用的内置工具列表 | +| `knowledges` | list[str] | 关联的知识库列表 | +| `mcps` | list[str] | 启用的 MCP 服务器名称 | + +```python +from src.agents.common import BaseContext + +@dataclass(kw_only=True) +class MyAgentContext(BaseContext): + # 继承所有 BaseContext 字段 + # 可在此添加智能体特有的额外配置 + custom_field: str = "默认值" +``` + +如需自定义工具选项(如 ReporterAgent 的 MySQL 工具),可覆盖 `tools` 字段的 `options` 元数据: + +```python +from src.agents.common import BaseContext, gen_tool_info +from src.agents.common.tools import get_buildin_tools +from src.agents.common.toolkits.mysql import get_mysql_tools + +@dataclass(kw_only=True) +class ReporterContext(BaseContext): + tools: Annotated[list[dict], {"__template_metadata__": {"kind": "tools"}}] = field( + default_factory=lambda: [t.name for t in get_mysql_tools()], + metadata={ + "name": "工具", + "options": lambda: gen_tool_info(get_buildin_tools() + get_mysql_tools()), + "description": "包含内置工具和 MySQL 工具包。", + }, + ) + + def __post_init__(self): + self.mcps = ["mcp-server-chart"] # 默认启用图表 MCP +``` 智能体实例的生命周期交给管理器处理,会在自动发现时完成初始化并缓存单例,以便快速响应请求。在容器内热重载时,只要保存文件即可触发重新导入;需要强制刷新可调用 `agent_manager.get_agent(, reload=True)`。 @@ -81,7 +141,7 @@ from src.agents.common.middlewares import inject_attachment_context async def get_graph(self): graph = create_agent( model=load_chat_model("..."), - tools=get_tools(), + tools=tools, middleware=[ inject_attachment_context, # 添加附件中间件 context_aware_prompt, # 其他中间件... diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index 3d005689..95ae1e4f 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -19,7 +19,6 @@ from server.utils.auth_middleware import get_db, get_required_user from src import executor from src import config as conf from src.agents import agent_manager -from src.agents.common.tools import gen_tool_info, get_buildin_tools from src.models import select_model from src.plugins.guard import content_guard from src.services.doc_converter import ( @@ -728,26 +727,6 @@ async def update_chat_models(model_provider: str, model_names: list[str], curren return {"models": conf.model_names[model_provider].models} -@chat.get("/tools") -async def get_tools(agent_id: str, current_user: User = Depends(get_required_user)): - """获取所有可用工具(需要登录)""" - logger.error("[DEPRECATED] 该接口已被弃用,将在未来版本中移除") - # 获取Agent实例和配置类 - if not (agent := agent_manager.get_agent(agent_id)): - raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在") - - if hasattr(agent, "get_tools") and callable(agent.get_tools): - if asyncio.iscoroutinefunction(agent.get_tools): - tools = await agent.get_tools() - else: - tools = agent.get_tools() - else: - tools = get_buildin_tools() - - tools_info = gen_tool_info(tools) - return {"tools": {tool["id"]: tool for tool in tools_info}} - - @chat.post("/agent/{agent_id}/resume") async def resume_agent_chat( agent_id: str, diff --git a/src/agents/__init__.py b/src/agents/__init__.py index 89bbe0ab..abbe0bda 100644 --- a/src/agents/__init__.py +++ b/src/agents/__init__.py @@ -54,7 +54,7 @@ class AgentManager(metaclass=SingletonMeta): # 遍历所有子目录 for item in agents_dir.iterdir(): - logger.info(f"尝试导入模块:{item}") + # logger.info(f"尝试导入模块:{item}") # 跳过非目录、common 目录、__pycache__ 等 if not item.is_dir() or item.name.startswith("_") or item.name == "common": continue diff --git a/src/agents/chatbot/context.py b/src/agents/chatbot/context.py deleted file mode 100644 index 2f772e36..00000000 --- a/src/agents/chatbot/context.py +++ /dev/null @@ -1,42 +0,0 @@ -from dataclasses import dataclass, field -from typing import Annotated - -from src.agents.common import BaseContext, gen_tool_info -from src.knowledge import knowledge_base -from src.services.mcp_service import get_mcp_server_names - -from .tools import get_tools - - -@dataclass(kw_only=True) -class Context(BaseContext): - tools: Annotated[list[dict], {"__template_metadata__": {"kind": "tools"}}] = field( - default_factory=list, - metadata={ - "name": "工具", - "options": lambda: gen_tool_info(get_tools()), # 这里的选择是所有的工具 - "description": "内置的部分工具,包含 common 工具和本智能体的特有工具(不含 MCP)。", - }, - ) - - knowledges: list[str] = field( - default_factory=list, - metadata={ - "name": "知识库", - "options": lambda: [k["name"] for k in knowledge_base.get_retrievers().values()], - "description": "知识库列表,可以在左侧知识库页面中创建知识库。", - "type": "list", # Explicitly mark as list type for frontend if needed - }, - ) - - mcps: list[str] = field( - default_factory=list, - metadata={ - "name": "MCP服务器", - "options": lambda: get_mcp_server_names(), - "description": ( - "MCP服务器列表,建议使用支持 SSE 的 MCP 服务器," - "如果需要使用 uvx 或 npx 运行的服务器,也请在项目外部启动 MCP 服务器,并在项目中配置 MCP 服务器。" - ), - }, - ) diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index ce6ce771..810bdb1a 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -5,46 +5,16 @@ from src.agents.common import BaseAgent, load_chat_model from src.agents.common.middlewares import ( inject_attachment_context, ) -from src.agents.common.tools import get_kb_based_tools -from src.services.mcp_service import get_enabled_mcp_tools - -from .context import Context -from .tools import get_tools +from src.agents.common.tools import get_tools_from_context class ChatbotAgent(BaseAgent): name = "智能体助手" - description = "基础的对话机器人,可以回答问题,默认不使用任何工具,可在配置中启用需要的工具。" + description = "基础的对话机器人,可以回答问题,可在配置中启用需要的工具。" capabilities = ["file_upload"] # 支持文件上传功能 def __init__(self, **kwargs): super().__init__(**kwargs) - self.context_schema = Context - - async def get_tools(self, tools: list[str] = None, mcps=None, knowledges=None): - # 1. 基础工具 (从 context.tools 中筛选) - all_basic_tools = get_tools() - selected_tools = [] - - if tools: - # 创建工具映射表 - tools_map = {t.name: t for t in all_basic_tools} - for tool_name in tools: - if tool_name in tools_map: - selected_tools.append(tools_map[tool_name]) - - # 2. 知识库工具 - if knowledges: - kb_tools = get_kb_based_tools(db_names=knowledges) - selected_tools.extend(kb_tools) - - # 3. MCP 工具(使用统一入口,自动过滤 disabled_tools) - if mcps: - for server_name in mcps: - mcp_tools = await get_enabled_mcp_tools(server_name) - selected_tools.extend(mcp_tools) - - return selected_tools async def get_graph(self, **kwargs): """构建图""" @@ -57,7 +27,7 @@ class ChatbotAgent(BaseAgent): # 使用 create_agent 创建智能体 graph = create_agent( model=load_chat_model(context.model), # 使用 context 中的模型配置 - tools=await self.get_tools(context.tools, context.mcps, context.knowledges), + tools=await get_tools_from_context(context), system_prompt=context.system_prompt, middleware=[ inject_attachment_context, # 附件上下文注入 diff --git a/src/agents/chatbot/tools.py b/src/agents/chatbot/tools.py deleted file mode 100644 index ead197cd..00000000 --- a/src/agents/chatbot/tools.py +++ /dev/null @@ -1,56 +0,0 @@ -import os -import uuid -from typing import Any - -import requests -from langchain.tools import tool - -from src.agents.common import get_buildin_tools -from src.agents.common.subagents import calc_agent_tool -from src.storage.minio import aupload_file_to_minio -from src.utils import logger - - -@tool -async def text_to_img_demo(text: str) -> str: - """【测试用】使用模型生成图片, 会返回图片的URL""" - - url = "https://api.siliconflow.cn/v1/images/generations" - - payload = { - "model": "Qwen/Qwen-Image", - "prompt": text, - } - headers = {"Authorization": f"Bearer {os.getenv('SILICONFLOW_API_KEY')}", "Content-Type": "application/json"} - - try: - response = requests.post(url, json=payload, headers=headers) - response_json = response.json() - except Exception as e: - logger.error(f"Failed to generate image with: {e}") - raise ValueError(f"Image generation failed: {e}") - - try: - image_url = response_json["images"][0]["url"] - except (KeyError, IndexError, TypeError) as e: - logger.error(f"Failed to parse image URL from response: {e}, {response_json=}") - raise ValueError(f"Image URL extraction failed: {e}") - - # 2. Upload to MinIO (Simplified) - response = requests.get(image_url) - file_data = response.content - - file_name = f"{uuid.uuid4()}.jpg" - image_url = await aupload_file_to_minio( - bucket_name="generated-images", file_name=file_name, data=file_data, file_extension="jpg" - ) - logger.info(f"Image uploaded. URL: {image_url}") - return image_url - - -def get_tools() -> list[Any]: - """获取所有可运行的工具(给大模型使用)""" - tools = get_buildin_tools() - tools.append(text_to_img_demo) - tools.append(calc_agent_tool) - return tools diff --git a/src/agents/common/context.py b/src/agents/common/context.py index 789bbd60..190fb7b2 100644 --- a/src/agents/common/context.py +++ b/src/agents/common/context.py @@ -9,8 +9,12 @@ from typing import Annotated, get_args, get_origin import yaml from src import config as sys_config +from src.knowledge import knowledge_base +from src.services.mcp_service import get_mcp_server_names from src.utils import logger +from .tools import gen_tool_info, get_buildin_tools + @dataclass(kw_only=True) class BaseContext: @@ -53,6 +57,37 @@ class BaseContext: }, ) + tools: Annotated[list[dict], {"__template_metadata__": {"kind": "tools"}}] = field( + default_factory=list, + metadata={ + "name": "工具", + "options": lambda: gen_tool_info(get_buildin_tools()), + "description": "内置的工具。", + }, + ) + + knowledges: list[str] = field( + default_factory=list, + metadata={ + "name": "知识库", + "options": lambda: [k["name"] for k in knowledge_base.get_retrievers().values()], + "description": "知识库列表,可以在左侧知识库页面中创建知识库。", + "type": "list", # Explicitly mark as list type for frontend if needed + }, + ) + + mcps: list[str] = field( + default_factory=list, + metadata={ + "name": "MCP服务器", + "options": lambda: get_mcp_server_names(), + "description": ( + "MCP服务器列表,建议使用支持 SSE 的 MCP 服务器," + "如果需要使用 uvx 或 npx 运行的服务器,也请在项目外部启动 MCP 服务器,并在项目中配置 MCP 服务器。" + ), + }, + ) + @classmethod def from_file(cls, module_name: str, input_context: dict = None) -> "BaseContext": """Load configuration from a YAML file. 用于持久化配置""" diff --git a/src/agents/common/subagents/calc_agent.py b/src/agents/common/subagents/calc_agent.py index 7c9a11b5..6fac054a 100644 --- a/src/agents/common/subagents/calc_agent.py +++ b/src/agents/common/subagents/calc_agent.py @@ -8,7 +8,7 @@ from src.agents.common.tools import calculator calc_agent = create_agent( model=load_chat_model(config.default_model), tools=[calculator], - system_prompt="你可以使用计算器工具,处理各种数学计算任务。", + system_prompt="你可以使用计算器工具,处理各种数学计算任务。最终仅返回计算结果,不需要任何额外的解释。", ) diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index ea7bd724..d4dc0791 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -1,13 +1,18 @@ import asyncio +import os import traceback +import uuid 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.services.mcp_service import get_enabled_mcp_tools +from src.storage.minio import aupload_file_to_minio from src.utils import logger # Lazy initialization for TavilySearch (only when TAVILY_API_KEY is available) @@ -45,6 +50,43 @@ def calculator(a: float, b: float, operation: str) -> float: raise +@tool +async def text_to_img_demo(text: str) -> str: + """【测试用】使用模型生成图片, 会返回图片的URL""" + + url = "https://api.siliconflow.cn/v1/images/generations" + + payload = { + "model": "Qwen/Qwen-Image", + "prompt": text, + } + headers = {"Authorization": f"Bearer {os.getenv('SILICONFLOW_API_KEY')}", "Content-Type": "application/json"} + + try: + response = requests.post(url, json=payload, headers=headers) + response_json = response.json() + except Exception as e: + logger.error(f"Failed to generate image with: {e}") + raise ValueError(f"Image generation failed: {e}") + + try: + image_url = response_json["images"][0]["url"] + except (KeyError, IndexError, TypeError) as e: + logger.error(f"Failed to parse image URL from response: {e}, {response_json=}") + raise ValueError(f"Image URL extraction failed: {e}") + + # 2. Upload to MinIO (Simplified) + response = requests.get(image_url) + file_data = response.content + + file_name = f"{uuid.uuid4()}.jpg" + image_url = await aupload_file_to_minio( + bucket_name="generated-images", file_name=file_name, data=file_data, file_extension="jpg" + ) + logger.info(f"Image uploaded. URL: {image_url}") + return image_url + + @tool(name_or_callable="人工审批工具(Debug)", description="请求人工审批工具,用于在执行重要操作前获得人类确认。") def get_approved_user_goal( operation_description: str, @@ -105,19 +147,6 @@ def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge return f"知识图谱查询失败: {str(e)}" -def get_static_tools() -> list: - """注册静态工具""" - static_tools = [query_knowledge_graph, get_approved_user_goal, calculator] - - # 检查是否启用网页搜索 - if config.enable_web_search: - tavily_search = get_tavily_search() - if tavily_search: - static_tools.append(tavily_search) - - return static_tools - - class KnowledgeRetrieverModel(BaseModel): query_text: str = Field( description=( @@ -257,23 +286,6 @@ def get_kb_based_tools(db_names: list[str] | None = None) -> list: return kb_tools -def get_buildin_tools() -> list: - """获取所有可运行的工具(给大模型使用)""" - tools = [] - - try: - tools.extend(get_static_tools()) - - from src.agents.common.toolkits.mysql.tools import get_mysql_tools - - tools.extend(get_mysql_tools()) - - except Exception as e: - logger.error(f"Failed to get knowledge base retrievers: {e}") - - return tools - - def gen_tool_info(tools) -> list[dict[str, Any]]: """获取所有工具的信息(用于前端展示)""" tools_info = [] @@ -323,3 +335,53 @@ def gen_tool_info(tools) -> list[dict[str, Any]]: logger.info(f"Successfully extracted info for {len(tools_info)} tools") return tools_info + + +def get_buildin_tools() -> list: + """注册静态工具""" + static_tools = [ + query_knowledge_graph, + get_approved_user_goal, + calculator, + text_to_img_demo, + ] + + # subagents 工具 + from .subagents import calc_agent_tool + + static_tools.append(calc_agent_tool) + + # 检查是否启用网页搜索(即是否配置了 API_KEY) + if config.enable_web_search: + tavily_search = get_tavily_search() + if tavily_search: + static_tools.append(tavily_search) + + return static_tools + + +async def get_tools_from_context(context, extra_tools=None) -> list: + """从上下文配置中获取工具列表""" + # 1. 基础工具 (从 context.tools 中筛选) + all_basic_tools = get_buildin_tools() + (extra_tools or []) + selected_tools = [] + + if context.tools: + # 创建工具映射表 + tools_map = {t.name: t for t in all_basic_tools} + for tool_name in context.tools: + if tool_name in tools_map: + selected_tools.append(tools_map[tool_name]) + + # 2. 知识库工具 + if context.knowledges: + kb_tools = get_kb_based_tools(db_names=context.knowledges) + selected_tools.extend(kb_tools) + + # 3. MCP 工具(使用统一入口,自动过滤 disabled_tools) + if context.mcps: + for server_name in context.mcps: + mcp_tools = await get_enabled_mcp_tools(server_name) + selected_tools.extend(mcp_tools) + + return selected_tools diff --git a/src/agents/mini_agent/graph.py b/src/agents/mini_agent/graph.py index 21d068d3..6271508b 100644 --- a/src/agents/mini_agent/graph.py +++ b/src/agents/mini_agent/graph.py @@ -1,7 +1,7 @@ from langchain.agents import create_agent -from src import config -from src.agents.common import BaseAgent, get_buildin_tools, load_chat_model +from src.agents.common import BaseAgent, load_chat_model +from src.agents.common.tools import get_tools_from_context class MiniAgent(BaseAgent): @@ -11,9 +11,6 @@ class MiniAgent(BaseAgent): def __init__(self, **kwargs): super().__init__(**kwargs) - def get_tools(self): - return get_buildin_tools() - async def get_graph(self, **kwargs): if self.graph: return self.graph @@ -22,9 +19,9 @@ class MiniAgent(BaseAgent): # 创建 MiniAgent graph = create_agent( - model=load_chat_model(config.default_model), + model=load_chat_model(context.model), system_prompt=context.system_prompt, - tools=self.get_tools(), + tools=await get_tools_from_context(context), checkpointer=await self._get_checkpointer(), ) diff --git a/src/agents/reporter/graph.py b/src/agents/reporter/graph.py index e688d378..d104b681 100644 --- a/src/agents/reporter/graph.py +++ b/src/agents/reporter/graph.py @@ -1,23 +1,41 @@ +from dataclasses import dataclass, field +from typing import Annotated + from langchain.agents import create_agent -from src.agents.common import BaseAgent, load_chat_model +from src.agents.common import BaseAgent, BaseContext, load_chat_model +from src.agents.common.tools import gen_tool_info, get_buildin_tools from src.agents.common.toolkits.mysql import get_mysql_tools -from src.services.mcp_service import get_mcp_tools +from src.agents.common.tools import get_tools_from_context from src.utils import logger +@dataclass(kw_only=True) +class ReporterContext(BaseContext): + + # 覆盖默认的工具列表,添加 MySQL 工具包 + tools: Annotated[list[dict], {"__template_metadata__": {"kind": "tools"}}] = field( + default_factory=lambda: [t.name for t in get_mysql_tools()], + metadata={ + "name": "工具", + # 添加额外的 MySQL 工具包选项 + "options": lambda: gen_tool_info(get_buildin_tools() + get_mysql_tools()), + "description": "包含内置的工具,以及用于数据库报表生成的 MySQL 工具包。", + }, + ) + + def __post_init__(self): + self.mcps = ["mcp-server-chart"] # 默认启用 Charts MCPs + + class SqlReporterAgent(BaseAgent): name = "数据库报表助手" description = "一个能够生成 SQL 查询报告的智能体助手。同时调用 Charts MCP 生成图表。" + context_schema = ReporterContext def __init__(self, **kwargs): super().__init__(**kwargs) - async def get_tools(self): - mysql_tools = get_mysql_tools() - chart_tools = await get_mcp_tools("mcp-server-chart") - return mysql_tools + chart_tools - async def get_graph(self, **kwargs): if self.graph: return self.graph @@ -28,7 +46,7 @@ class SqlReporterAgent(BaseAgent): graph = create_agent( model=load_chat_model(context.model), # 使用 context 中的模型配置 system_prompt=context.system_prompt, - tools=await self.get_tools(), + tools=await get_tools_from_context(context, extra_tools=get_mysql_tools()), checkpointer=await self._get_checkpointer(), ) diff --git a/src/services/evaluation_service.py b/src/services/evaluation_service.py index b74cfddb..604099a0 100644 --- a/src/services/evaluation_service.py +++ b/src/services/evaluation_service.py @@ -769,7 +769,7 @@ class EvaluationService: self, db_id: str, task_id: str, page: int = 1, page_size: int = 20, error_only: bool = False ) -> dict[str, Any]: # Validate task_id format to prevent path traversal - if not re.match(r'^eval_[a-f0-9]{8}$', task_id): + if not re.match(r"^eval_[a-f0-9]{8}$", task_id): raise ValueError("Invalid task_id format") result_file_path = os.path.join(self._get_result_dir(db_id), f"{task_id}.json") if not os.path.exists(result_file_path): @@ -838,7 +838,7 @@ class EvaluationService: async def delete_evaluation_result_by_db(self, db_id: str, task_id: str) -> None: # Validate task_id format to prevent path traversal - if not re.match(r'^eval_[a-f0-9]{8}$', task_id): + if not re.match(r"^eval_[a-f0-9]{8}$", task_id): raise ValueError("Invalid task_id format") result_file_path = os.path.join(self._get_result_dir(db_id), f"{task_id}.json") if os.path.exists(result_file_path):