本页目录
Middleware / 中间件Note 92

06 类中间件与 `state_schema`

AgentMiddleware 类写法适合把多个 hook 和它们共享的配置放在一起。stateschema 用来扩展 agent state,让 middleware 能把计数、标记、审计结果这类运行时信息写回最终 state。它不是全局变量,也不是系统提示词,作用域是这次 agent run 的状态。

AgentMiddleware 类写法适合把多个 hook 和它们共享的配置放在一起。state_schema 用来扩展 agent state,让 middleware 能把计数、标记、审计结果这类运行时信息写回最终 state。它不是全局变量,也不是系统提示词,作用域是这次 agent run 的状态。

最小代码

代码在:

deepagent_src/middleware_teach/06_class_state_schema.py

核心逻辑:

class AuditState(AgentState):
    model_call_count: NotRequired[int]
    last_model_text: NotRequired[str]


class AuditStateMiddleware(AgentMiddleware[AuditState, Any]):
    state_schema = AuditState

    def before_model(self, state: AuditState, runtime: Runtime[Any]):
        return {"model_call_count": state.get("model_call_count", 0) + 1}

    def after_model(self, state: AuditState, runtime: Runtime[Any]):
        return {"last_model_text": state["messages"][-1].text}

before_model 在模型调用前把计数加一,after_model 在模型返回后记录最后的模型文本。两个 hook 共享同一个扩展状态结构。

运行

uv run python -m deepagent_src.middleware_teach.06_class_state_schema

预期输出包含:

model_call_count: 1
last_model_text: STATE_SCHEMA_OK
class middleware state_schema local check ok

常见误区

不要把 state_schema 当数据库用。它只描述 agent state 中允许出现的字段,适合保存一次运行里的轻量状态;跨线程、跨会话持久化要用 checkpointer、store、Deep Agents memory 或后端存储。

和前面章节的关系

  • 装饰器写法:适合一个简单 hook。
  • 类写法:适合多个 hook、共享配置、复用成生产 middleware。
  • state_schema:让 middleware 的状态字段有明确结构,避免到处塞魔法 key,艹,那种隐式字典最容易把人坑死。

相关资源

  • 查看示例代码:deepagent_src/middleware_teach/06_class_state_schema.py
    from __future__ import annotations
    
    from typing import Any, NotRequired
    
    from langchain.agents import create_agent
    from langchain.agents.middleware import AgentMiddleware, AgentState
    from langchain_core.language_models.fake_chat_models import FakeListChatModel
    from langgraph.runtime import Runtime
    
    
    class AuditState(AgentState):
        model_call_count: NotRequired[int]
        last_model_text: NotRequired[str]
    
    
    class AuditStateMiddleware(AgentMiddleware[AuditState, Any]):
        state_schema = AuditState
    
        def before_model(
            self,
            state: AuditState,
            runtime: Runtime[Any],
        ) -> dict[str, Any] | None:
            return {"model_call_count": state.get("model_call_count", 0) + 1}
    
        def after_model(
            self,
            state: AuditState,
            runtime: Runtime[Any],
        ) -> dict[str, Any] | None:
            return {"last_model_text": state["messages"][-1].text}
    
    
    def main() -> None:
        agent = create_agent(
            model=FakeListChatModel(responses=["STATE_SCHEMA_OK"]),
            tools=[],
            middleware=[AuditStateMiddleware()],
        )
        state = agent.invoke({"messages": [{"role": "user", "content": "hi"}]})
    
        print("model_call_count:", state["model_call_count"])
        print("last_model_text:", state["last_model_text"])
        print("final:", state["messages"][-1].text)
    
        assert state["model_call_count"] == 1, state
        assert state["last_model_text"] == "STATE_SCHEMA_OK", state
        assert state["messages"][-1].text == "STATE_SCHEMA_OK"
        print("class middleware state_schema local check ok")
    
    
    if __name__ == "__main__":
        main()