diff --git a/libs/langgraph/tests/__snapshots__/test_large_cases.ambr b/libs/langgraph/tests/__snapshots__/test_large_cases.ambr index 34e4ed64c..3df0b4159 100644 --- a/libs/langgraph/tests/__snapshots__/test_large_cases.ambr +++ b/libs/langgraph/tests/__snapshots__/test_large_cases.ambr @@ -175,10 +175,10 @@ ''' # --- # name: test_prebuilt_tool_chat - '{"$defs": {"BaseMessage": {"additionalProperties": true, "description": "Base abstract message class.\\n\\nMessages are the inputs and outputs of ChatModels.", "properties": {"content": {"anyOf": [{"type": "string"}, {"items": {"anyOf": [{"type": "string"}, {"additionalProperties": true, "type": "object"}]}, "type": "array"}], "title": "Content"}, "additional_kwargs": {"additionalProperties": true, "title": "Additional Kwargs", "type": "object"}, "response_metadata": {"additionalProperties": true, "title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "type": "string"}, "name": {"anyOf": [{"type": "string"}, {"type": "null"}], "default": null, "title": "Name"}, "id": {"anyOf": [{"type": "string"}, {"type": "null"}], "default": null, "title": "Id"}}, "required": ["content", "type"], "title": "BaseMessage", "type": "object"}}, "description": "The state of the agent.", "properties": {"messages": {"items": {"$ref": "#/$defs/BaseMessage"}, "title": "Messages", "type": "array"}, "remaining_steps": {"title": "Remaining Steps", "type": "integer"}}, "required": ["messages"], "title": "AgentState", "type": "object"}' + '{"$defs": {"BaseMessage": {"additionalProperties": true, "description": "Base abstract message class.\\n\\nMessages are the inputs and outputs of ChatModels.", "properties": {"content": {"anyOf": [{"type": "string"}, {"items": {"anyOf": [{"type": "string"}, {"additionalProperties": true, "type": "object"}]}, "type": "array"}], "title": "Content"}, "additional_kwargs": {"additionalProperties": true, "title": "Additional Kwargs", "type": "object"}, "response_metadata": {"additionalProperties": true, "title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "type": "string"}, "name": {"anyOf": [{"type": "string"}, {"type": "null"}], "default": null, "title": "Name"}, "id": {"anyOf": [{"type": "string"}, {"type": "null"}], "default": null, "title": "Id"}}, "required": ["content", "type"], "title": "BaseMessage", "type": "object"}}, "description": "The state of the agent.", "properties": {"messages": {"items": {"$ref": "#/$defs/BaseMessage"}, "title": "Messages", "type": "array"}, "remaining_steps": {"title": "Remaining Steps", "type": "integer"}, "structured_response": {"title": "Structured Response", "type": "null"}}, "required": ["messages"], "title": "AgentState", "type": "object"}' # --- # name: test_prebuilt_tool_chat.1 - '{"$defs": {"BaseMessage": {"additionalProperties": true, "description": "Base abstract message class.\\n\\nMessages are the inputs and outputs of ChatModels.", "properties": {"content": {"anyOf": [{"type": "string"}, {"items": {"anyOf": [{"type": "string"}, {"additionalProperties": true, "type": "object"}]}, "type": "array"}], "title": "Content"}, "additional_kwargs": {"additionalProperties": true, "title": "Additional Kwargs", "type": "object"}, "response_metadata": {"additionalProperties": true, "title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "type": "string"}, "name": {"anyOf": [{"type": "string"}, {"type": "null"}], "default": null, "title": "Name"}, "id": {"anyOf": [{"type": "string"}, {"type": "null"}], "default": null, "title": "Id"}}, "required": ["content", "type"], "title": "BaseMessage", "type": "object"}}, "description": "The state of the agent.", "properties": {"messages": {"items": {"$ref": "#/$defs/BaseMessage"}, "title": "Messages", "type": "array"}, "remaining_steps": {"title": "Remaining Steps", "type": "integer"}}, "required": ["messages"], "title": "AgentState", "type": "object"}' + '{"$defs": {"BaseMessage": {"additionalProperties": true, "description": "Base abstract message class.\\n\\nMessages are the inputs and outputs of ChatModels.", "properties": {"content": {"anyOf": [{"type": "string"}, {"items": {"anyOf": [{"type": "string"}, {"additionalProperties": true, "type": "object"}]}, "type": "array"}], "title": "Content"}, "additional_kwargs": {"additionalProperties": true, "title": "Additional Kwargs", "type": "object"}, "response_metadata": {"additionalProperties": true, "title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "type": "string"}, "name": {"anyOf": [{"type": "string"}, {"type": "null"}], "default": null, "title": "Name"}, "id": {"anyOf": [{"type": "string"}, {"type": "null"}], "default": null, "title": "Id"}}, "required": ["content", "type"], "title": "BaseMessage", "type": "object"}}, "description": "The state of the agent.", "properties": {"messages": {"items": {"$ref": "#/$defs/BaseMessage"}, "title": "Messages", "type": "array"}, "remaining_steps": {"title": "Remaining Steps", "type": "integer"}, "structured_response": {"title": "Structured Response", "type": "null"}}, "required": ["messages"], "title": "AgentState", "type": "object"}' # --- # name: test_prebuilt_tool_chat.2 ''' @@ -198,7 +198,7 @@ } }, { - "id": "agent", + "id": "model", "type": "runnable", "data": { "id": [ @@ -207,7 +207,7 @@ "_runnable", "RunnableCallable" ], - "name": "agent" + "name": "model" } }, { @@ -230,21 +230,21 @@ "edges": [ { "source": "__start__", - "target": "agent" + "target": "model" }, { - "source": "agent", + "source": "model", "target": "__end__", "conditional": true }, { - "source": "agent", + "source": "model", "target": "tools", "conditional": true }, { "source": "tools", - "target": "agent" + "target": "model" } ] } @@ -253,10 +253,10 @@ # name: test_prebuilt_tool_chat.3 ''' graph TD; - __start__ --> agent; - agent -.-> __end__; - agent -.-> tools; - tools --> agent; + __start__ --> model; + model -.-> __end__; + model -.-> tools; + tools --> model; ''' # --- diff --git a/libs/langgraph/tests/conftest.py b/libs/langgraph/tests/conftest.py index 4ef30b918..465cdc9d5 100644 --- a/libs/langgraph/tests/conftest.py +++ b/libs/langgraph/tests/conftest.py @@ -69,7 +69,7 @@ def cache(request: pytest.FixtureRequest) -> Iterator[BaseCache]: elif request.param == "redis": # Get worker ID for parallel test isolation worker_id = getattr(request.config, "workerinput", {}).get("workerid", "master") - + redis_client = redis.Redis( host="localhost", port=6379, db=0, decode_responses=False ) diff --git a/libs/langgraph/tests/test_large_cases.py b/libs/langgraph/tests/test_large_cases.py index dbba69966..a7b5c36c7 100644 --- a/libs/langgraph/tests/test_large_cases.py +++ b/libs/langgraph/tests/test_large_cases.py @@ -1390,11 +1390,11 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None: ), { "langgraph_step": 1, - "langgraph_node": "agent", - "langgraph_triggers": ("branch:to:agent",), - "langgraph_path": (PULL, "agent"), - "langgraph_checkpoint_ns": AnyStr("agent:"), - "checkpoint_ns": AnyStr("agent:"), + "langgraph_node": "model", + "langgraph_triggers": ("branch:to:model",), + "langgraph_path": (PULL, "model"), + "langgraph_checkpoint_ns": AnyStr("model:"), + "checkpoint_ns": AnyStr("model:"), "ls_provider": "fakechatmodel", "ls_model_type": "chat", }, @@ -1449,11 +1449,11 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None: ), { "langgraph_step": 3, - "langgraph_node": "agent", - "langgraph_triggers": ("branch:to:agent",), - "langgraph_path": (PULL, "agent"), - "langgraph_checkpoint_ns": AnyStr("agent:"), - "checkpoint_ns": AnyStr("agent:"), + "langgraph_node": "model", + "langgraph_triggers": ("branch:to:model",), + "langgraph_path": (PULL, "model"), + "langgraph_checkpoint_ns": AnyStr("model:"), + "checkpoint_ns": AnyStr("model:"), "ls_provider": "fakechatmodel", "ls_model_type": "chat", }, @@ -1497,11 +1497,11 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None: ), { "langgraph_step": 5, - "langgraph_node": "agent", - "langgraph_triggers": ("branch:to:agent",), - "langgraph_path": (PULL, "agent"), - "langgraph_checkpoint_ns": AnyStr("agent:"), - "checkpoint_ns": AnyStr("agent:"), + "langgraph_node": "model", + "langgraph_triggers": ("branch:to:model",), + "langgraph_path": (PULL, "model"), + "langgraph_checkpoint_ns": AnyStr("model:"), + "checkpoint_ns": AnyStr("model:"), "ls_provider": "fakechatmodel", "ls_model_type": "chat", }, @@ -1533,7 +1533,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None: for output in (invoke_updates_events, stream_updates_events): assert output[:3] == [ { - "agent": { + "model": { "messages": [ _AnyIdAIMessage( content="", @@ -1560,7 +1560,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None: } }, { - "agent": { + "model": { "messages": [ _AnyIdAIMessage( content="", @@ -1606,7 +1606,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None: }, ) assert output[5:] == [ - {"agent": {"messages": [_AnyIdAIMessage(content="answer")]}} + {"model": {"messages": [_AnyIdAIMessage(content="answer")]}} ] diff --git a/libs/langgraph/tests/test_large_cases_async.py b/libs/langgraph/tests/test_large_cases_async.py index 6af84c0d6..274b2c711 100644 --- a/libs/langgraph/tests/test_large_cases_async.py +++ b/libs/langgraph/tests/test_large_cases_async.py @@ -1143,11 +1143,11 @@ async def test_prebuilt_tool_chat() -> None: ), { "langgraph_step": 1, - "langgraph_node": "agent", - "langgraph_triggers": ("branch:to:agent",), - "langgraph_path": (PULL, "agent"), - "langgraph_checkpoint_ns": AnyStr("agent:"), - "checkpoint_ns": AnyStr("agent:"), + "langgraph_node": "model", + "langgraph_triggers": ("branch:to:model",), + "langgraph_path": (PULL, "model"), + "langgraph_checkpoint_ns": AnyStr("model:"), + "checkpoint_ns": AnyStr("model:"), "ls_provider": "fakechatmodel", "ls_model_type": "chat", }, @@ -1202,11 +1202,11 @@ async def test_prebuilt_tool_chat() -> None: ), { "langgraph_step": 3, - "langgraph_node": "agent", - "langgraph_triggers": ("branch:to:agent",), - "langgraph_path": (PULL, "agent"), - "langgraph_checkpoint_ns": AnyStr("agent:"), - "checkpoint_ns": AnyStr("agent:"), + "langgraph_node": "model", + "langgraph_triggers": ("branch:to:model",), + "langgraph_path": (PULL, "model"), + "langgraph_checkpoint_ns": AnyStr("model:"), + "checkpoint_ns": AnyStr("model:"), "ls_provider": "fakechatmodel", "ls_model_type": "chat", }, @@ -1250,11 +1250,11 @@ async def test_prebuilt_tool_chat() -> None: ), { "langgraph_step": 5, - "langgraph_node": "agent", - "langgraph_triggers": ("branch:to:agent",), - "langgraph_path": (PULL, "agent"), - "langgraph_checkpoint_ns": AnyStr("agent:"), - "checkpoint_ns": AnyStr("agent:"), + "langgraph_node": "model", + "langgraph_triggers": ("branch:to:model",), + "langgraph_path": (PULL, "model"), + "langgraph_checkpoint_ns": AnyStr("model:"), + "checkpoint_ns": AnyStr("model:"), "ls_provider": "fakechatmodel", "ls_model_type": "chat", }, @@ -1269,7 +1269,7 @@ async def test_prebuilt_tool_chat() -> None: ] assert stream_updates_events[:3] == [ { - "agent": { + "model": { "messages": [ _AnyIdAIMessage( content="", @@ -1296,7 +1296,7 @@ async def test_prebuilt_tool_chat() -> None: } }, { - "agent": { + "model": { "messages": [ _AnyIdAIMessage( content="", @@ -1342,7 +1342,7 @@ async def test_prebuilt_tool_chat() -> None: }, ) assert stream_updates_events[5:] == [ - {"agent": {"messages": [_AnyIdAIMessage(content="answer")]}} + {"model": {"messages": [_AnyIdAIMessage(content="answer")]}} ] diff --git a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py index a8390ff80..e31d99fdd 100644 --- a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py @@ -8,7 +8,6 @@ from typing import ( Generic, Optional, Sequence, - Type, Union, cast, get_type_hints, @@ -21,7 +20,6 @@ from langchain_core.language_models import ( ) from langchain_core.messages import ( AIMessage, - AnyMessage, BaseMessage, SystemMessage, ToolMessage, @@ -52,51 +50,21 @@ from langgraph.prebuilt.tool_node import ToolNode from langgraph.runtime import Runtime from langgraph.store.base import BaseStore from langgraph.types import Checkpointer, Command, Send -from langgraph.typing import ContextT - -F = TypeVar("F", bound=Callable[..., Any]) +from langgraph.typing import ContextT, StateT StructuredResponseT = TypeVar( "StructuredResponseT", bound=Union[dict, BaseModel, None], default=None ) -# We create the AgentState that we will pass around -# This simply involves a list of messages -# We want steps to return messages to append to the list -# So we annotate the messages attribute with `add_messages` reducer -class AgentState(TypedDict): +class AgentState(TypedDict, Generic[StructuredResponseT]): """The state of the agent.""" messages: Annotated[Sequence[BaseMessage], add_messages] remaining_steps: NotRequired[RemainingSteps] - -class AgentStatePydantic(BaseModel): - """The state of the agent.""" - - messages: Annotated[Sequence[BaseMessage], add_messages] - - remaining_steps: RemainingSteps = 25 - - -class AgentStateWithStructuredResponse(AgentState, Generic[StructuredResponseT]): - """The state of the agent with a structured response.""" - - structured_response: StructuredResponseT - - -class AgentStateWithStructuredResponsePydantic( - AgentStatePydantic, Generic[StructuredResponseT] -): - """The state of the agent with a structured response.""" - - structured_response: StructuredResponseT - - -StateSchema = TypeVar("StateSchema", bound=Union[AgentState, AgentStatePydantic]) -StateSchemaType = Type[StateSchema] + structured_response: NotRequired[StructuredResponseT] PROMPT_RUNNABLE_NAME = "Prompt" @@ -104,12 +72,12 @@ PROMPT_RUNNABLE_NAME = "Prompt" Prompt = Union[ SystemMessage, str, - Callable[[StateSchema], LanguageModelInput], - Runnable[StateSchema, LanguageModelInput], + Callable[[StateT], LanguageModelInput], + Runnable[StateT, LanguageModelInput], ] -def _get_state_value(state: StateSchema, key: str, default: Any = None) -> Any: +def _get_state_value(state: StateT, key: str, default: Any = None) -> Any: return ( state.get(key, default) if isinstance(state, dict) @@ -184,15 +152,17 @@ def _validate_chat_history( raise ValueError(error_message) -class _AgentBuilder(Generic[StructuredResponseT]): +class _AgentBuilder(Generic[StateT, ContextT, StructuredResponseT]): """Internal builder class for constructing and agent.""" + _final_state_schema: type[StateT] + def __init__( self, model: Union[ str, BaseChatModel, - SyncOrAsync[[StateSchema, Runtime[ContextT]], BaseModel], + SyncOrAsync[[StateT, Runtime[ContextT]], BaseChatModel], ], tools: Union[Sequence[Union[BaseTool, Callable, dict[str, Any]]], ToolNode], *, @@ -200,8 +170,8 @@ class _AgentBuilder(Generic[StructuredResponseT]): response_format: Optional[ResponseFormat[StructuredResponseT]] = None, pre_model_hook: Optional[RunnableLike] = None, post_model_hook: Optional[RunnableLike] = None, - state_schema: Optional[StateSchemaType] = None, - context_schema: Optional[Type[Any]] = None, + state_schema: Optional[type[StateT]] = None, + context_schema: Optional[type[ContextT]] = None, name: Optional[str] = None, store: Optional[BaseStore] = None, ): @@ -292,11 +262,7 @@ class _AgentBuilder(Generic[StructuredResponseT]): self._final_state_schema = self.state_schema else: - self._final_state_schema = ( - AgentStateWithStructuredResponse - if self.response_format is not None - else AgentState - ) + self._final_state_schema = cast(type[StateT], AgentState) def _handle_structured_response_tool_calls( self, response: AIMessage @@ -404,7 +370,7 @@ class _AgentBuilder(Generic[StructuredResponseT]): self._static_model = None def _resolve_model( - self, state: StateSchema, runtime: Runtime[ContextT] + self, state: StateT, runtime: Runtime[ContextT] ) -> LanguageModelLike: """Resolve the model to use, handling both static and dynamic models.""" if self._is_dynamic_model: @@ -413,12 +379,12 @@ class _AgentBuilder(Generic[StructuredResponseT]): return self._static_model async def _aresolve_model( - self, state: StateSchema, runtime: Runtime[ContextT] + self, state: StateT, runtime: Runtime[ContextT] ) -> LanguageModelLike: """Async resolve the model to use, handling both static and dynamic models.""" if self._is_async_dynamic_model: dynamic_model = cast( - Callable[[StateSchema, Runtime[ContextT]], Awaitable[BaseChatModel]], + Callable[[StateT, Runtime[ContextT]], Awaitable[BaseChatModel]], self.model, ) resolved_model = await dynamic_model(state, runtime) @@ -431,21 +397,11 @@ class _AgentBuilder(Generic[StructuredResponseT]): def create_model_node(self) -> RunnableCallable: """Create the 'agent' node that calls the LLM.""" - def _get_model_input_state(state: StateSchema) -> StateSchema: - if self.pre_model_hook is not None: - messages = _get_state_value( - state, "llm_input_messages" - ) or _get_state_value(state, "messages") - error_msg = ( - f"Expected input to call_model to have 'llm_input_messages' " - f"or 'messages' key, but got {state}" - ) - else: - messages = _get_state_value(state, "messages") - error_msg = ( - f"Expected input to call_model to " - f"have 'messages' key, but got {state}" - ) + def _get_model_input_state(state: StateT) -> StateT: + messages = _get_state_value(state, "messages") + error_msg = ( + f"Expected input to call_model to have 'messages' key, but got {state}" + ) if messages is None: raise ValueError(error_msg) @@ -462,7 +418,7 @@ class _AgentBuilder(Generic[StructuredResponseT]): state["messages"] = messages # type: ignore return state - def _are_more_steps_needed(state: StateSchema, response: BaseMessage) -> bool: + def _are_more_steps_needed(state: StateT, response: BaseMessage) -> bool: has_tool_calls = isinstance(response, AIMessage) and response.tool_calls all_tools_return_direct = ( all( @@ -481,7 +437,7 @@ class _AgentBuilder(Generic[StructuredResponseT]): return False def call_model( - state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig + state: StateT, runtime: Runtime[ContextT], config: RunnableConfig ) -> dict[str, Any] | Command: """Call the model with the current state and return the response.""" if self._is_async_dynamic_model: @@ -519,7 +475,7 @@ class _AgentBuilder(Generic[StructuredResponseT]): return {"messages": [response]} async def acall_model( - state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig + state: StateT, runtime: Runtime[ContextT], config: RunnableConfig ) -> dict[str, Any] | Command: """Call the model with the current state and return the response.""" model_input = _get_model_input_state(state) @@ -555,32 +511,10 @@ class _AgentBuilder(Generic[StructuredResponseT]): return RunnableCallable(call_model, acall_model) - def _get_input_schema(self) -> StateSchemaType: - """Get input schema for model node.""" - if self.pre_model_hook is not None: - if isinstance(self._final_state_schema, type) and issubclass( - self._final_state_schema, BaseModel - ): - from pydantic import create_model - - return create_model( - "CallModelInputSchema", - llm_input_messages=(list[AnyMessage], ...), - __base__=self._final_state_schema, - ) - else: - - class CallModelInputSchema(self._final_state_schema): # type: ignore - llm_input_messages: list[AnyMessage] - - return CallModelInputSchema - else: - return self._final_state_schema - - def create_model_router(self) -> Callable[[StateSchema], Union[str, list[Send]]]: + def create_model_router(self) -> Callable[[StateT], Union[str, list[Send]]]: """Create routing function for model node conditional edges.""" - def should_continue(state: StateSchema) -> Union[str, list[Send]]: + def should_continue(state: StateT) -> Union[str, list[Send]]: messages = _get_state_value(state, "messages") last_message = messages[-1] @@ -615,10 +549,10 @@ class _AgentBuilder(Generic[StructuredResponseT]): def create_post_model_hook_router( self, - ) -> Callable[[StateSchema], Union[str, list[Send]]]: + ) -> Callable[[StateT], Union[str, list[Send]]]: """Create a routing function for post_model_hook node conditional edges.""" - def post_model_hook_router(state: StateSchema) -> Union[str, list[Send]]: + def post_model_hook_router(state: StateT) -> Union[str, list[Send]]: messages = _get_state_value(state, "messages") # Check if the last message is a ToolMessage from a structured tool. @@ -655,12 +589,12 @@ class _AgentBuilder(Generic[StructuredResponseT]): return post_model_hook_router - def create_tools_router(self) -> Optional[Callable[[StateSchema], str]]: + def create_tools_router(self) -> Optional[Callable[[StateT], str]]: """Create a routing function for tools node conditional edges.""" if not self._should_return_direct: return None - def route_tool_responses(state: StateSchema) -> str: + def route_tool_responses(state: StateT) -> str: messages = _get_state_value(state, "messages") for m in reversed(messages): if not isinstance(m, ToolMessage): @@ -713,9 +647,7 @@ class _AgentBuilder(Generic[StructuredResponseT]): workflow.set_entry_point(self._get_entry_point()) # Add nodes - workflow.add_node( - "model", self.create_model_node(), input_schema=self._get_input_schema() - ) + workflow.add_node("model", self.create_model_node()) if self._tool_calling_enabled: workflow.add_node("tools", self._tool_node) @@ -775,7 +707,7 @@ def create_agent( model: Union[ str, BaseChatModel, - SyncOrAsync[[StateSchema, Runtime[ContextT]], BaseModel], + SyncOrAsync[[StateT, Runtime[ContextT]], BaseChatModel], ], tools: Union[Sequence[Union[BaseTool, Callable, dict[str, Any]]], ToolNode], *, @@ -785,8 +717,8 @@ def create_agent( ] = None, pre_model_hook: Optional[RunnableLike] = None, post_model_hook: Optional[RunnableLike] = None, - state_schema: Optional[StateSchemaType] = None, - context_schema: Optional[Type[Any]] = None, + state_schema: Optional[type[StateT]] = None, + context_schema: Optional[type[ContextT]] = None, checkpointer: Optional[Checkpointer] = None, store: Optional[BaseStore] = None, interrupt_before: Optional[list[str]] = None, @@ -874,24 +806,17 @@ def create_agent( Useful for managing long message histories (e.g., message trimming, summarization, etc.). Pre-model hook must be a callable or a runnable that takes in current graph state and returns a state update in the form of ```python - # At least one of `messages` or `llm_input_messages` MUST be provided + # Where `messages` MUST be provided { - # If provided, will UPDATE the `messages` in the state + # will UPDATE the `messages` in the state "messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), ...], - # If provided, will be used as the input to the LLM, - # and will NOT UPDATE `messages` in the state - "llm_input_messages": [...], # Any other state keys that need to be propagated ... } ``` - !!! Important - At least one of `messages` or `llm_input_messages` MUST be provided and will be used as an input to the `agent` node. - The rest of the keys will be added to the graph state. - !!! Warning - If you are returning `messages` in the pre-model hook, you should OVERWRITE the `messages` key by doing the following: + you should OVERWRITE the `messages` key by doing the following: ```python { @@ -980,7 +905,7 @@ def create_agent( response_format = cast(Optional[ToolOutput], response_format) # Create and configure the agent builder - builder = _AgentBuilder( + builder = _AgentBuilder[StateT, ContextT, StructuredResponseT]( model=model, tools=tools, prompt=prompt, @@ -1008,7 +933,4 @@ def create_agent( __all__ = [ "create_agent", "AgentState", - "AgentStatePydantic", - "AgentStateWithStructuredResponse", - "AgentStateWithStructuredResponsePydantic", ] diff --git a/libs/prebuilt/tests/test_react_agent.py b/libs/prebuilt/tests/test_react_agent.py index 7dd787ff0..9674e36f4 100644 --- a/libs/prebuilt/tests/test_react_agent.py +++ b/libs/prebuilt/tests/test_react_agent.py @@ -2,7 +2,6 @@ import dataclasses import inspect from typing import ( Annotated, - Optional, Union, ) @@ -32,8 +31,6 @@ from langgraph.prebuilt import ( ) from langgraph.prebuilt.chat_agent_executor import ( AgentState, - AgentStatePydantic, - StateSchemaType, _validate_chat_history, ) from langgraph.prebuilt.tool_node import ( @@ -466,14 +463,8 @@ class CustomState(AgentState): user_name: str -class CustomStatePydantic(AgentStatePydantic): - user_name: Optional[str] = None - - -@pytest.mark.parametrize("state_schema", [CustomState, CustomStatePydantic]) def test_react_agent_update_state( sync_checkpointer: BaseCheckpointSaver, - state_schema: StateSchemaType, ) -> None: @dec_tool def get_user_name(tool_call_id: Annotated[str, InjectedToolCallId]): @@ -490,31 +481,20 @@ def test_react_agent_update_state( } ) - if issubclass(state_schema, AgentStatePydantic): + def prompt(state: CustomState): + user_name = state.get("user_name") + if user_name is None: + return state["messages"] - def prompt(state: CustomStatePydantic): - user_name = state.user_name - if user_name is None: - return state.messages - - system_msg = f"User name is {user_name}" - return [{"role": "system", "content": system_msg}] + state.messages - else: - - def prompt(state: CustomState): - user_name = state.get("user_name") - if user_name is None: - return state["messages"] - - system_msg = f"User name is {user_name}" - return [{"role": "system", "content": system_msg}] + state["messages"] + system_msg = f"User name is {user_name}" + return [{"role": "system", "content": system_msg}] + state["messages"] tool_calls = [[{"args": {}, "id": "1", "name": "get_user_name"}]] model = FakeToolCallingModel(tool_calls=tool_calls) agent = create_agent( model, [get_user_name], - state_schema=state_schema, + state_schema=CustomState, prompt=prompt, checkpointer=sync_checkpointer, ) @@ -604,39 +584,20 @@ class AgentStateExtraKey(AgentState): foo: int -class AgentStateExtraKeyPydantic(AgentStatePydantic): - foo: int - - -@pytest.mark.parametrize( - "state_schema", [AgentStateExtraKey, AgentStateExtraKeyPydantic] -) -def test_create_react_agent_inject_vars(state_schema: StateSchemaType) -> None: +def test_create_react_agent_inject_vars() -> None: """Test that the agent can inject state and store into tool functions.""" store = InMemoryStore() namespace = ("test",) store.put(namespace, "test_key", {"bar": 3}) - if issubclass(state_schema, AgentStatePydantic): - - def tool1( - some_val: int, - state: Annotated[AgentStateExtraKeyPydantic, InjectedState], - store: Annotated[BaseStore, InjectedStore()], - ) -> str: - """Tool 1 docstring.""" - store_val = store.get(namespace, "test_key").value["bar"] - return some_val + state.foo + store_val - else: - - def tool1( - some_val: int, - state: Annotated[dict, InjectedState], - store: Annotated[BaseStore, InjectedStore()], - ) -> str: - """Tool 1 docstring.""" - store_val = store.get(namespace, "test_key").value["bar"] - return some_val + state["foo"] + store_val + def tool1( + some_val: int, + state: Annotated[dict, InjectedState], + store: Annotated[BaseStore, InjectedStore()], + ) -> str: + """Tool 1 docstring.""" + store_val = store.get(namespace, "test_key").value["bar"] + return some_val + state["foo"] + store_val tool_call = { "name": "tool1", @@ -648,7 +609,7 @@ def test_create_react_agent_inject_vars(state_schema: StateSchemaType) -> None: agent = create_agent( model, ToolNode([tool1], handle_tool_errors=False), - state_schema=state_schema, + state_schema=AgentStateExtraKey, store=store, ) result = agent.invoke({"messages": [{"role": "user", "content": "hi"}], "foo": 2}) @@ -1474,6 +1435,7 @@ async def test_dynamic_model_receives_correct_state_async(): assert received_state["messages"][0].content == "hello async" +@pytest.mark.skip(reason="TODO: support with prepare call") def test_pre_model_hook() -> None: model = FakeToolCallingModel(tool_calls=[]) @@ -1658,36 +1620,19 @@ def test_post_model_hook_with_structured_output() -> None: ] -@pytest.mark.parametrize( - "state_schema", [AgentStateExtraKey, AgentStateExtraKeyPydantic] -) -def test_create_react_agent_inject_vars_with_post_model_hook( - state_schema: StateSchemaType, -) -> None: +def test_create_react_agent_inject_vars_with_post_model_hook() -> None: store = InMemoryStore() namespace = ("test",) store.put(namespace, "test_key", {"bar": 3}) - if issubclass(state_schema, AgentStatePydantic): - - def tool1( - some_val: int, - state: Annotated[AgentStateExtraKeyPydantic, InjectedState], - store: Annotated[BaseStore, InjectedStore()], - ) -> str: - """Tool 1 docstring.""" - store_val = store.get(namespace, "test_key").value["bar"] - return some_val + state.foo + store_val - else: - - def tool1( - some_val: int, - state: Annotated[dict, InjectedState], - store: Annotated[BaseStore, InjectedStore()], - ) -> str: - """Tool 1 docstring.""" - store_val = store.get(namespace, "test_key").value["bar"] - return some_val + state["foo"] + store_val + def tool1( + some_val: int, + state: Annotated[dict, InjectedState], + store: Annotated[BaseStore, InjectedStore()], + ) -> str: + """Tool 1 docstring.""" + store_val = store.get(namespace, "test_key").value["bar"] + return some_val + state["foo"] + store_val tool_call = { "name": "tool1", @@ -1704,7 +1649,7 @@ def test_create_react_agent_inject_vars_with_post_model_hook( agent = create_agent( model, ToolNode([tool1], handle_tool_errors=False), - state_schema=state_schema, + state_schema=AgentStateExtraKey, store=store, post_model_hook=post_model_hook, )