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
### 工具系统
系统提供统一的工具获取函数 `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)`
@ -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, # 其他中间件...

View File

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

View File

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

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 (
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, # 附件上下文注入

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
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. 用于持久化配置"""

View File

@ -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="你可以使用计算器工具,处理各种数学计算任务。最终仅返回计算结果,不需要任何额外的解释。",
)

View File

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

View File

@ -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(),
)

View File

@ -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(),
)

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
) -> 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):