乐于分享
好东西不私藏

LangChain 源码剖析-自定义中间件详解(Custom middleware)

LangChain 源码剖析-自定义中间件详解(Custom middleware)

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[strAny] | 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[strAny] | 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[strAny] | 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[strAny] | 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 AnyCallable@before_modeldef log_before_model(state: AgentState, runtime: Runtime) -> dict[strAny] | 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 AnyCallableclass LoggingMiddleware(AgentMiddleware):    def before_model(self, state: AgentState, runtime: Runtime) -> dict[strAny] | None:        print(f"About to call model with {len(state['messages'])} messages")        return None    def after_model(self, state: AgentState, runtime: Runtime) -> dict[strAny] | 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[strAny] | 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[strAny] | 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[strAny] | 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[strAny] | 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()
  • Wrap钩子嵌套式函数调用:
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()

关键规则:

  • before_* hooks: 从第一到最后
  • after_* hooks: 从最后到第一
  • 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[strAny] | 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[strAny] | 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参数