diff --git a/README.md b/README.md index a9746a0a..27a00aa6 100644 --- a/README.md +++ b/README.md @@ -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) diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index 9ec98af4..f0d93b33 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -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() diff --git a/src/agents/common/mcp.py b/src/agents/common/mcp.py index 60476461..1fa36171 100644 --- a/src/agents/common/mcp.py +++ b/src/agents/common/mcp.py @@ -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 = [] diff --git a/src/agents/common/toolkits/mysql/tools.py b/src/agents/common/toolkits/mysql/tools.py index 8e86c27e..d4bb401a 100644 --- a/src/agents/common/toolkits/mysql/tools.py +++ b/src/agents/common/toolkits/mysql/tools.py @@ -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, diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index 0023d9b8..7d9fa2aa 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -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 diff --git a/src/agents/react/graph.py b/src/agents/react/graph.py index 831905fb..eb79c87a 100644 --- a/src/agents/react/graph.py +++ b/src/agents/react/graph.py @@ -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 diff --git a/src/agents/reporter/__init__.py b/src/agents/reporter/__init__.py new file mode 100644 index 00000000..03af7748 --- /dev/null +++ b/src/agents/reporter/__init__.py @@ -0,0 +1,4 @@ +from .graph import SqlReporterAgent + + +__all__ = ["SqlReporterAgent"] diff --git a/src/agents/reporter/graph.py b/src/agents/reporter/graph.py new file mode 100644 index 00000000..8402722a --- /dev/null +++ b/src/agents/reporter/graph.py @@ -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 \ No newline at end of file