diff --git a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py index 117a2bc91..ceed3ef47 100644 --- a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py @@ -10,6 +10,7 @@ from typing import ( TypeVar, Union, cast, + get_type_hints, ) from langchain_core.language_models import ( @@ -57,13 +58,27 @@ class AgentState(TypedDict): remaining_steps: RemainingSteps +class AgentStatePydantic(BaseModel): + """The state of the agent.""" + + messages: Annotated[Sequence[BaseMessage], add_messages] + + remaining_steps: RemainingSteps = 25 + + class AgentStateWithStructuredResponse(AgentState): """The state of the agent with a structured response.""" structured_response: StructuredResponse -StateSchema = TypeVar("StateSchema", bound=AgentState) +class AgentStateWithStructuredResponsePydantic(AgentStatePydantic): + """The state of the agent with a structured response.""" + + structured_response: StructuredResponse + + +StateSchema = TypeVar("StateSchema", bound=Union[AgentState, AgentStatePydantic]) StateSchemaType = Type[StateSchema] PROMPT_RUNNABLE_NAME = "Prompt" @@ -76,21 +91,29 @@ Prompt = Union[ ] +def _get_state_value(state: StateSchema, key: str, default: Any = None) -> Any: + return ( + state.get(key, default) + if isinstance(state, dict) + else getattr(state, key, default) + ) + + def _get_prompt_runnable(prompt: Optional[Prompt]) -> Runnable: prompt_runnable: Runnable if prompt is None: prompt_runnable = RunnableCallable( - lambda state: state["messages"], name=PROMPT_RUNNABLE_NAME + lambda state: _get_state_value(state, "messages"), name=PROMPT_RUNNABLE_NAME ) elif isinstance(prompt, str): _system_message: BaseMessage = SystemMessage(content=prompt) prompt_runnable = RunnableCallable( - lambda state: [_system_message] + state["messages"], + lambda state: [_system_message] + _get_state_value(state, "messages"), name=PROMPT_RUNNABLE_NAME, ) elif isinstance(prompt, SystemMessage): prompt_runnable = RunnableCallable( - lambda state: [prompt] + state["messages"], + lambda state: [prompt] + _get_state_value(state, "messages"), name=PROMPT_RUNNABLE_NAME, ) elif inspect.iscoroutinefunction(prompt): @@ -283,7 +306,7 @@ def create_react_agent( The graph will make a separate call to the LLM to generate the structured response after the agent loop is finished. This is not the only strategy to get structured responses, see more options in [this guide](https://langchain-ai.github.io/langgraph/how-tos/react-agent-structured-output/). state_schema: An optional state schema that defines graph state. - Must have `messages` and `is_last_step` keys. + Must have `messages` and `remaining_steps` keys. Defaults to `AgentState` that defines those two keys. config_schema: An optional schema for configuration. Use this to expose configurable parameters via agent.config_specs. @@ -595,7 +618,8 @@ def create_react_agent( if response_format is not None: required_keys.add("structured_response") - if missing_keys := required_keys - set(state_schema.__annotations__): + schema_keys = set(get_type_hints(state_schema)) + if missing_keys := required_keys - set(schema_keys): raise ValueError(f"Missing required key(s) {missing_keys} in state_schema") if state_schema is None: @@ -636,35 +660,34 @@ def create_react_agent( # our graph needs to check if these were called should_return_direct = {t.name for t in tool_classes if t.return_direct} - # Define the function that calls the model - def call_model(state: AgentState, config: RunnableConfig) -> AgentState: - _validate_chat_history(state["messages"]) - response = cast(AIMessage, model_runnable.invoke(state, config)) - # add agent name to the AIMessage - response.name = name + def _are_more_steps_needed(state: StateSchema, response: BaseMessage) -> bool: has_tool_calls = isinstance(response, AIMessage) and response.tool_calls all_tools_return_direct = ( all(call["name"] in should_return_direct for call in response.tool_calls) if isinstance(response, AIMessage) else False ) - if ( - ( - "remaining_steps" not in state - and state.get("is_last_step", False) - and has_tool_calls - ) + remaining_steps = _get_state_value(state, "remaining_steps", None) + is_last_step = _get_state_value(state, "is_last_step", False) + return ( + (remaining_steps is None and is_last_step and has_tool_calls) or ( - "remaining_steps" in state - and state["remaining_steps"] < 1 + remaining_steps is not None + and remaining_steps < 1 and all_tools_return_direct ) - or ( - "remaining_steps" in state - and state["remaining_steps"] < 2 - and has_tool_calls - ) - ): + or (remaining_steps is not None and remaining_steps < 2 and has_tool_calls) + ) + + # Define the function that calls the model + def call_model(state: StateSchema, config: RunnableConfig) -> StateSchema: + messages = _get_state_value(state, "messages") + _validate_chat_history(messages) + response = cast(AIMessage, model_runnable.invoke(state, config)) + # add agent name to the AIMessage + response.name = name + + if _are_more_steps_needed(state, response): return { "messages": [ AIMessage( @@ -676,34 +699,13 @@ def create_react_agent( # We return a list, because this will get added to the existing list return {"messages": [response]} - async def acall_model(state: AgentState, config: RunnableConfig) -> AgentState: - _validate_chat_history(state["messages"]) + async def acall_model(state: StateSchema, config: RunnableConfig) -> StateSchema: + messages = _get_state_value(state, "messages") + _validate_chat_history(messages) response = cast(AIMessage, await model_runnable.ainvoke(state, config)) # add agent name to the AIMessage response.name = name - has_tool_calls = isinstance(response, AIMessage) and response.tool_calls - all_tools_return_direct = ( - all(call["name"] in should_return_direct for call in response.tool_calls) - if isinstance(response, AIMessage) - else False - ) - if ( - ( - "remaining_steps" not in state - and state.get("is_last_step", False) - and has_tool_calls - ) - or ( - "remaining_steps" in state - and state["remaining_steps"] < 1 - and all_tools_return_direct - ) - or ( - "remaining_steps" in state - and state["remaining_steps"] < 2 - and has_tool_calls - ) - ): + if _are_more_steps_needed(state, response): return { "messages": [ AIMessage( @@ -716,11 +718,11 @@ def create_react_agent( return {"messages": [response]} def generate_structured_response( - state: AgentState, config: RunnableConfig - ) -> AgentState: + state: StateSchema, config: RunnableConfig + ) -> StateSchema: # NOTE: we exclude the last message because there is enough information # for the LLM to generate the structured response - messages = state["messages"][:-1] + messages = _get_state_value(state, "messages")[:-1] structured_response_schema = response_format if isinstance(response_format, tuple): system_prompt, structured_response_schema = response_format @@ -733,11 +735,11 @@ def create_react_agent( return {"structured_response": response} async def agenerate_structured_response( - state: AgentState, config: RunnableConfig - ) -> AgentState: + state: StateSchema, config: RunnableConfig + ) -> StateSchema: # NOTE: we exclude the last message because there is enough information # for the LLM to generate the structured response - messages = state["messages"][:-1] + messages = _get_state_value(state, "messages")[:-1] structured_response_schema = response_format if isinstance(response_format, tuple): system_prompt, structured_response_schema = response_format @@ -773,8 +775,8 @@ def create_react_agent( ) # Define the function that determines whether to continue or not - def should_continue(state: AgentState) -> Union[str, list]: - messages = state["messages"] + def should_continue(state: StateSchema) -> Union[str, list]: + messages = _get_state_value(state, "messages") last_message = messages[-1] # If there is no function call, then we finish if not isinstance(last_message, AIMessage) or not last_message.tool_calls: @@ -824,8 +826,8 @@ def create_react_agent( path_map=should_continue_destinations, ) - def route_tool_responses(state: AgentState) -> Literal["agent", "__end__"]: - for m in reversed(state["messages"]): + def route_tool_responses(state: StateSchema) -> Literal["agent", "__end__"]: + for m in reversed(_get_state_value(state, "messages")): if not isinstance(m, ToolMessage): break if m.name in should_return_direct: diff --git a/libs/prebuilt/tests/test_react_agent.py b/libs/prebuilt/tests/test_react_agent.py index d5f405384..9d8b42e67 100644 --- a/libs/prebuilt/tests/test_react_agent.py +++ b/libs/prebuilt/tests/test_react_agent.py @@ -5,6 +5,7 @@ from functools import partial from typing import ( Annotated, List, + Optional, Type, TypeVar, Union, @@ -35,6 +36,8 @@ from langgraph.prebuilt import ( ) from langgraph.prebuilt.chat_agent_executor import ( AgentState, + AgentStatePydantic, + StateSchemaType, _get_model, _should_bind_tools, _validate_chat_history, @@ -528,22 +531,31 @@ def test_react_agent_with_structured_response(version: str) -> None: assert response["messages"][-2].content == "The weather is sunny and 75°F." +class CustomState(AgentState): + user_name: str + + +class CustomStatePydantic(AgentStatePydantic): + user_name: Optional[str] = None + + @pytest.mark.skipif( not IS_LANGCHAIN_CORE_030_OR_GREATER, reason="Langchain core 0.3.0 or greater is required", ) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) +@pytest.mark.parametrize("state_schema", [CustomState, CustomStatePydantic]) def test_react_agent_update_state( - request: pytest.FixtureRequest, checkpointer_name: str, version: str + request: pytest.FixtureRequest, + checkpointer_name: str, + version: str, + state_schema: StateSchemaType, ) -> None: checkpointer: BaseCheckpointSaver = request.getfixturevalue( "checkpointer_" + checkpointer_name ) - class State(AgentState): - user_name: str - @dec_tool def get_user_name(tool_call_id: Annotated[str, InjectedToolCallId]): """Retrieve user name""" @@ -559,20 +571,31 @@ def test_react_agent_update_state( } ) - def prompt(state: State): - user_name = state.get("user_name") - if user_name is None: - return state["messages"] + if issubclass(state_schema, AgentStatePydantic): - system_msg = f"User name is {user_name}" - return [{"role": "system", "content": system_msg}] + 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"] tool_calls = [[{"args": {}, "id": "1", "name": "get_user_name"}]] model = FakeToolCallingModel(tool_calls=tool_calls) agent = create_react_agent( model, [get_user_name], - state_schema=State, + state_schema=state_schema, prompt=prompt, checkpointer=checkpointer, version=version, @@ -802,23 +825,45 @@ def test_tool_node_inject_state(schema_: Type[T]) -> None: assert tool_message.content == "hi?" -@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) -def test_create_react_agent_inject_vars(version: str) -> None: - class AgentStateExtraKey(AgentState): - foo: int +class AgentStateExtraKey(AgentState): + foo: int + +class AgentStateExtraKeyPydantic(AgentStatePydantic): + foo: int + + +@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) +@pytest.mark.parametrize( + "state_schema", [AgentStateExtraKey, AgentStateExtraKeyPydantic] +) +def test_create_react_agent_inject_vars( + version: str, state_schema: StateSchemaType +) -> None: store = InMemoryStore() namespace = ("test",) store.put(namespace, "test_key", {"bar": 3}) - 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 + 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 tool_call = { "name": "tool1", @@ -830,7 +875,7 @@ def test_create_react_agent_inject_vars(version: str) -> None: agent = create_react_agent( model, [tool1], - state_schema=AgentStateExtraKey, + state_schema=state_schema, store=store, version=version, )