feat: 添加 middleware, subagents 支持
- 统一导入路径,优化代码结构,删除冗余文件和函数 - refactor: 将 chatbot agent 使用 create_agent 重构,大幅简化处理逻辑
This commit is contained in:
parent
159f4819f0
commit
09fffb233c
@ -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
|
||||
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
@ -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)
|
||||
@ -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
|
||||
|
||||
39
src/agents/common/__init__.py
Normal file
39
src/agents/common/__init__.py
Normal 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",
|
||||
]
|
||||
8
src/agents/common/middlewares/__init__.py
Normal file
8
src/agents/common/middlewares/__init__.py
Normal 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",
|
||||
]
|
||||
25
src/agents/common/middlewares/context_middlewares.py
Normal file
25
src/agents/common/middlewares/context_middlewares.py
Normal 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)
|
||||
68
src/agents/common/middlewares/dynamic_tool_middleware.py
Normal file
68
src/agents/common/middlewares/dynamic_tool_middleware.py
Normal 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)
|
||||
@ -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
|
||||
|
||||
3
src/agents/common/subagents/__init__.py
Normal file
3
src/agents/common/subagents/__init__.py
Normal file
@ -0,0 +1,3 @@
|
||||
from .calc_agent import calc_agent, calc_agent_tool
|
||||
|
||||
__all__ = ["calc_agent", "calc_agent_tool"]
|
||||
19
src/agents/common/subagents/calc_agent.py
Normal file
19
src/agents/common/subagents/calc_agent.py
Normal 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
|
||||
@ -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:
|
||||
|
||||
@ -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):
|
||||
|
||||
@ -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 生成图表。"
|
||||
|
||||
Loading…
Reference in New Issue
Block a user