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