feat: 添加 SqlReporterAgent,支持生成 SQL 查询报告并调用图表工具

This commit is contained in:
Wenjie Zhang 2025-10-25 22:49:20 +08:00
parent a8491c5beb
commit 2936f84ba1
8 changed files with 94 additions and 19 deletions

View File

@ -17,6 +17,8 @@
语析是一个功能强大的智能问答平台,融合了 RAG 知识库与知识图谱技术,基于 LangGraph + Vue.js + FastAPI + LightRAG 架构构建。
---
🙏 感谢 Star ~ ⭐⭐⭐⭐⭐
详细文档请查看全新的 [**📄文档中心**](https://xerrors.github.io/Yuxi-Know/) [📽️ 点击查看视频演示 v0.2](https://www.bilibili.com/video/BV1ETedzREgY/?share_source=copy_web&vd_source=37b0bdbf95b72ea38b2dc959cfadc4d8)

View File

@ -408,8 +408,11 @@ async def get_tools(agent_id: str, current_user: User = Depends(get_required_use
if not (agent := agent_manager.get_agent(agent_id)):
raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在")
if hasattr(agent, "get_tools"):
tools = agent.get_tools()
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()

View File

@ -58,7 +58,7 @@ async def get_mcp_client(
return None
async def get_mcp_tools(server_name: str) -> list[Callable[..., Any]]:
async def get_mcp_tools(server_name: str, additional_servers: dict[str, dict] = None) -> list[Callable[..., Any]]:
"""Get MCP tools for a specific server, initializing client if needed."""
global _mcp_tools_cache
@ -66,9 +66,11 @@ async def get_mcp_tools(server_name: str) -> list[Callable[..., Any]]:
if server_name in _mcp_tools_cache:
return _mcp_tools_cache[server_name]
mcp_servers = MCP_SERVERS | (additional_servers or {})
try:
assert server_name in MCP_SERVERS, f"Server {server_name} not found in MCP_SERVERS"
client = await get_mcp_client({server_name: MCP_SERVERS[server_name]})
assert server_name in mcp_servers, f"Server {server_name} not found in MCP_SERVERS"
client = await get_mcp_client({server_name: mcp_servers[server_name]})
if client is None:
return []
@ -86,7 +88,6 @@ async def get_mcp_tools(server_name: str) -> list[Callable[..., Any]]:
logger.error(f"Failed to load tools from MCP server '{server_name}': {e}")
return []
async def get_all_mcp_tools() -> list[Callable[..., Any]]:
"""Get all tools from all configured MCP servers."""
all_tools = []

View File

@ -44,7 +44,7 @@ class TableListModel(BaseModel):
pass
@tool(args_schema=TableListModel)
@tool(name_or_callable="查询表名", args_schema=TableListModel)
def mysql_list_tables() -> str:
"""获取数据库中的所有表名
@ -94,7 +94,7 @@ class TableDescribeModel(BaseModel):
table_name: str = Field(description="要查询的表名", example="users")
@tool(args_schema=TableDescribeModel)
@tool(name_or_callable="描述表", args_schema=TableDescribeModel)
def mysql_describe_table(table_name: Annotated[str, "要查询结构的表名"]) -> str:
"""获取指定表的详细结构信息
@ -168,7 +168,7 @@ class QueryModel(BaseModel):
timeout: int | None = Field(default=10, description="查询超时时间默认10秒最大60秒", ge=1, le=60)
@tool(args_schema=QueryModel)
@tool(name_or_callable="执行 SQL 查询", args_schema=QueryModel)
def mysql_query(
sql: Annotated[str, "要执行的SQL查询语句只能是SELECT语句"],
limit: Annotated[int | None, "限制返回的最大行数默认100最大1000"] = 100,

View File

@ -146,7 +146,11 @@ def gen_tool_info(tools) -> list[dict[str, Any]]:
}
if hasattr(tool_obj, "args_schema") and tool_obj.args_schema:
schema = tool_obj.args_schema.schema()
if isinstance(tool_obj.args_schema, dict):
schema = tool_obj.args_schema
else:
schema = tool_obj.args_schema.schema()
for arg_name, arg_info in schema.get("properties", {}).items():
info["args"].append(
{
@ -161,7 +165,8 @@ def gen_tool_info(tools) -> list[dict[str, Any]]:
except Exception as e:
logger.error(
f"Failed to process tool {getattr(tool_obj, 'name', 'unknown')}: {e}\n{traceback.format_exc()}"
f"Failed to process tool {getattr(tool_obj, 'name', 'unknown')}: {e}\n{traceback.format_exc()}. "
f"Details: {dict(tool_obj.__dict__)}"
)
continue

View File

@ -3,14 +3,11 @@ from pathlib import Path
from langchain.agents import create_agent
from langchain.agents.middleware import ModelRequest, ModelResponse, dynamic_prompt, wrap_model_call
from src import config as sys_config
from src.agents.common.base import BaseAgent
from src.agents.common.models import load_chat_model
from src.agents.common.tools import get_buildin_tools
from src.utils import logger
model = load_chat_model("siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507")
@dynamic_prompt
def context_aware_prompt(request: ModelRequest) -> str:
@ -34,9 +31,6 @@ class ReActAgent(BaseAgent):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.graph = None
self.workdir = Path(sys_config.save_dir) / "agents" / self.module_name
self.workdir.mkdir(parents=True, exist_ok=True)
def get_tools(self):
return get_buildin_tools()
@ -47,12 +41,11 @@ class ReActAgent(BaseAgent):
# 创建 ReActAgent
graph = create_agent(
model=model,
model=load_chat_model("siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507"), # 实际会被覆盖
tools=self.get_tools(),
middleware=[context_aware_prompt, context_based_model],
checkpointer=await self._get_checkpointer(),
)
self.graph = graph
logger.info("ReActAgent 使用内存 checkpointer 构建成功")
return graph

View File

@ -0,0 +1,4 @@
from .graph import SqlReporterAgent
__all__ = ["SqlReporterAgent"]

View File

@ -0,0 +1,67 @@
import textwrap
from pathlib import Path
from langchain.agents import create_agent
from langchain.agents.middleware import ModelRequest, ModelResponse, dynamic_prompt, wrap_model_call
from src.agents.common.base import BaseAgent
from src.agents.common.models import load_chat_model
from src.agents.common.mcp import get_mcp_tools
from src.agents.common.toolkits.mysql import get_mysql_tools
from src.utils import logger
_mcp_servers = {
"mcp-server-chart": {
"url": "https://mcp.api-inference.modelscope.net/9993ae42524c4c/mcp",
"transport": "streamable_http",
},
}
@dynamic_prompt
def context_aware_prompt(request: ModelRequest) -> str:
user_prompt = request.runtime.context.system_prompt
agent_prompt = user_prompt + textwrap.dedent("""
You are an SQL reporting assistant. Your task is to generate SQL queries based on user requests
and provide insights from the database. Use the tools provided to you to answer the questions.
""")
return agent_prompt
@wrap_model_call
async def context_based_model(request: ModelRequest, handler) -> ModelResponse:
# 从 runtime context 读取配置
model_spec = request.runtime.context.model
model = load_chat_model(model_spec)
request = request.override(model=model)
return await handler(request)
class SqlReporterAgent(BaseAgent):
name = "SQL 报告助手"
description = "一个能够生成 SQL 查询报告的智能体助手。同时调用 Charts MCP 生成图表。"
def __init__(self, **kwargs):
super().__init__(**kwargs)
async def get_tools(self):
chart_tools = await get_mcp_tools("mcp-server-chart", additional_servers=_mcp_servers)
mysql_tools = get_mysql_tools()
return chart_tools + mysql_tools
async def get_graph(self, **kwargs):
if self.graph:
return self.graph
# 创建 SqlReporterAgent
graph = create_agent(
model=load_chat_model("siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507"),
tools=await self.get_tools(),
middleware=[context_aware_prompt, context_based_model],
checkpointer=await self._get_checkpointer(),
)
self.graph = graph
logger.info("SqlReporterAgent 构建成功")
return graph