feat: 添加 middleware, subagents 支持

- 统一导入路径,优化代码结构,删除冗余文件和函数
- refactor: 将 chatbot agent 使用 create_agent 重构,大幅简化处理逻辑
This commit is contained in:
Wenjie Zhang 2025-11-05 02:04:34 +08:00
parent 159f4819f0
commit 09fffb233c
15 changed files with 219 additions and 188 deletions

View File

@ -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

View File

@ -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

View File

@ -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

View File

@ -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)

View File

@ -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: 计算操作符号可以是addsubtractmultiplydivide
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

View File

@ -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",
]

View File

@ -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",
]

View File

@ -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)

View File

@ -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)

View File

@ -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

View File

@ -0,0 +1,3 @@
from .calc_agent import calc_agent, calc_agent_tool
__all__ = ["calc_agent", "calc_agent_tool"]

View File

@ -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

View File

@ -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:

View File

@ -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):

View File

@ -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 生成图表。"