LangChain 源码剖析-自定义中间件详解(Custom middleware)
- 通过实现代理执行流程中特定点运行的钩子来构建自定义中间件。
Hook(钩子)
Node-style hooks(节点样式挂钩)
- 在特定执行点按顺序运行。用于日志记录、验证和状态更新。
before_agent - 在代理启动之前(每次调用一次)before_model - 每次模型调用之前after_model - 每个模型响应后after_agent - 代理完成后(每次调用一次)
from langchain.agents.middleware import before_model, after_model, AgentStatefrom langchain.messages import AIMessagefrom langgraph.runtime import Runtimefrom typing import Any@before_model(can_jump_to=["end"])def check_message_limit(state: AgentState, runtime: Runtime) -> dict[str, Any] | None: if len(state["messages"]) >= 50: return { "messages": [AIMessage("Conversation limit reached.")], "jump_to": "end" } return None@after_modeldef log_response(state: AgentState, runtime: Runtime) -> dict[str, Any] | None: print(f"Model returned: {state['messages'][-1].content}") return None
from langchain.agents.middleware import AgentMiddleware, AgentState, hook_configfrom langchain.messages import AIMessagefrom langgraph.runtime import Runtimefrom typing import Anyclass MessageLimitMiddleware(AgentMiddleware): def __init__(self, max_messages: int = 50): super().__init__() self.max_messages = max_messages @hook_config(can_jump_to=["end"]) def before_model(self, state: AgentState, runtime: Runtime) -> dict[str, Any] | None: if len(state["messages"]) == self.max_messages: return { "messages": [AIMessage("Conversation limit reached.")], "jump_to": "end" } return None def after_model(self, state: AgentState, runtime: Runtime) -> dict[str, Any] | None: print(f"Model returned: {state['messages'][-1].content}") return None
Wrap-style hooks(装饰器样式挂钩)
@wrap_model_call - 用自定义逻辑包装每个模型调用@wrap_tool_call - 用自定义逻辑包装每个工具调用
Convenience(动态便利性)
@dynamic_prompt - 生成动态系统提示
from langchain.agents.middleware import ( before_model, wrap_model_call, AgentState, ModelRequest, ModelResponse,)from langchain.agents import create_agentfrom langgraph.runtime import Runtimefrom typing import Any, Callable@before_modeldef log_before_model(state: AgentState, runtime: Runtime) -> dict[str, Any] | None: print(f"About to call model with {len(state['messages'])} messages") return None@wrap_model_calldef retry_model( request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse],) -> ModelResponse: for attempt in range(3): try: return handler(request) except Exception as e: if attempt == 2: raise print(f"Retry {attempt + 1}/3 after error: {e}")agent = create_agent( model="gpt-4o", middleware=[log_before_model, retry_model], tools=[...],)
何时使用装饰器:
基于类的中间件
- 对于具有多个钩子或配置的复杂中间件更强大。当您需要为同一个钩子定义同步和异步实现时,或者当您想在单个中间件中组合多个钩子时,请使用类。
from langchain.agents.middleware import ( AgentMiddleware, AgentState, ModelRequest, ModelResponse,)from langgraph.runtime import Runtimefrom typing import Any, Callableclass LoggingMiddleware(AgentMiddleware): def before_model(self, state: AgentState, runtime: Runtime) -> dict[str, Any] | None: print(f"About to call model with {len(state['messages'])} messages") return None def after_model(self, state: AgentState, runtime: Runtime) -> dict[str, Any] | None: print(f"Model returned: {state['messages'][-1].content}") return Noneagent = create_agent( model="gpt-4o", middleware=[LoggingMiddleware()], tools=[...],)
何时使用类:
自定义状态架构
- 中间件可以使用自定义属性扩展代理的状态。这使得中间件能够:
- 跨执行跟踪状态:维护在代理执行生命周期中持续存在的计数器、标志或其他值- 在钩子之间共享数据:将信息从before_model传递到after_model或在不同的中间件实例之间传递- 实现跨领域关注点:添加限速、使用跟踪、用户上下文或审计日志等功能,而无需修改核心代理逻辑- 做出条件决策:使用累积状态来确定是继续执行、跳转到不同节点还是动态修改行为
from langchain.agents import create_agentfrom langchain.messages import HumanMessagefrom langchain.agents.middleware import AgentState, before_model, after_modelfrom typing_extensions import NotRequiredfrom typing import Anyfrom langgraph.runtime import Runtimeclass CustomState(AgentState): model_call_count: NotRequired[int] user_id: NotRequired[str]@before_model(state_schema=CustomState, can_jump_to=["end"])def check_call_limit(state: CustomState, runtime: Runtime) -> dict[str, Any] | None: count = state.get("model_call_count", 0) if count > 10: return {"jump_to": "end"} return None@after_model(state_schema=CustomState)def increment_counter(state: CustomState, runtime: Runtime) -> dict[str, Any] | None: return {"model_call_count": state.get("model_call_count", 0) + 1}agent = create_agent( model="gpt-4o", middleware=[check_call_limit, increment_counter], tools=[],)# Invoke with custom stateresult = agent.invoke({ "messages": [HumanMessage("Hello")], "model_call_count": 0, "user_id": "user-123",})
from langchain.agents import create_agentfrom langchain.messages import HumanMessagefrom langchain.agents.middleware import AgentState, AgentMiddlewarefrom typing_extensions import NotRequiredfrom typing import Anyclass CustomState(AgentState): model_call_count: NotRequired[int] user_id: NotRequired[str]class CallCounterMiddleware(AgentMiddleware[CustomState]): state_schema = CustomState def before_model(self, state: CustomState, runtime) -> dict[str, Any] | None: count = state.get("model_call_count", 0) if count > 10: return {"jump_to": "end"} return None def after_model(self, state: CustomState, runtime) -> dict[str, Any] | None: return {"model_call_count": state.get("model_call_count", 0) + 1}agent = create_agent( model="gpt-4o", middleware=[CallCounterMiddleware()], tools=[],)# Invoke with custom stateresult = agent.invoke({ "messages": [HumanMessage("Hello")], "model_call_count": 0, "user_id": "user-123",})
中间件执行顺序
agent = create_agent( model="gpt-4o", middleware=[middleware1, middleware2, middleware3], tools=[...],)
1. middleware1.before_agent()2. middleware2.before_agent()3. middleware3.before_agent()
4. middleware1.before_model()5. middleware2.before_model()6. middleware3.before_model()
7. middleware1.wrap_model_call() → middleware2.wrap_model_call() → middleware3.wrap_model_call() → model
8. middleware3.after_model()9. middleware2.after_model()10. middleware1.after_model()
11. middleware3.after_agent()12. middleware2.after_agent()13. middleware1.after_agent()
关键规则:
- wrap_* hooks: 嵌套(第一个中间件包裹所有其他中间件)
代理跳转执行
- 要提前退出中间件,请使用jump_To返回一个字典:
'end':跳转到代理执行的末尾(或第一个after_agent钩子)'tools':跳转到工具节点'model':跳转到model节点(或第一个before_model钩子)
from langchain.agents.middleware import after_model, hook_config, AgentStatefrom langchain.messages import AIMessagefrom langgraph.runtime import Runtimefrom typing import Any@after_model@hook_config(can_jump_to=["end"])def check_for_blocked(state: AgentState, runtime: Runtime) -> dict[str, Any] | None: last_message = state["messages"][-1] if "BLOCKED" in last_message.content: return { "messages": [AIMessage("I cannot respond to that request.")], "jump_to": "end" } return None
from langchain.agents.middleware import AgentMiddleware, hook_config, AgentStatefrom langchain.messages import AIMessagefrom langgraph.runtime import Runtimefrom typing import Anyclass BlockedContentMiddleware(AgentMiddleware): @hook_config(can_jump_to=["end"]) def after_model(self, state: AgentState, runtime: Runtime) -> dict[str, Any] | None: last_message = state["messages"][-1] if "BLOCKED" in last_message.content: return { "messages": [AIMessage("I cannot respond to that request.")], "jump_to": "end" } return None
最佳实践
动态模型选择中间件
from langchain.agents.middleware import wrap_model_call, ModelRequest, ModelResponsefrom langchain.chat_models import init_chat_modelfrom typing import Callablecomplex_model = init_chat_model("gpt-4o")simple_model = init_chat_model("gpt-4o-mini")@wrap_model_calldef dynamic_model( request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse],) -> ModelResponse: # Use different model based on conversation length if len(request.messages) > 10: model = complex_model else: model = simple_model return handler(request.override(model=model))
from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponsefrom langchain.chat_models import init_chat_modelfrom typing import Callablecomplex_model = init_chat_model("gpt-4o")simple_model = init_chat_model("gpt-4o-mini")class DynamicModelMiddleware(AgentMiddleware): def wrap_model_call( self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelResponse: # Use different model based on conversation length if len(request.messages) > 10: model = complex_model else: model = simple_model return handler(request.override(model=model))
工具调用监控
from langchain.agents.middleware import wrap_tool_callfrom langchain.tools.tool_node import ToolCallRequestfrom langchain.messages import ToolMessagefrom langgraph.types import Commandfrom typing import Callable@wrap_tool_calldef monitor_tool( request: ToolCallRequest, handler: Callable[[ToolCallRequest], ToolMessage | Command],) -> ToolMessage | Command: print(f"Executing tool: {request.tool_call['name']}") print(f"Arguments: {request.tool_call['args']}") try: result = handler(request) print(f"Tool completed successfully") return result except Exception as e: print(f"Tool failed: {e}") raise
from langchain.tools.tool_node import ToolCallRequestfrom langchain.agents.middleware import AgentMiddlewarefrom langchain.messages import ToolMessagefrom langgraph.types import Commandfrom typing import Callableclass ToolMonitoringMiddleware(AgentMiddleware): def wrap_tool_call( self, request: ToolCallRequest, handler: Callable[[ToolCallRequest], ToolMessage | Command], ) -> ToolMessage | Command: print(f"Executing tool: {request.tool_call['name']}") print(f"Arguments: {request.tool_call['args']}") try: result = handler(request) print(f"Tool completed successfully") return result except Exception as e: print(f"Tool failed: {e}") raise
动态选择工具
from langchain.agents import create_agentfrom langchain.agents.middleware import wrap_model_call, ModelRequest, ModelResponsefrom typing import Callable@wrap_model_calldef select_tools( request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse],) -> ModelResponse: """Middleware to select relevant tools based on state/context.""" # Select a small, relevant subset of tools based on state/context relevant_tools = select_relevant_tools(request.state, request.runtime) return handler(request.override(tools=relevant_tools))agent = create_agent( model="gpt-4o", tools=all_tools, # All available tools need to be registered upfront middleware=[select_tools],)
from langchain.agents import create_agentfrom langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponsefrom typing import Callableclass ToolSelectorMiddleware(AgentMiddleware): def wrap_model_call( self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelResponse: """Middleware to select relevant tools based on state/context.""" # Select a small, relevant subset of tools based on state/context relevant_tools = select_relevant_tools(request.state, request.runtime) return handler(request.override(tools=relevant_tools))agent = create_agent( model="gpt-4o", tools=all_tools, # All available tools need to be registered upfront middleware=[ToolSelectorMiddleware()],)
处理系统消息
- 使用ModelRequest上的system_message字段修改中间件中的系统消息。
- system_message字段包含一个SystemMessage对象(即使代理是用字符串system_prompt创建的)。
向系统消息添加上下文:
from langchain.agents.middleware import wrap_model_call, ModelRequest, ModelResponsefrom langchain.messages import SystemMessagefrom typing import Callable@wrap_model_calldef add_context( request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse],) -> ModelResponse: # Always work with content blocks new_content = list(request.system_message.content_blocks) + [ {"type": "text", "text": "Additional context."} ] new_system_message = SystemMessage(content=new_content) return handler(request.override(system_message=new_system_message))
from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponsefrom langchain.messages import SystemMessagefrom typing import Callableclass ContextMiddleware(AgentMiddleware): def wrap_model_call( self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelResponse: # Always work with content blocks new_content = list(request.system_message.content_blocks) + [ {"type": "text", "text": "Additional context."} ] new_system_message = SystemMessage(content=new_content) return handler(request.override(system_message=new_system_message))
使用缓存控制(Anthropic)
- 使用Anthropic模型时,您可以使用带有缓存控制指令的结构化内容块来缓存大型系统提示:
from langchain.agents.middleware import wrap_model_call, ModelRequest, ModelResponsefrom langchain.messages import SystemMessagefrom typing import Callable@wrap_model_calldef add_cached_context( request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse],) -> ModelResponse: # Always work with content blocks new_content = list(request.system_message.content_blocks) + [ { "type": "text", "text": "Here is a large document to analyze:\n\n<document>...</document>", # content up until this point is cached "cache_control": {"type": "ephemeral"} } ] new_system_message = SystemMessage(content=new_content) return handler(request.override(system_message=new_system_message))
from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponsefrom langchain.messages import SystemMessagefrom typing import Callableclass CachedContextMiddleware(AgentMiddleware): def wrap_model_call( self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelResponse: # Always work with content blocks new_content = list(request.system_message.content_blocks) + [ { "type": "text", "text": "Here is a large document to analyze:\n\n<document>...</document>", "cache_control": {"type": "ephemeral"} # This content will be cached } ] new_system_message = SystemMessage(content=new_content) return handler(request.override(system_message=new_system_message))
注意
- ModelRequest.system_message始终是SystemMessage对象,即使代理是使用system_prompt=“string”创建的
- 使用SystemMessage.content_blocks以块列表的形式访问内容,无论原始内容是字符串还是列表
- 修改系统消息时,使用content_blocks并附加新块以保留现有结构
- 对于缓存控制等高级用例,您可以将SystemMessage对象直接传递给create_agent的system_prompt参数