refactor: 重构工具获取逻辑,移除过时的工具接口,优化上下文工具管理

This commit is contained in:
Wenjie Zhang 2026-01-15 16:04:19 +08:00
parent 9d1e2b92f5
commit 339cd0f3b2
12 changed files with 225 additions and 202 deletions

View File

@ -36,7 +36,67 @@
<<< @/../src/agents/reporter/graph.py <<< @/../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(<id>, reload=True)` 智能体实例的生命周期交给管理器处理,会在自动发现时完成初始化并缓存单例,以便快速响应请求。在容器内热重载时,只要保存文件即可触发重新导入;需要强制刷新可调用 `agent_manager.get_agent(<id>, reload=True)`
@ -81,7 +141,7 @@ from src.agents.common.middlewares import inject_attachment_context
async def get_graph(self): async def get_graph(self):
graph = create_agent( graph = create_agent(
model=load_chat_model("..."), model=load_chat_model("..."),
tools=get_tools(), tools=tools,
middleware=[ middleware=[
inject_attachment_context, # 添加附件中间件 inject_attachment_context, # 添加附件中间件
context_aware_prompt, # 其他中间件... context_aware_prompt, # 其他中间件...

View File

@ -19,7 +19,6 @@ from server.utils.auth_middleware import get_db, get_required_user
from src import executor from src import executor
from src import config as conf from src import config as conf
from src.agents import agent_manager 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.models import select_model
from src.plugins.guard import content_guard from src.plugins.guard import content_guard
from src.services.doc_converter import ( 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} 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") @chat.post("/agent/{agent_id}/resume")
async def resume_agent_chat( async def resume_agent_chat(
agent_id: str, agent_id: str,

View File

@ -54,7 +54,7 @@ class AgentManager(metaclass=SingletonMeta):
# 遍历所有子目录 # 遍历所有子目录
for item in agents_dir.iterdir(): for item in agents_dir.iterdir():
logger.info(f"尝试导入模块:{item}") # logger.info(f"尝试导入模块:{item}")
# 跳过非目录、common 目录、__pycache__ 等 # 跳过非目录、common 目录、__pycache__ 等
if not item.is_dir() or item.name.startswith("_") or item.name == "common": if not item.is_dir() or item.name.startswith("_") or item.name == "common":
continue continue

View File

@ -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 服务器。"
),
},
)

View File

@ -5,46 +5,16 @@ from src.agents.common import BaseAgent, load_chat_model
from src.agents.common.middlewares import ( from src.agents.common.middlewares import (
inject_attachment_context, inject_attachment_context,
) )
from src.agents.common.tools import get_kb_based_tools from src.agents.common.tools import get_tools_from_context
from src.services.mcp_service import get_enabled_mcp_tools
from .context import Context
from .tools import get_tools
class ChatbotAgent(BaseAgent): class ChatbotAgent(BaseAgent):
name = "智能体助手" name = "智能体助手"
description = "基础的对话机器人,可以回答问题,默认不使用任何工具,可在配置中启用需要的工具。" description = "基础的对话机器人,可以回答问题,可在配置中启用需要的工具。"
capabilities = ["file_upload"] # 支持文件上传功能 capabilities = ["file_upload"] # 支持文件上传功能
def __init__(self, **kwargs): def __init__(self, **kwargs):
super().__init__(**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): async def get_graph(self, **kwargs):
"""构建图""" """构建图"""
@ -57,7 +27,7 @@ class ChatbotAgent(BaseAgent):
# 使用 create_agent 创建智能体 # 使用 create_agent 创建智能体
graph = create_agent( graph = create_agent(
model=load_chat_model(context.model), # 使用 context 中的模型配置 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, system_prompt=context.system_prompt,
middleware=[ middleware=[
inject_attachment_context, # 附件上下文注入 inject_attachment_context, # 附件上下文注入

View File

@ -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

View File

@ -9,8 +9,12 @@ from typing import Annotated, get_args, get_origin
import yaml import yaml
from src import config as sys_config 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 src.utils import logger
from .tools import gen_tool_info, get_buildin_tools
@dataclass(kw_only=True) @dataclass(kw_only=True)
class BaseContext: 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 @classmethod
def from_file(cls, module_name: str, input_context: dict = None) -> "BaseContext": def from_file(cls, module_name: str, input_context: dict = None) -> "BaseContext":
"""Load configuration from a YAML file. 用于持久化配置""" """Load configuration from a YAML file. 用于持久化配置"""

View File

@ -8,7 +8,7 @@ from src.agents.common.tools import calculator
calc_agent = create_agent( calc_agent = create_agent(
model=load_chat_model(config.default_model), model=load_chat_model(config.default_model),
tools=[calculator], tools=[calculator],
system_prompt="你可以使用计算器工具,处理各种数学计算任务。", system_prompt="你可以使用计算器工具,处理各种数学计算任务。最终仅返回计算结果,不需要任何额外的解释。",
) )

View File

@ -1,13 +1,18 @@
import asyncio import asyncio
import os
import traceback import traceback
import uuid
from typing import Annotated, Any from typing import Annotated, Any
import requests
from langchain.tools import tool from langchain.tools import tool
from langchain_core.tools import StructuredTool from langchain_core.tools import StructuredTool
from langgraph.types import interrupt from langgraph.types import interrupt
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from src import config, graph_base, knowledge_base 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 from src.utils import logger
# Lazy initialization for TavilySearch (only when TAVILY_API_KEY is available) # 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 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="请求人工审批工具,用于在执行重要操作前获得人类确认。") @tool(name_or_callable="人工审批工具(Debug)", description="请求人工审批工具,用于在执行重要操作前获得人类确认。")
def get_approved_user_goal( def get_approved_user_goal(
operation_description: str, operation_description: str,
@ -105,19 +147,6 @@ def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge
return f"知识图谱查询失败: {str(e)}" 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): class KnowledgeRetrieverModel(BaseModel):
query_text: str = Field( query_text: str = Field(
description=( description=(
@ -257,23 +286,6 @@ def get_kb_based_tools(db_names: list[str] | None = None) -> list:
return kb_tools 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]]: def gen_tool_info(tools) -> list[dict[str, Any]]:
"""获取所有工具的信息(用于前端展示)""" """获取所有工具的信息(用于前端展示)"""
tools_info = [] 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") logger.info(f"Successfully extracted info for {len(tools_info)} tools")
return tools_info 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

View File

@ -1,7 +1,7 @@
from langchain.agents import create_agent from langchain.agents import create_agent
from src import config from src.agents.common import BaseAgent, load_chat_model
from src.agents.common import BaseAgent, get_buildin_tools, load_chat_model from src.agents.common.tools import get_tools_from_context
class MiniAgent(BaseAgent): class MiniAgent(BaseAgent):
@ -11,9 +11,6 @@ class MiniAgent(BaseAgent):
def __init__(self, **kwargs): def __init__(self, **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
def get_tools(self):
return get_buildin_tools()
async def get_graph(self, **kwargs): async def get_graph(self, **kwargs):
if self.graph: if self.graph:
return self.graph return self.graph
@ -22,9 +19,9 @@ class MiniAgent(BaseAgent):
# 创建 MiniAgent # 创建 MiniAgent
graph = create_agent( graph = create_agent(
model=load_chat_model(config.default_model), model=load_chat_model(context.model),
system_prompt=context.system_prompt, system_prompt=context.system_prompt,
tools=self.get_tools(), tools=await get_tools_from_context(context),
checkpointer=await self._get_checkpointer(), checkpointer=await self._get_checkpointer(),
) )

View File

@ -1,23 +1,41 @@
from dataclasses import dataclass, field
from typing import Annotated
from langchain.agents import create_agent 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.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 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): class SqlReporterAgent(BaseAgent):
name = "数据库报表助手" name = "数据库报表助手"
description = "一个能够生成 SQL 查询报告的智能体助手。同时调用 Charts MCP 生成图表。" description = "一个能够生成 SQL 查询报告的智能体助手。同时调用 Charts MCP 生成图表。"
context_schema = ReporterContext
def __init__(self, **kwargs): def __init__(self, **kwargs):
super().__init__(**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): async def get_graph(self, **kwargs):
if self.graph: if self.graph:
return self.graph return self.graph
@ -28,7 +46,7 @@ class SqlReporterAgent(BaseAgent):
graph = create_agent( graph = create_agent(
model=load_chat_model(context.model), # 使用 context 中的模型配置 model=load_chat_model(context.model), # 使用 context 中的模型配置
system_prompt=context.system_prompt, 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(), checkpointer=await self._get_checkpointer(),
) )

View File

@ -769,7 +769,7 @@ class EvaluationService:
self, db_id: str, task_id: str, page: int = 1, page_size: int = 20, error_only: bool = False self, db_id: str, task_id: str, page: int = 1, page_size: int = 20, error_only: bool = False
) -> dict[str, Any]: ) -> dict[str, Any]:
# Validate task_id format to prevent path traversal # 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") raise ValueError("Invalid task_id format")
result_file_path = os.path.join(self._get_result_dir(db_id), f"{task_id}.json") result_file_path = os.path.join(self._get_result_dir(db_id), f"{task_id}.json")
if not os.path.exists(result_file_path): 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: async def delete_evaluation_result_by_db(self, db_id: str, task_id: str) -> None:
# Validate task_id format to prevent path traversal # 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") raise ValueError("Invalid task_id format")
result_file_path = os.path.join(self._get_result_dir(db_id), f"{task_id}.json") result_file_path = os.path.join(self._get_result_dir(db_id), f"{task_id}.json")
if os.path.exists(result_file_path): if os.path.exists(result_file_path):