from collections.abc import Callable from typing import Any from langchain.agents import create_agent from langchain.agents.middleware import ( AgentMiddleware, AgentState, ModelRequest, ModelResponse, ) from langchain.messages import SystemMessage from langchain_core.language_models.fake_chat_models import FakeListChatModel from langgraph.runtime import Runtime from typing_extensions import NotRequired class AgentReadyFakeModel(FakeListChatModel): def bind_tools(self, tools, *, tool_choice=None, **kwargs): if tools: raise NotImplementedError("This smoke test does not exercise tools.") return self class TraceState(AgentState): model_call_count: NotRequired[int] last_user_message: NotRequired[str] class TraceModelCallMiddleware(AgentMiddleware[TraceState]): state_schema = TraceState def before_model( self, state: TraceState, runtime: Runtime ) -> dict[str, Any] | None: count = state.get("model_call_count", 0) + 1 user_messages = [ message for message in state["messages"] if getattr(message, "type", None) == "human" ] latest = str(user_messages[-1].content) if user_messages else "" print(f"middleware before_model: call={count}, latest={latest!r}") return { "model_call_count": count, "last_user_message": latest, } def wrap_model_call( self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelResponse: content_blocks = list(request.system_message.content_blocks) content_blocks.append( {"type": "text", "text": "Add the middleware trace marker."} ) patched_request = request.override( system_message=SystemMessage(content=content_blocks) ) print(f"middleware wrap_model_call: messages={len(request.messages)}") return handler(patched_request) model = AgentReadyFakeModel( responses=["middleware trace marker: custom hook ran."] ) agent = create_agent( model=model, tools=[], system_prompt="Answer briefly.", middleware=[TraceModelCallMiddleware()], ) result = agent.invoke( { "messages": [ {"role": "user", "content": "Test the custom middleware."} ], "model_call_count": 0, } ) print(f"Agent reply: {result['messages'][-1].content}") print(f"Model calls recorded: {result['model_call_count']}") print(f"Last user message: {result['last_user_message']}")