refactor: 重构工具获取逻辑,移除过时的工具接口,优化上下文工具管理
This commit is contained in:
parent
9d1e2b92f5
commit
339cd0f3b2
@ -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, # 其他中间件...
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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 服务器。"
|
||||
),
|
||||
},
|
||||
)
|
||||
@ -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, # 附件上下文注入
|
||||
|
||||
@ -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
|
||||
@ -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. 用于持久化配置"""
|
||||
|
||||
@ -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="你可以使用计算器工具,处理各种数学计算任务。最终仅返回计算结果,不需要任何额外的解释。",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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(),
|
||||
)
|
||||
|
||||
|
||||
@ -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(),
|
||||
)
|
||||
|
||||
|
||||
@ -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):
|
||||
|
||||
Loading…
Reference in New Issue
Block a user