Add State property

This commit is contained in:
Nuno Campos
2025-09-03 14:43:46 +01:00
parent 0c929e62eb
commit 46ce6ad927
2 changed files with 22 additions and 19 deletions
+10 -2
View File
@@ -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 = (
+12 -17
View File
@@ -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