From 67e8f8fc115daa108e07592ed82559c5c13c2073 Mon Sep 17 00:00:00 2001 From: Vadym Barda Date: Wed, 2 Apr 2025 11:41:31 -0400 Subject: [PATCH] prebuilt: add optional pre-model hook that runs before calling LLM in create_react_agent (#4059) Example: ```python from typing import Any from langchain_openai import ChatOpenAI from langchain_community.tools.tavily_search import TavilySearchResults from langchain_core.messages import AnyMessage from langchain_core.messages.utils import count_tokens_approximately from langgraph.graph import MessagesState from langgraph.prebuilt.chat_agent_executor import create_react_agent, AgentState from langgraph.checkpoint.memory import InMemorySaver from langmem.short_term import SummarizationNode, RunningSummary class State(MessagesState): context: dict[str, Any] search = TavilySearchResults(max_results=3) tools = [search] model = ChatOpenAI(model="gpt-4o") summarization_model = model.bind(max_tokens=256) summarization_node = SummarizationNode( token_counter=count_tokens_approximately, model=summarization_model, max_tokens=2048, max_summary_tokens=256, output_messages_key="messages" # output_messages_key="llm_input_messages" ) checkpointer = InMemorySaver() class State(AgentState): user_language: str # summarization-related keys context: dict[str, Any] def prompt(state): language = state["user_language"] system_msg = f"Always respond in {language}" return [{"role": "system", "content": system_msg}] + state["messages"] graph = create_react_agent( model, tools, prompt=prompt, pre_model_hook=summarization_node, state_schema=State, checkpointer=checkpointer ) ``` --- .../langgraph/prebuilt/chat_agent_executor.py | 133 ++++++++++++++++-- libs/prebuilt/tests/test_react_agent.py | 35 +++++ 2 files changed, 154 insertions(+), 14 deletions(-) diff --git a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py index 0ca0d42b1..667e9961a 100644 --- a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py @@ -18,7 +18,13 @@ from langchain_core.language_models import ( LanguageModelInput, LanguageModelLike, ) -from langchain_core.messages import AIMessage, BaseMessage, SystemMessage, ToolMessage +from langchain_core.messages import ( + AIMessage, + AnyMessage, + BaseMessage, + SystemMessage, + ToolMessage, +) from langchain_core.runnables import ( Runnable, RunnableBinding, @@ -37,7 +43,7 @@ from langgraph.managed import IsLastStep, RemainingSteps from langgraph.prebuilt.tool_node import ToolNode from langgraph.store.base import BaseStore from langgraph.types import Checkpointer, Send -from langgraph.utils.runnable import RunnableCallable +from langgraph.utils.runnable import RunnableCallable, RunnableLike StructuredResponse = Union[dict, BaseModel] StructuredResponseSchema = Union[dict, type[BaseModel]] @@ -263,6 +269,7 @@ def create_react_agent( response_format: Optional[ Union[StructuredResponseSchema, tuple[str, StructuredResponseSchema]] ] = None, + pre_model_hook: Optional[RunnableLike] = None, state_schema: Optional[StateSchemaType] = None, config_schema: Optional[Type[Any]] = None, checkpointer: Optional[Checkpointer] = None, @@ -305,6 +312,36 @@ def create_react_agent( !!! Note 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/). + + pre_model_hook: An optional node to add before the `agent` node (i.e., the node that calls the LLM). + 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 + { + # If provided, 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: + + ```python + { + "messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), *new_messages] + ... + } + ``` state_schema: An optional state schema that defines graph state. Must have `messages` and `remaining_steps` keys. Defaults to `AgentState` that defines those two keys. @@ -678,10 +715,33 @@ def create_react_agent( or (remaining_steps is not None and remaining_steps < 2 and has_tool_calls) ) + def _get_model_input_state(state: StateSchema) -> StateSchema: + if 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' or 'messages' key, but got {state}" + else: + 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) + + _validate_chat_history(messages) + # we're passing messages under `messages` key, as this is expected by the prompt + if isinstance(state_schema, type) and issubclass(state_schema, BaseModel): + state.messages = messages # type: ignore + else: + state["messages"] = messages # type: ignore + + return state + # 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) + state = _get_model_input_state(state) response = cast(AIMessage, model_runnable.invoke(state, config)) # add agent name to the AIMessage response.name = name @@ -699,8 +759,7 @@ def create_react_agent( return {"messages": [response]} async def acall_model(state: StateSchema, config: RunnableConfig) -> StateSchema: - messages = _get_state_value(state, "messages") - _validate_chat_history(messages) + state = _get_model_input_state(state) response = cast(AIMessage, await model_runnable.ainvoke(state, config)) # add agent name to the AIMessage response.name = name @@ -716,6 +775,27 @@ def create_react_agent( # We return a list, because this will get added to the existing list return {"messages": [response]} + input_schema: StateSchemaType + if pre_model_hook is not None: + # Dynamically create a schema that inherits from state_schema and adds 'llm_input_messages' + if isinstance(state_schema, type) and issubclass(state_schema, BaseModel): + # For Pydantic schemas + from pydantic import create_model + + input_schema = create_model( + "CallModelInputSchema", + llm_input_messages=(list[AnyMessage], ...), + __base__=state_schema, + ) + else: + # For TypedDict schemas + class CallModelInputSchema(state_schema): # type: ignore + llm_input_messages: list[AnyMessage] + + input_schema = CallModelInputSchema + else: + input_schema = state_schema + def generate_structured_response( state: StateSchema, config: RunnableConfig ) -> StateSchema: @@ -749,8 +829,20 @@ def create_react_agent( if not tool_calling_enabled: # Define a new graph workflow = StateGraph(state_schema, config_schema=config_schema) - workflow.add_node("agent", RunnableCallable(call_model, acall_model)) - workflow.set_entry_point("agent") + workflow.add_node( + "agent", + RunnableCallable(call_model, acall_model), + input=input_schema, + ) + if pre_model_hook is not None: + workflow.add_node("pre_model_hook", pre_model_hook) + workflow.add_edge("pre_model_hook", "agent") + entrypoint = "pre_model_hook" + else: + entrypoint = "agent" + + workflow.set_entry_point(entrypoint) + if response_format is not None: workflow.add_node( "generate_structured_response", @@ -791,12 +883,23 @@ def create_react_agent( workflow = StateGraph(state_schema or AgentState, config_schema=config_schema) # Define the two nodes we will cycle between - workflow.add_node("agent", RunnableCallable(call_model, acall_model)) + workflow.add_node( + "agent", RunnableCallable(call_model, acall_model), input=input_schema + ) workflow.add_node("tools", tool_node) + # Optionally add a pre-model hook node that will be called + # every time before the "agent" (LLM-calling node) + if pre_model_hook is not None: + workflow.add_node("pre_model_hook", pre_model_hook) + workflow.add_edge("pre_model_hook", "agent") + entrypoint = "pre_model_hook" + else: + entrypoint = "agent" + # Set the entrypoint as `agent` # This means that this node is the first one called - workflow.set_entry_point("agent") + workflow.set_entry_point(entrypoint) # Add a structured output node if response_format is provided if response_format is not None: @@ -821,18 +924,20 @@ def create_react_agent( path_map=should_continue_destinations, ) - def route_tool_responses(state: StateSchema) -> Literal["agent", "__end__"]: + def route_tool_responses(state: StateSchema) -> str: for m in reversed(_get_state_value(state, "messages")): if not isinstance(m, ToolMessage): break if m.name in should_return_direct: return END - return "agent" + return entrypoint if should_return_direct: - workflow.add_conditional_edges("tools", route_tool_responses) + workflow.add_conditional_edges( + "tools", route_tool_responses, path_map=[entrypoint, END] + ) else: - workflow.add_edge("tools", "agent") + workflow.add_edge("tools", entrypoint) # Finally, we compile it! # This compiles it into a LangChain Runnable, diff --git a/libs/prebuilt/tests/test_react_agent.py b/libs/prebuilt/tests/test_react_agent.py index f816875ad..4b469d077 100644 --- a/libs/prebuilt/tests/test_react_agent.py +++ b/libs/prebuilt/tests/test_react_agent.py @@ -16,6 +16,7 @@ from langchain_core.messages import ( AIMessage, AnyMessage, HumanMessage, + RemoveMessage, SystemMessage, ToolCall, ToolMessage, @@ -29,6 +30,7 @@ from typing_extensions import TypedDict from langgraph.checkpoint.base import BaseCheckpointSaver from langgraph.graph import START, MessagesState, StateGraph, add_messages +from langgraph.graph.message import REMOVE_ALL_MESSAGES from langgraph.prebuilt import ( ToolNode, create_react_agent, @@ -1432,3 +1434,36 @@ def test_get_model() -> None: with pytest.raises(TypeError): _get_model(RunnableLambda(lambda message: message)) + + +def test_pre_model_hook() -> None: + model = FakeToolCallingModel(tool_calls=[]) + + # Test `llm_input_messages` + def pre_model_hook(state: AgentState): + return {"llm_input_messages": [HumanMessage("Hello!")]} + + agent = create_react_agent(model, [], pre_model_hook=pre_model_hook) + assert "pre_model_hook" in agent.nodes + result = agent.invoke({"messages": [HumanMessage("hi?")]}) + assert result == { + "messages": [ + _AnyIdHumanMessage(content="hi?"), + AIMessage(content="Hello!", id="0"), + ] + } + + # Test `messages` + def pre_model_hook(state: AgentState): + return { + "messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), HumanMessage("Hello!")] + } + + agent = create_react_agent(model, [], pre_model_hook=pre_model_hook) + result = agent.invoke({"messages": [HumanMessage("hi?")]}) + assert result == { + "messages": [ + _AnyIdHumanMessage(content="Hello!"), + AIMessage(content="Hello!", id="1"), + ] + }