feat: 添加 SqlReporterAgent,支持生成 SQL 查询报告并调用图表工具
This commit is contained in:
parent
a8491c5beb
commit
2936f84ba1
@ -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)
|
||||
|
||||
@ -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()
|
||||
|
||||
|
||||
@ -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 = []
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
4
src/agents/reporter/__init__.py
Normal file
4
src/agents/reporter/__init__.py
Normal file
@ -0,0 +1,4 @@
|
||||
from .graph import SqlReporterAgent
|
||||
|
||||
|
||||
__all__ = ["SqlReporterAgent"]
|
||||
67
src/agents/reporter/graph.py
Normal file
67
src/agents/reporter/graph.py
Normal 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
|
||||
Loading…
Reference in New Issue
Block a user