refactor(agents): 实现统一工具加载逻辑并使用中间件动态配置

将各agent的工具加载逻辑统一改为通过RuntimeConfigMiddleware动态配置
移除直接的工具参数传递,改为从MCP服务获取工具并合并到中间件
在chat_stream_service中添加知识库权限过滤功能

TODO:目前历史记录的加载导致会话混乱
This commit is contained in:
Wenjie Zhang 2026-01-22 06:54:14 +08:00
parent 01ea590f76
commit 5f90385ad0
6 changed files with 103 additions and 48 deletions

View File

@ -6,7 +6,7 @@ from src.agents.common.middlewares import (
RuntimeConfigMiddleware, RuntimeConfigMiddleware,
inject_attachment_context, inject_attachment_context,
) )
from src.agents.common.tools import get_buildin_tools from src.services.mcp_service import get_tools_from_all_servers
class ChatbotAgent(BaseAgent): class ChatbotAgent(BaseAgent):
@ -20,15 +20,16 @@ class ChatbotAgent(BaseAgent):
async def get_graph(self, **kwargs): async def get_graph(self, **kwargs):
"""构建图""" """构建图"""
context = self.context_schema() context = self.context_schema()
all_mcp_tools = await get_tools_from_all_servers() # 因为异步加载,无法放在 RuntimeConfigMiddleware 的 __init__ 中
# 使用 create_agent 创建智能体 # 使用 create_agent 创建智能体
# 注意tools 参数由 RuntimeConfigMiddleware 在 wrap_model_call 中动态设置
graph = create_agent( graph = create_agent(
model=load_chat_model(context.model), model=load_chat_model(context.model),
tools=get_buildin_tools(),
system_prompt=context.system_prompt, system_prompt=context.system_prompt,
middleware=[ middleware=[
inject_attachment_context, # 附件上下文注入 inject_attachment_context, # 附件上下文注入
RuntimeConfigMiddleware(), # 运行时配置应用(模型/工具/知识库/MCP/提示词) RuntimeConfigMiddleware(extra_tools=all_mcp_tools), # 运行时配置应用(模型/工具/知识库/MCP/提示词)
ModelRetryMiddleware(), # 模型重试中间件 ModelRetryMiddleware(), # 模型重试中间件
], ],
checkpointer=await self._get_checkpointer(), checkpointer=await self._get_checkpointer(),

View File

@ -1,14 +1,14 @@
from __future__ import annotations from __future__ import annotations
from collections.abc import Callable from collections.abc import Callable
from dataclasses import replace
from typing import Any from typing import Any
from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse
from src import knowledge_base from src.utils.logging_config import logger
from src.agents.common import load_chat_model from src.agents.common import load_chat_model
from src.agents.common.tools import get_tools_from_context from src.agents.common.tools import get_kb_based_tools, get_buildin_tools
from src.services.mcp_service import get_enabled_mcp_tools
def _is_system_message(msg: Any) -> bool: def _is_system_message(msg: Any) -> bool:
@ -28,38 +28,32 @@ def _get_message_content(msg: Any) -> str | None:
class RuntimeConfigMiddleware(AgentMiddleware): class RuntimeConfigMiddleware(AgentMiddleware):
"""运行时配置中间件 - 应用模型/工具/知识库/MCP/提示词配置
注意所有可能用到的知识库工具必须在初始化时预加载并注册到 self.tools
运行时根据配置从 self.tools 中筛选工具不能动态添加新工具
"""
def __init__(self, *, extra_tools: list[Any] | None = None):
"""初始化中间件
Args:
extra_tools: 额外工具列表 create_agent tools 参数传入
"""
super().__init__()
# 这里的工具只是提供给 langchain 调用,并不是真正的绑定在模型上
self.kb_tools = get_kb_based_tools()
self.tools = self.kb_tools + (extra_tools or [])
logger.debug(f"Initialized tools: {len(self.tools)}")
async def awrap_model_call( async def awrap_model_call(
self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse]
) -> ModelResponse: ) -> ModelResponse:
# 虽然功能实现了, 但是总感觉怪怪的 TODO 优化下
runtime_context = request.runtime.context runtime_context = request.runtime.context
effective_context = runtime_context
blocked_knowledges: list[str] = []
requested_knowledges = getattr(runtime_context, "knowledges", None)
department_id = getattr(runtime_context, "department_id", None)
if department_id and isinstance(requested_knowledges, list) and requested_knowledges:
user_info = {"role": "user", "department_id": department_id}
accessible_databases = await knowledge_base.get_databases_by_user(user_info)
accessible_kb_names = {
db.get("name")
for db in accessible_databases.get("databases", [])
if isinstance(db, dict) and db.get("name")
}
filtered_knowledges = [kb for kb in requested_knowledges if kb in accessible_kb_names]
blocked_knowledges = [kb for kb in requested_knowledges if kb not in accessible_kb_names]
if blocked_knowledges:
effective_context = replace(runtime_context)
effective_context.knowledges = filtered_knowledges
model = load_chat_model(getattr(runtime_context, "model", None)) model = load_chat_model(getattr(runtime_context, "model", None))
tools = await get_tools_from_context(effective_context) enabled_tools = await self.get_tools_from_context(runtime_context)
system_prompt = getattr(runtime_context, "system_prompt", None) system_prompt = getattr(runtime_context, "system_prompt", None)
notice = None
if blocked_knowledges:
notice = f"注意:已自动过滤无权访问的知识库:{', '.join(blocked_knowledges)}"
existing_systems: list[Any] = [] existing_systems: list[Any] = []
remaining: list[Any] = [] remaining: list[Any] = []
@ -83,10 +77,33 @@ class RuntimeConfigMiddleware(AgentMiddleware):
new_systems.append(existing_systems.pop(idx)) new_systems.append(existing_systems.pop(idx))
existing_contents.pop(idx) existing_contents.pop(idx)
if notice and notice not in existing_contents:
new_systems.append({"role": "system", "content": notice})
messages = [*new_systems, *existing_systems, *remaining] messages = [*new_systems, *existing_systems, *remaining]
request = request.override(model=model, tools=tools, messages=messages) request = request.override(model=model, tools=enabled_tools, messages=messages)
return await handler(request) return await handler(request)
async def get_tools_from_context(self, context) -> list:
"""从上下文配置中获取工具列表"""
# 1. 基础工具 (从 context.tools 中筛选)
selected_tools = []
if context.tools:
# 创建工具映射表
tools_map = {t.name: t for t in self.tools}
for tool_name in context.tools:
if tool_name in tools_map:
selected_tools.append(tools_map[tool_name])
# 2. 知识库工具
if context.knowledges:
kb_tools = get_kb_based_tools(db_names=context.knowledges)
selected_tools.extend(kb_tools)
# 3. MCP 工具(使用统一入口,自动过滤 disabled_tools
if context.mcps:
for server_name in context.mcps:
mcp_tools = await get_enabled_mcp_tools(server_name)
selected_tools.extend(mcp_tools)
return selected_tools

View File

@ -4,11 +4,12 @@ from deepagents.middleware.filesystem import FilesystemMiddleware
from deepagents.middleware.patch_tool_calls import PatchToolCallsMiddleware from deepagents.middleware.patch_tool_calls import PatchToolCallsMiddleware
from deepagents.middleware.subagents import SubAgentMiddleware from deepagents.middleware.subagents import SubAgentMiddleware
from langchain.agents import create_agent from langchain.agents import create_agent
from langchain.agents.middleware import ModelRequest, SummarizationMiddleware, TodoListMiddleware, dynamic_prompt from langchain.agents.middleware import ModelRequest, ModelRetryMiddleware, SummarizationMiddleware, TodoListMiddleware, dynamic_prompt
from src.agents.common import BaseAgent, load_chat_model from src.agents.common import BaseAgent, load_chat_model
from src.agents.common.middlewares import inject_attachment_context from src.agents.common.middlewares import RuntimeConfigMiddleware, inject_attachment_context
from src.agents.common.tools import get_tavily_search from src.agents.common.tools import get_tavily_search
from src.services.mcp_service import get_tools_from_all_servers
from .context import DeepContext from .context import DeepContext
from .prompts import DEEP_PROMPT from .prompts import DEEP_PROMPT
@ -93,24 +94,26 @@ class DeepAgent(BaseAgent):
model = load_chat_model(context.model) model = load_chat_model(context.model)
sub_model = load_chat_model(context.subagents_model) sub_model = load_chat_model(context.subagents_model)
tools = await self.get_tools() search_tools = await self.get_tools()
all_mcp_tools = await get_tools_from_all_servers()
# 合并搜索工具和 MCP 工具
# Build subagents with search tools # Build subagents with search tools
research_sub_agent = _get_research_sub_agent(tools) research_sub_agent = _get_research_sub_agent(search_tools)
# 使用 create_deep_agent 创建深度智能体 # 使用 create_deep_agent 创建深度智能体
graph = create_agent( graph = create_agent(
model=model, model=model,
tools=tools,
system_prompt=context.system_prompt, system_prompt=context.system_prompt,
middleware=[ middleware=[
context_aware_prompt, # 动态系统提示词 context_aware_prompt, # 动态系统提示词
inject_attachment_context, # 附件上下文注入 inject_attachment_context, # 附件上下文注入
RuntimeConfigMiddleware(extra_tools=all_mcp_tools),
TodoListMiddleware(), TodoListMiddleware(),
FilesystemMiddleware(), FilesystemMiddleware(),
SubAgentMiddleware( SubAgentMiddleware(
default_model=sub_model, default_model=sub_model,
default_tools=tools, default_tools=search_tools,
subagents=[critique_sub_agent, research_sub_agent], subagents=[critique_sub_agent, research_sub_agent],
default_middleware=[ default_middleware=[
TodoListMiddleware(), # 子智能体也有 todo 列表 TodoListMiddleware(), # 子智能体也有 todo 列表

View File

@ -1,7 +1,10 @@
from langchain.agents import create_agent from langchain.agents import create_agent
from src.agents.common import BaseAgent, load_chat_model from src.agents.common import BaseAgent, load_chat_model
from src.agents.common.tools import get_tools_from_context from src.agents.common.middlewares import (
RuntimeConfigMiddleware,
)
from src.services.mcp_service import get_tools_from_all_servers
class MiniAgent(BaseAgent): class MiniAgent(BaseAgent):
@ -12,13 +15,16 @@ class MiniAgent(BaseAgent):
super().__init__(**kwargs) super().__init__(**kwargs)
async def get_graph(self, **kwargs): async def get_graph(self, **kwargs):
"""构建图"""
context = self.context_schema.from_file(module_name=self.module_name) context = self.context_schema.from_file(module_name=self.module_name)
all_mcp_tools = await get_tools_from_all_servers()
# 创建 MiniAgent
graph = create_agent( graph = create_agent(
model=load_chat_model(context.model), model=load_chat_model(context.model),
system_prompt=context.system_prompt, system_prompt=context.system_prompt,
tools=await get_tools_from_context(context), middleware=[
RuntimeConfigMiddleware(extra_tools=all_mcp_tools),
],
checkpointer=await self._get_checkpointer(), checkpointer=await self._get_checkpointer(),
) )

View File

@ -2,10 +2,15 @@ from dataclasses import dataclass, field
from typing import Annotated from typing import Annotated
from langchain.agents import create_agent from langchain.agents import create_agent
from langchain.agents.middleware import ModelRetryMiddleware
from src.agents.common import BaseAgent, BaseContext, load_chat_model from src.agents.common import BaseAgent, BaseContext, load_chat_model
from src.agents.common.middlewares import (
RuntimeConfigMiddleware,
)
from src.agents.common.toolkits.mysql import get_mysql_tools from src.agents.common.toolkits.mysql import get_mysql_tools
from src.agents.common.tools import gen_tool_info, get_buildin_tools, get_tools_from_context from src.agents.common.tools import gen_tool_info, get_buildin_tools
from src.services.mcp_service import get_tools_from_all_servers
from src.utils import logger from src.utils import logger
@ -35,13 +40,18 @@ class SqlReporterAgent(BaseAgent):
super().__init__(**kwargs) super().__init__(**kwargs)
async def get_graph(self, **kwargs): async def get_graph(self, **kwargs):
"""构建图"""
context = self.context_schema.from_file(module_name=self.module_name) context = self.context_schema.from_file(module_name=self.module_name)
all_mcp_tools = await get_tools_from_all_servers()
# 合并 MySQL 工具和 MCP 工具
extra_tools = get_mysql_tools() + all_mcp_tools
# 创建 SqlReporterAgent
graph = create_agent( graph = create_agent(
model=load_chat_model(context.model), # 使用 context 中的模型配置 model=load_chat_model(context.model),
system_prompt=context.system_prompt, system_prompt=context.system_prompt,
tools=await get_tools_from_context(context, extra_tools=get_mysql_tools()), middleware=[
RuntimeConfigMiddleware(extra_tools=extra_tools),
],
checkpointer=await self._get_checkpointer(), checkpointer=await self._get_checkpointer(),
) )

View File

@ -331,6 +331,24 @@ async def stream_agent_chat(
logger.error(f"Error loading attachments for thread_id={thread_id}: {e}") logger.error(f"Error loading attachments for thread_id={thread_id}: {e}")
input_context["attachments"] = [] input_context["attachments"] = []
# 根据用户权限过滤知识库
requested_knowledge_names = input_context.get("knowledges")
if requested_knowledge_names and isinstance(requested_knowledge_names, list) and requested_knowledge_names:
user_info = {"role": "user", "department_id": department_id}
accessible_databases = await knowledge_base.get_databases_by_user(user_info)
accessible_kb_names = {
db.get("name")
for db in accessible_databases.get("databases", [])
if isinstance(db, dict) and db.get("name")
}
filtered_knowledge_names = [kb for kb in requested_knowledge_names if kb in accessible_kb_names]
blocked_knowledge_names = [kb for kb in requested_knowledge_names if kb not in accessible_kb_names]
if blocked_knowledge_names:
logger.warning(
f"用户 {user_id} 无权访问知识库: {blocked_knowledge_names}, 已自动过滤"
)
input_context["knowledges"] = filtered_knowledge_names
full_msg = None full_msg = None
accumulated_content = [] accumulated_content = []
langgraph_config = {"configurable": {"thread_id": thread_id, "user_id": user_id}} langgraph_config = {"configurable": {"thread_id": thread_id, "user_id": user_id}}