From 09fffb233cb5d9561b580750ba0b5853bdb96c98 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Wed, 5 Nov 2025 02:04:34 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0=20middleware,=20suba?= =?UTF-8?q?gents=20=E6=94=AF=E6=8C=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 统一导入路径,优化代码结构,删除冗余文件和函数 - refactor: 将 chatbot agent 使用 create_agent 重构,大幅简化处理逻辑 --- src/agents/__init__.py | 2 +- src/agents/chatbot/context.py | 3 +- src/agents/chatbot/graph.py | 103 +++++------------- src/agents/chatbot/state.py | 20 ---- src/agents/chatbot/tools.py | 37 +------ src/agents/common/__init__.py | 39 +++++++ src/agents/common/middlewares/__init__.py | 8 ++ .../common/middlewares/context_middlewares.py | 25 +++++ .../middlewares/dynamic_tool_middleware.py | 68 ++++++++++++ src/agents/common/models.py | 10 -- src/agents/common/subagents/__init__.py | 3 + src/agents/common/subagents/calc_agent.py | 19 ++++ src/agents/common/tools.py | 23 +++- src/agents/mini_agent/graph.py | 21 +--- src/agents/reporter/graph.py | 26 +---- 15 files changed, 219 insertions(+), 188 deletions(-) delete mode 100644 src/agents/chatbot/state.py create mode 100644 src/agents/common/__init__.py create mode 100644 src/agents/common/middlewares/__init__.py create mode 100644 src/agents/common/middlewares/context_middlewares.py create mode 100644 src/agents/common/middlewares/dynamic_tool_middleware.py create mode 100644 src/agents/common/subagents/__init__.py create mode 100644 src/agents/common/subagents/calc_agent.py diff --git a/src/agents/__init__.py b/src/agents/__init__.py index 2b2746fb..210ed251 100644 --- a/src/agents/__init__.py +++ b/src/agents/__init__.py @@ -4,7 +4,7 @@ import inspect from pathlib import Path from server.utils.singleton import SingletonMeta -from src.agents.common.base import BaseAgent +from src.agents.common import BaseAgent from src.utils import logger diff --git a/src/agents/chatbot/context.py b/src/agents/chatbot/context.py index 7bb1adad..4e12e44a 100644 --- a/src/agents/chatbot/context.py +++ b/src/agents/chatbot/context.py @@ -1,9 +1,8 @@ from dataclasses import dataclass, field from typing import Annotated -from src.agents.common.context import BaseContext +from src.agents.common import BaseContext, gen_tool_info from src.agents.common.mcp import MCP_SERVERS -from src.agents.common.tools import gen_tool_info from .tools import get_tools diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index 12b263aa..28ffa0ef 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -1,17 +1,11 @@ -from typing import Any, cast +from langchain.agents import create_agent -from langchain.messages import AIMessage, ToolMessage -from langgraph.graph import END, START, StateGraph -from langgraph.prebuilt import ToolNode, tools_condition -from langgraph.runtime import Runtime - -from src.agents.common.base import BaseAgent -from src.agents.common.mcp import get_mcp_tools -from src.agents.common.models import load_chat_model -from src.utils import logger +from src.agents.common import BaseAgent, load_chat_model +from src.agents.common.mcp import MCP_SERVERS +from src.agents.common.middlewares import DynamicToolMiddleware, context_aware_prompt, context_based_model +from src.agents.common.subagents import calc_agent_tool from .context import Context -from .state import State from .tools import get_tools @@ -24,81 +18,38 @@ class ChatbotAgent(BaseAgent): self.graph = None self.checkpointer = None self.context_schema = Context - self.agent_tools = None def get_tools(self): - return get_tools() - - async def _get_invoke_tools(self, selected_tools: list[str], selected_mcps: list[str]): - """根据配置获取工具。 - 默认不使用任何工具。 - 如果配置为列表,则使用列表中的工具。 - """ - enabled_tools = [] - self.agent_tools = self.agent_tools or self.get_tools() - if selected_tools and isinstance(selected_tools, list) and len(selected_tools) > 0: - # 使用配置中指定的工具 - enabled_tools = [tool for tool in self.agent_tools if tool.name in selected_tools] - - if selected_mcps and isinstance(selected_mcps, list) and len(selected_mcps) > 0: - for mcp in selected_mcps: - enabled_tools.extend(await get_mcp_tools(mcp)) - - return enabled_tools - - async def llm_call(self, state: State, runtime: Runtime[Context] = None) -> dict[str, Any]: - """调用 llm 模型 - 异步版本以支持异步工具""" - model = load_chat_model(runtime.context.model) - - # 这里要根据配置动态获取工具 - available_tools = await self._get_invoke_tools(runtime.context.tools, runtime.context.mcps) - logger.info(f"LLM binded ({len(available_tools)}) available_tools: {[tool.name for tool in available_tools]}") - - if available_tools: - model = model.bind_tools(available_tools) - - # 使用异步调用 - response = cast( - AIMessage, - await model.ainvoke([{"role": "system", "content": runtime.context.system_prompt}, *state.messages]), - ) - return {"messages": [response]} - - async def dynamic_tools_node(self, state: State, runtime: Runtime[Context]) -> dict[str, list[ToolMessage]]: - """Execute tools dynamically based on configuration. - - This function gets the available tools based on the current configuration - and executes the requested tool calls from the last message. - """ - # Get available tools based on configuration - available_tools = await self._get_invoke_tools(runtime.context.tools, runtime.context.mcps) - - # Create a ToolNode with the available tools - tool_node = ToolNode(available_tools) - - # Execute the tool node - result = await tool_node.ainvoke(state) - - return cast(dict[str, list[ToolMessage]], result) + """返回基本工具""" + base_tools = get_tools() + base_tools.append(calc_agent_tool) + return base_tools async def get_graph(self, **kwargs): """构建图""" if self.graph: return self.graph - builder = StateGraph(State, context_schema=self.context_schema) - builder.add_node("chatbot", self.llm_call) - builder.add_node("tools", self.dynamic_tools_node) - builder.add_edge(START, "chatbot") - builder.add_conditional_edges( - "chatbot", - tools_condition, + # 创建动态工具中间件实例,并传入所有可用的 MCP 服务器列表 + dynamic_tool_middleware = DynamicToolMiddleware( + base_tools=self.get_tools(), mcp_servers=list(MCP_SERVERS.keys()) + ) + + # 预加载所有 MCP 工具并注册到 middleware.tools + await dynamic_tool_middleware.initialize_mcp_tools() + + # 使用 create_agent 创建智能体,并传入 middleware + graph = create_agent( + model=load_chat_model("siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507"), # 默认模型,会被 middleware 覆盖 + tools=get_tools(), # 注册基础工具 + middleware=[ + context_aware_prompt, # 动态系统提示词 + context_based_model, # 动态模型选择 + dynamic_tool_middleware, # 动态工具选择(支持 MCP 工具注册) + ], + checkpointer=await self._get_checkpointer(), ) - builder.add_edge("tools", "chatbot") - builder.add_edge("chatbot", END) - self.checkpointer = await self._get_checkpointer() - graph = builder.compile(checkpointer=self.checkpointer, name=self.name) self.graph = graph return graph diff --git a/src/agents/chatbot/state.py b/src/agents/chatbot/state.py deleted file mode 100644 index f46bfb9c..00000000 --- a/src/agents/chatbot/state.py +++ /dev/null @@ -1,20 +0,0 @@ -"""Define the state structures for the agent.""" - -from __future__ import annotations - -from collections.abc import Sequence -from dataclasses import dataclass, field -from typing import Annotated - -from langchain.messages import AnyMessage -from langgraph.graph import add_messages - - -@dataclass -class State: - """Defines the input state for the agent, representing a narrower interface to the outside world. - - This class is used to define the initial state and structure of incoming data. - """ - - messages: Annotated[Sequence[AnyMessage], add_messages] = field(default_factory=list) diff --git a/src/agents/chatbot/tools.py b/src/agents/chatbot/tools.py index edfc1b10..4ae144a0 100644 --- a/src/agents/chatbot/tools.py +++ b/src/agents/chatbot/tools.py @@ -4,47 +4,15 @@ from typing import Any import requests from langchain.tools import tool +from src.agents.common import get_buildin_tools from src.agents.common.toolkits.mysql import get_mysql_tools -from src.agents.common.tools import get_buildin_tools from src.storage.minio import upload_image_to_minio from src.utils import logger -# TODO:[已完成]修改了tool定义的示例,使用更符合langgraph调用的方式 -@tool(name_or_callable="全能计算器", description="可以对给定的2个数字选择进行加减乘除四种计算") -def calculator(a: float, b: float, operation: str) -> float: - """ - 可以对给定的2个数字选择进行加减乘除四种计算 - - Args: - a: 第一个数字 - b: 第二个数字 - operation: 计算操作符号,可以是add,subtract,multiply,divide - - Returns: - float: 最终的计算结果 - """ - try: - if operation == "add": - return a + b - elif operation == "subtract": - return a - b - elif operation == "multiply": - return a * b - elif operation == "divide": - if b == 0: - raise ZeroDivisionError("除数不能为零") - return a / b - else: - raise ValueError(f"不支持的运算类型: {operation},仅支持 add, subtract, multiply, divide") - except Exception as e: - logger.error(f"Calculator error: {e}") - raise - - @tool async def text_to_img_qwen(text: str) -> str: - """(用来测试文件存储)使用Kolors模型生成图片, 会返回图片的URL""" + """(用来测试文件存储)使用模型生成图片, 会返回图片的URL""" url = "https://api.siliconflow.cn/v1/images/generations" @@ -79,7 +47,6 @@ async def text_to_img_qwen(text: str) -> str: def get_tools() -> list[Any]: """获取所有可运行的工具(给大模型使用)""" tools = get_buildin_tools() - tools.append(calculator) tools.append(text_to_img_qwen) tools.extend(get_mysql_tools()) return tools diff --git a/src/agents/common/__init__.py b/src/agents/common/__init__.py new file mode 100644 index 00000000..8c155072 --- /dev/null +++ b/src/agents/common/__init__.py @@ -0,0 +1,39 @@ +""" +Common utilities and base classes for agents. + +This module provides a unified namespace for commonly used base classes and utilities, +allowing simplified imports like: + from src.agents.common import BaseAgent, BaseContext, BaseState + +For other specific functions, use the original import style: + from src.agents.common.tools import query_knowledge_graph + from src.agents.common.mcp import MCP_SERVERS +""" + +# Base classes - 核心基类 +from src.agents.common.base import BaseAgent +from src.agents.common.context import BaseContext +from src.agents.common.state import BaseState + +# Model utilities - 模型加载 +from src.agents.common.models import load_chat_model + +# Tools - 核心工具函数 +from src.agents.common.tools import gen_tool_info, get_buildin_tools + +# MCP - 核心 MCP 函数 +from src.agents.common.mcp import get_mcp_tools + +__all__ = [ + # Base classes + "BaseAgent", + "BaseContext", + "BaseState", + # Model utilities + "load_chat_model", + # Core tools + "get_buildin_tools", + "gen_tool_info", + # Core MCP + "get_mcp_tools", +] diff --git a/src/agents/common/middlewares/__init__.py b/src/agents/common/middlewares/__init__.py new file mode 100644 index 00000000..fc571db1 --- /dev/null +++ b/src/agents/common/middlewares/__init__.py @@ -0,0 +1,8 @@ +from .context_middlewares import context_aware_prompt, context_based_model +from .dynamic_tool_middleware import DynamicToolMiddleware + +__all__ = [ + "DynamicToolMiddleware", + "context_aware_prompt", + "context_based_model", +] diff --git a/src/agents/common/middlewares/context_middlewares.py b/src/agents/common/middlewares/context_middlewares.py new file mode 100644 index 00000000..056ee996 --- /dev/null +++ b/src/agents/common/middlewares/context_middlewares.py @@ -0,0 +1,25 @@ +"""通用的 Context 相关中间件""" + +from collections.abc import Callable + +from langchain.agents.middleware import ModelRequest, ModelResponse, dynamic_prompt, wrap_model_call + +from src.agents.common import load_chat_model + + +@dynamic_prompt +def context_aware_prompt(request: ModelRequest) -> str: + """从 runtime context 动态生成系统提示词""" + return request.runtime.context.system_prompt + + +@wrap_model_call +async def context_based_model( + request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] +) -> ModelResponse: + """从 runtime context 动态选择模型""" + model_spec = request.runtime.context.model + model = load_chat_model(model_spec) + + request = request.override(model=model) + return await handler(request) diff --git a/src/agents/common/middlewares/dynamic_tool_middleware.py b/src/agents/common/middlewares/dynamic_tool_middleware.py new file mode 100644 index 00000000..acce5f4d --- /dev/null +++ b/src/agents/common/middlewares/dynamic_tool_middleware.py @@ -0,0 +1,68 @@ +from collections.abc import Callable +from typing import Any + +from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse + +from src.agents.common import get_mcp_tools +from src.utils import logger + + +class DynamicToolMiddleware(AgentMiddleware): + """动态工具选择中间件 - 支持 MCP 工具的动态加载和注册 + + 注意:所有可能用到的 MCP 工具必须在初始化时预加载并注册到 self.tools + 运行时只是根据配置筛选工具,不能动态添加新工具 + """ + + def __init__(self, base_tools: list[Any], mcp_servers: list[str] | None = None): + """初始化中间件 + + Args: + base_tools: 基础工具列表 + mcp_servers: 需要预加载的 MCP 服务器列表(可选) + """ + super().__init__() + self.tools: list[Any] = base_tools + self._all_mcp_tools: dict[str, list[Any]] = {} # 所有已加载的 MCP 工具 + self._mcp_servers = mcp_servers or [] + + async def initialize_mcp_tools(self) -> None: + """异步初始化:预加载所有可能用到的 MCP 工具""" + for mcp_name in self._mcp_servers: + if mcp_name not in self._all_mcp_tools: + logger.info(f"Pre-loading MCP tools from: {mcp_name}") + mcp_tools = await get_mcp_tools(mcp_name) + self._all_mcp_tools[mcp_name] = mcp_tools + # 将 MCP 工具注册到 middleware.tools + self.tools.extend(mcp_tools) + logger.info(f"Registered {len(mcp_tools)} tools from {mcp_name}") + + async def awrap_model_call( + self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] + ) -> ModelResponse: + """根据配置动态选择工具(从已注册的工具中筛选)""" + # 从 runtime context 获取配置 + selected_tools = request.runtime.context.tools + selected_mcps = request.runtime.context.mcps + + enabled_tools = [] + + # 根据配置筛选基础工具 + if selected_tools and isinstance(selected_tools, list) and len(selected_tools) > 0: + enabled_tools = [tool for tool in self.tools if tool.name in selected_tools] + + # 根据配置筛选 MCP 工具(从已注册的工具中选择) + if selected_mcps and isinstance(selected_mcps, list) and len(selected_mcps) > 0: + for mcp in selected_mcps: + if mcp in self._all_mcp_tools: + enabled_tools.extend(self._all_mcp_tools[mcp]) + else: + logger.warning(f"MCP server '{mcp}' not pre-loaded. Please add it to mcp_servers list.") + + logger.info( + f"Dynamic tool selection: {len(enabled_tools)} tools enabled: {[tool.name for tool in enabled_tools]}" + ) + + # 更新 request 中的工具列表 + request = request.override(tools=enabled_tools) + return await handler(request) diff --git a/src/agents/common/models.py b/src/agents/common/models.py index 2a6bebc4..2aaeaf0f 100644 --- a/src/agents/common/models.py +++ b/src/agents/common/models.py @@ -37,16 +37,6 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: stream_usage=True, ) - # elif provider == "together": - # from langchain_together import ChatTogether - - # return ChatTogether( - # model=model, - # api_key=SecretStr(api_key), - # base_url=base_url, - # stream_usage=True, - # ) - else: try: # 其他模型,默认使用OpenAIBase, like openai, zhipuai from langchain_openai import ChatOpenAI diff --git a/src/agents/common/subagents/__init__.py b/src/agents/common/subagents/__init__.py new file mode 100644 index 00000000..ebbf765c --- /dev/null +++ b/src/agents/common/subagents/__init__.py @@ -0,0 +1,3 @@ +from .calc_agent import calc_agent, calc_agent_tool + +__all__ = ["calc_agent", "calc_agent_tool"] \ No newline at end of file diff --git a/src/agents/common/subagents/calc_agent.py b/src/agents/common/subagents/calc_agent.py new file mode 100644 index 00000000..918bf06f --- /dev/null +++ b/src/agents/common/subagents/calc_agent.py @@ -0,0 +1,19 @@ +from langchain.agents import create_agent +from langchain.tools import tool + +from src import config +from src.agents.common import load_chat_model +from src.agents.common.tools import calculator + + +calc_agent = create_agent( + model=load_chat_model(config.default_model), + tools=[calculator], + system_prompt="你可以使用计算器工具,处理各种数学计算任务。", +) + +@tool(name_or_callable="calc_agent_tool", description="使用 CalcAgent 进行计算任务,输入是数学表达式或计算描述,输出是计算结果。") +async def calc_agent_tool(description: str) -> str: + """CalcAgent 工具 - 使用子智能体 CalcAgent 进行计算任务""" + response = await calc_agent.ainvoke({"messages": [("user", description)]}) + return response["messages"][-1].content \ No newline at end of file diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index f240c507..27210ac6 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -12,7 +12,28 @@ from src import config, graph_base, knowledge_base from src.utils import logger -@tool(name_or_callable="人工审批工具", description="请求人工审批工具,用于在执行重要操作前获得人类确认。") + +@tool(name_or_callable="计算器", description="可以对给定的2个数字选择进行 add, subtract, multiply, divide 运算") +def calculator(a: float, b: float, operation: str) -> float: + try: + if operation == "add": + return a + b + elif operation == "subtract": + return a - b + elif operation == "multiply": + return a * b + elif operation == "divide": + if b == 0: + raise ZeroDivisionError("除数不能为零") + return a / b + else: + raise ValueError(f"不支持的运算类型: {operation},仅支持 add, subtract, multiply, divide") + except Exception as e: + logger.error(f"Calculator error: {e}") + raise + + +@tool(name_or_callable="人工审批工具(Debug)", description="请求人工审批工具,用于在执行重要操作前获得人类确认。") def get_approved_user_goal( operation_description: str, )->dict: diff --git a/src/agents/mini_agent/graph.py b/src/agents/mini_agent/graph.py index e11f7751..ab78188d 100644 --- a/src/agents/mini_agent/graph.py +++ b/src/agents/mini_agent/graph.py @@ -1,25 +1,8 @@ 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.tools import get_buildin_tools - - -@dynamic_prompt -def context_aware_prompt(request: ModelRequest) -> str: - runtime = request.runtime - return runtime.context.system_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) +from src.agents.common import BaseAgent, load_chat_model, get_buildin_tools +from src.agents.common.middlewares import context_aware_prompt, context_based_model class MiniAgent(BaseAgent): diff --git a/src/agents/reporter/graph.py b/src/agents/reporter/graph.py index 29c12b69..74ce52d9 100644 --- a/src/agents/reporter/graph.py +++ b/src/agents/reporter/graph.py @@ -3,9 +3,8 @@ import textwrap 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.mcp import get_mcp_tools -from src.agents.common.models import load_chat_model +from src.agents.common import BaseAgent, load_chat_model, get_mcp_tools +from src.agents.common.middlewares import context_aware_prompt, context_based_model from src.agents.common.toolkits.mysql import get_mysql_tools from src.utils import logger @@ -17,27 +16,6 @@ _mcp_servers = { } -@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 = "数据库报表助手" description = "一个能够生成 SQL 查询报告的智能体助手。同时调用 Charts MCP 生成图表。"