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()