diff --git a/libs/langgraph/langgraph/agent/__init__.py b/libs/langgraph/langgraph/agent/__init__.py index b665dc26e..5abe55da7 100644 --- a/libs/langgraph/langgraph/agent/__init__.py +++ b/libs/langgraph/langgraph/agent/__init__.py @@ -81,9 +81,17 @@ def create_agent( graph.add_node("tools", tool_node) for m in middleware: if m.__class__.before_model is not AgentMiddleware.before_model: - graph.add_node(f"{m.__class__.__name__}.before_model", m.before_model) + graph.add_node( + f"{m.__class__.__name__}.before_model", + m.before_model, + input_schema=m.State, + ) if m.__class__.after_model is not AgentMiddleware.after_model: - graph.add_node(f"{m.__class__.__name__}.after_model", m.after_model) + graph.add_node( + f"{m.__class__.__name__}.after_model", + m.after_model, + input_schema=m.State, + ) # add start edge first_node = ( diff --git a/libs/langgraph/langgraph/agent/types.py b/libs/langgraph/langgraph/agent/types.py index a1ce5ce27..38d907b9f 100644 --- a/libs/langgraph/langgraph/agent/types.py +++ b/libs/langgraph/langgraph/agent/types.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Sequence from dataclasses import dataclass -from typing import Annotated, Any, Literal, Self +from typing import Annotated, Any, Literal from langchain_core.language_models.chat_models import BaseChatModel from langchain_core.messages import AnyMessage @@ -27,22 +27,24 @@ class ModelRequest: response_format: ResponseFormat | None +@dataclass +class AgentState: + messages: Annotated[list[AnyMessage], add_messages] + jump_to: Annotated[JumpTo | None, EphemeralValue] = None + response: dict | None = None + + class AgentMiddleware: - def __init__(self) -> None: + class State(AgentState): pass - def __copy__(self) -> Self: - return self.__class__(**self.__dict__) - - def before_model(self, state: AgentState) -> AgentUpdate | AgentJump | None: + def before_model(self, state: State) -> AgentUpdate | AgentJump | None: pass - def modify_model_request( - self, request: ModelRequest, state: AgentState - ) -> ModelRequest: + def modify_model_request(self, request: ModelRequest, state: State) -> ModelRequest: return request - def after_model(self, state: AgentState) -> AgentUpdate | AgentJump | None: + def after_model(self, state: State) -> AgentUpdate | AgentJump | None: pass @@ -54,10 +56,3 @@ class AgentUpdate(TypedDict, total=False): class AgentJump(TypedDict, total=False): messages: Messages jump_to: JumpTo - - -@dataclass -class AgentState: - messages: Annotated[list[AnyMessage], add_messages] - jump_to: Annotated[JumpTo | None, EphemeralValue] = None - response: dict | None = None