From 2fed0e485241a4aaeb185926cedf97703586b8a4 Mon Sep 17 00:00:00 2001 From: Eugene Yurtsev Date: Tue, 12 Aug 2025 17:28:05 -0400 Subject: [PATCH] Internal refactor of create react-agent --- .../langgraph/prebuilt/chat_agent_executor.py | 969 ++++++++++-------- libs/prebuilt/tests/test_react_agent.py | 2 +- libs/prebuilt/tests/test_react_agent_graph.py | 12 +- 3 files changed, 537 insertions(+), 446 deletions(-) diff --git a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py index 8256bbc44..bae31c6ca 100644 --- a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py @@ -245,6 +245,515 @@ def _validate_chat_history( raise ValueError(error_message) +class _AgentBuilder: + """Internal builder class for constructing React agents with intuitive method-to-node mapping.""" + + def __init__( + self, + model: Union[ + str, + LanguageModelLike, + Callable[[StateSchema, Runtime[ContextT]], BaseChatModel], + Callable[[StateSchema, Runtime[ContextT]], Awaitable[BaseChatModel]], + Callable[ + [StateSchema, Runtime[ContextT]], + Runnable[LanguageModelInput, BaseMessage], + ], + Callable[ + [StateSchema, Runtime[ContextT]], + Awaitable[Runnable[LanguageModelInput, BaseMessage]], + ], + ], + tools: Union[Sequence[Union[BaseTool, Callable, dict[str, Any]]], ToolNode], + *, + prompt: Optional[Prompt] = None, + response_format: Optional[ + Union[StructuredResponseSchema, tuple[str, StructuredResponseSchema]] + ] = None, + pre_model_hook: Optional[RunnableLike] = None, + post_model_hook: Optional[RunnableLike] = None, + state_schema: Optional[StateSchemaType] = None, + context_schema: Optional[Type[Any]] = None, + version: Literal["v1", "v2"] = "v2", + name: Optional[str] = None, + store: Optional[BaseStore] = None, + ): + self.model = model + self.tools = tools + self.prompt = prompt + self.response_format = response_format + self.pre_model_hook = pre_model_hook + self.post_model_hook = post_model_hook + self.state_schema = state_schema + self.context_schema = context_schema + self.version = version + self.name = name + self.store = store + + self._setup_tools() + self._setup_state_schema() + self._setup_model() + + def _setup_tools(self) -> None: + """Setup tool-related attributes.""" + if isinstance(self.tools, ToolNode): + self._tool_classes = list(self.tools.tools_by_name.values()) + self._tool_node = self.tools + self._llm_builtin_tools = [] + else: + self._llm_builtin_tools = [t for t in self.tools if isinstance(t, dict)] + self._tool_node = ToolNode( + [t for t in self.tools if not isinstance(t, dict)] + ) + self._tool_classes = list(self._tool_node.tools_by_name.values()) + + self._should_return_direct = { + t.name for t in self._tool_classes if t.return_direct + } + self._tool_calling_enabled = len(self._tool_classes) > 0 + + def _setup_state_schema(self) -> None: + """Setup state schema with validation.""" + if self.state_schema is not None: + required_keys = {"messages", "remaining_steps"} + if self.response_format is not None: + required_keys.add("structured_response") + + schema_keys = set(get_type_hints(self.state_schema)) + if missing_keys := required_keys - schema_keys: + raise ValueError( + f"Missing required key(s) {missing_keys} in state_schema" + ) + + self._final_state_schema = self.state_schema + else: + self._final_state_schema = ( + AgentStateWithStructuredResponse + if self.response_format is not None + else AgentState + ) + + def _setup_model(self) -> None: + """Setup model-related attributes.""" + self._is_dynamic_model = not isinstance( + self.model, (str, Runnable) + ) and callable(self.model) + self._is_async_dynamic_model = ( + self._is_dynamic_model and inspect.iscoroutinefunction(self.model) + ) + + if not self._is_dynamic_model: + model = self.model + if isinstance(model, str): + try: + from langchain.chat_models import ( + init_chat_model, # type: ignore[import-not-found] + ) + except ImportError: + raise ImportError( + "Please install langchain (`pip install langchain`) to use ':' string syntax for `model` parameter." + ) + model = init_chat_model(model) + + if ( + _should_bind_tools( + model, # type: ignore[arg-type] + self._tool_classes, + num_builtin=len(self._llm_builtin_tools), + ) + and len(self._tool_classes + self._llm_builtin_tools) > 0 + ): + model = cast(BaseChatModel, model).bind_tools( + self._tool_classes + self._llm_builtin_tools # type: ignore[operator] + ) + + self._static_model: Optional[Runnable] = ( + _get_prompt_runnable(self.prompt) | model # type: ignore[operator] + ) + else: + self._static_model = None + + def _resolve_model( + self, state: StateSchema, runtime: Runtime[ContextT] + ) -> LanguageModelLike: + """Resolve the model to use, handling both static and dynamic models.""" + if self._is_dynamic_model: + return _get_prompt_runnable(self.prompt) | self.model(state, runtime) # type: ignore[arg-type, operator] + else: + return self._static_model + + async def _aresolve_model( + self, state: StateSchema, 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]], + self.model, + ) + resolved_model = await dynamic_model(state, runtime) + return _get_prompt_runnable(self.prompt) | resolved_model + elif self._is_dynamic_model: + return _get_prompt_runnable(self.prompt) | self.model(state, runtime) # type: ignore[arg-type,operator] + else: + return self._static_model + + 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}" + ) + + if messages is None: + raise ValueError(error_msg) + + _validate_chat_history(messages) + + if isinstance(self._final_state_schema, type) and issubclass( + self._final_state_schema, BaseModel + ): + # we're passing messages under `messages` key, as this is expected by the prompt + state.messages = messages # type: ignore + else: + state["messages"] = messages # type: ignore + return state + + 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 self._should_return_direct + for call in response.tool_calls + ) + if isinstance(response, AIMessage) + else False + ) + remaining_steps = _get_state_value(state, "remaining_steps", None) + if remaining_steps is not None: + if remaining_steps < 1 and all_tools_return_direct: + return True + elif remaining_steps < 2 and has_tool_calls: + return True + return False + + def call_model( + state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig + ) -> StateSchema: + if self._is_async_dynamic_model: + raise RuntimeError( + "Async model callable provided but agent invoked synchronously. " + "Use agent.ainvoke() or agent.astream(), or provide a sync model callable." + ) + + model_input = _get_model_input_state(state) + model = self._resolve_model(state, runtime) + response = cast(AIMessage, model.invoke(model_input, config)) # type: ignore[arg-type] + response.name = self.name + + if _are_more_steps_needed(state, response): + return { + "messages": [ + AIMessage( + id=response.id, + content="Sorry, need more steps to process this request.", + ) + ] + } + return {"messages": [response]} + + async def acall_model( + state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig + ) -> StateSchema: + model_input = _get_model_input_state(state) + + model = await self._aresolve_model(state, runtime) + response = cast( + AIMessage, + await model.ainvoke(model_input, config), # type: ignore[arg-type] + ) + response.name = self.name + if _are_more_steps_needed(state, response): + return { + "messages": [ + AIMessage( + id=response.id, + content="Sorry, need more steps to process this request.", + ) + ] + } + return {"messages": [response]} + + 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_structured_response_node(self) -> Optional[RunnableCallable]: + """Create the 'generate_structured_response' node if configured.""" + if self.response_format is None: + return None + + def generate_structured_response( + state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig + ) -> StateSchema: + if self._is_async_dynamic_model: + raise RuntimeError( + "Async model callable provided but agent invoked synchronously. " + "Use agent.ainvoke() or agent.astream(), or provide a sync model callable." + ) + + messages = _get_state_value(state, "messages") + structured_response_schema = self.response_format + if isinstance(self.response_format, tuple): + system_prompt, structured_response_schema = self.response_format + messages = [SystemMessage(content=system_prompt)] + list(messages) + + resolved_model = self._resolve_model(state, runtime) + model_with_structured_output = _get_model( + resolved_model + ).with_structured_output( + cast(StructuredResponseSchema, structured_response_schema) + ) + response = model_with_structured_output.invoke(messages, config) + return {"structured_response": response} + + async def agenerate_structured_response( + state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig + ) -> StateSchema: + messages = _get_state_value(state, "messages") + structured_response_schema = self.response_format + if isinstance(self.response_format, tuple): + system_prompt, structured_response_schema = self.response_format + messages = [SystemMessage(content=system_prompt)] + list(messages) + + resolved_model = await self._aresolve_model(state, runtime) + model_with_structured_output = _get_model( + resolved_model + ).with_structured_output( + cast(StructuredResponseSchema, structured_response_schema) + ) + response = await model_with_structured_output.ainvoke(messages, config) + return {"structured_response": response} + + return RunnableCallable( + generate_structured_response, agenerate_structured_response + ) + + def create_model_router(self) -> Callable[[StateSchema], Union[str, list[Send]]]: + """Create routing function for model node conditional edges.""" + + def should_continue(state: StateSchema) -> Union[str, list[Send]]: + messages = _get_state_value(state, "messages") + last_message = messages[-1] + + if not isinstance(last_message, AIMessage) or not last_message.tool_calls: + if self.post_model_hook is not None: + return "post_model_hook" + elif self.response_format is not None: + return "generate_structured_response" + else: + return END + else: + if self.version == "v1": + return "tools" + elif self.version == "v2": + if self.post_model_hook is not None: + return "post_model_hook" + tool_calls = [ + self._tool_node.inject_tool_args(call, state, self.store) # type: ignore[arg-type] + for call in last_message.tool_calls + ] + return [Send("tools", [tool_call]) for tool_call in tool_calls] + + return should_continue + + def create_post_model_hook_router( + self, + ) -> Callable[[StateSchema], Union[str, list[Send]]]: + """Create routing function for post_model_hook node conditional edges.""" + + def post_model_hook_router(state: StateSchema) -> Union[str, list[Send]]: + messages = _get_state_value(state, "messages") + tool_messages = [ + m.tool_call_id for m in messages if isinstance(m, ToolMessage) + ] + last_ai_message = next( + m for m in reversed(messages) if isinstance(m, AIMessage) + ) + pending_tool_calls = [ + c for c in last_ai_message.tool_calls if c["id"] not in tool_messages + ] + + if pending_tool_calls: + pending_tool_calls = [ + self._tool_node.inject_tool_args(call, state, self.store) # type: ignore[arg-type] + for call in pending_tool_calls + ] + return [Send("tools", [tool_call]) for tool_call in pending_tool_calls] + elif isinstance(messages[-1], ToolMessage): + return self._get_entry_point() + elif self.response_format is not None: + return "generate_structured_response" + else: + return END + + return post_model_hook_router + + def create_tools_router(self) -> Optional[Callable[[StateSchema], str]]: + """Create routing function for tools node conditional edges.""" + if not self._should_return_direct: + return None + + def route_tool_responses(state: StateSchema) -> str: + messages = _get_state_value(state, "messages") + for m in reversed(messages): + if not isinstance(m, ToolMessage): + break + if m.name in self._should_return_direct: + return END + + if isinstance(m, AIMessage) and m.tool_calls: + if any( + call["name"] in self._should_return_direct for call in m.tool_calls + ): + return END + + return self._get_entry_point() + + return route_tool_responses + + def _get_entry_point(self) -> str: + """Get the workflow entry point.""" + return "pre_model_hook" if self.pre_model_hook else "agent" + + def _get_model_paths(self) -> list[str]: + """Get possible edge destinations from model node.""" + paths = [] + if self._tool_calling_enabled: + paths.append("tools") + if self.response_format: + paths.append("generate_structured_response") + else: + paths.append(END) + + return paths + + def _get_post_model_hook_paths(self) -> list[str]: + """Get possible edge destinations from post_model_hook node.""" + paths = [] + if self._tool_calling_enabled: + paths = [self._get_entry_point(), "tools"] + if self.response_format is not None: + paths.append("generate_structured_response") + else: + paths.append(END) + return paths + + def build(self) -> StateGraph: + """Build the agent workflow graph (uncompiled).""" + workflow = StateGraph( + state_schema=self._final_state_schema, + context_schema=self.context_schema, + ) + + # Set entry point + workflow.set_entry_point(self._get_entry_point()) + + # Add nodes + workflow.add_node( + "agent", self.create_model_node(), input_schema=self._get_input_schema() + ) + + if self._tool_calling_enabled: + workflow.add_node("tools", self._tool_node) + + if self.pre_model_hook: + workflow.add_node("pre_model_hook", self.pre_model_hook) # type: ignore[arg-type] + + if self.post_model_hook: + workflow.add_node("post_model_hook", self.post_model_hook) # type: ignore[arg-type] + + structured_node = self.create_structured_response_node() + if structured_node: + workflow.add_node("generate_structured_response", structured_node) + + # Add edges + if self.pre_model_hook: + workflow.add_edge("pre_model_hook", "agent") + + if self.post_model_hook: + workflow.add_edge("agent", "post_model_hook") + post_hook_paths = self._get_post_model_hook_paths() + if len(post_hook_paths) == 1: + # No need for a conditional edge if there's only one path + workflow.add_edge("post_model_hook", post_hook_paths[0]) + else: + workflow.add_conditional_edges( + "post_model_hook", + self.create_post_model_hook_router(), + path_map=post_hook_paths, + ) + else: + model_paths = self._get_model_paths() + if len(model_paths) == 1: + # No need for a conditional edge if there's only one path + workflow.add_edge("agent", model_paths[0]) + else: + workflow.add_conditional_edges( + "agent", + self.create_model_router(), + path_map=model_paths, + ) + + if self._tool_calling_enabled: + # In some cases, tools can return directly. In these cases + # we add a conditional edge from the tools node to the END node + # instead of going to the entry point. + tools_router = self.create_tools_router() + if tools_router: + workflow.add_conditional_edges( + "tools", + tools_router, + path_map=[self._get_entry_point(), END], + ) + else: + workflow.add_edge("tools", self._get_entry_point()) + + return workflow + + def create_react_agent( model: Union[ str, @@ -462,6 +971,7 @@ def create_react_agent( print(chunk) ``` """ + # Handle deprecated config_schema parameter if ( config_schema := deprecated_kwargs.pop("config_schema", MISSING) ) is not MISSING: @@ -469,7 +979,6 @@ def create_react_agent( "`config_schema` is deprecated and will be removed. Please use `context_schema` instead.", category=LangGraphDeprecatedSinceV10, ) - if context_schema is None: context_schema = config_schema @@ -483,451 +992,23 @@ def create_react_agent( f"Invalid version {version}. Supported versions are 'v1' and 'v2'." ) - if state_schema is not None: - required_keys = {"messages", "remaining_steps"} - if response_format is not None: - required_keys.add("structured_response") - - 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: - state_schema = ( - AgentStateWithStructuredResponse - if response_format is not None - else AgentState - ) - - llm_builtin_tools: list[dict] = [] - if isinstance(tools, ToolNode): - tool_classes = list(tools.tools_by_name.values()) - tool_node = tools - else: - llm_builtin_tools = [t for t in tools if isinstance(t, dict)] - tool_node = ToolNode([t for t in tools if not isinstance(t, dict)]) - tool_classes = list(tool_node.tools_by_name.values()) - - is_dynamic_model = not isinstance(model, (str, Runnable)) and callable(model) - is_async_dynamic_model = is_dynamic_model and inspect.iscoroutinefunction(model) - - tool_calling_enabled = len(tool_classes) > 0 - - if not is_dynamic_model: - if isinstance(model, str): - try: - from langchain.chat_models import ( # type: ignore[import-not-found] - init_chat_model, - ) - except ImportError: - raise ImportError( - "Please install langchain (`pip install langchain`) to " - "use ':' string syntax for `model` parameter." - ) - - model = cast(BaseChatModel, init_chat_model(model)) - - if ( - _should_bind_tools(model, tool_classes, num_builtin=len(llm_builtin_tools)) # type: ignore[arg-type] - and len(tool_classes + llm_builtin_tools) > 0 - ): - model = cast(BaseChatModel, model).bind_tools( - tool_classes + llm_builtin_tools # type: ignore[operator] - ) - - static_model: Optional[Runnable] = _get_prompt_runnable(prompt) | model # type: ignore[operator] - else: - # For dynamic models, we'll create the runnable at runtime - static_model = None - - # If any of the tools are configured to return_directly after running, - # our graph needs to check if these were called - should_return_direct = {t.name for t in tool_classes if t.return_direct} - - def _resolve_model( - state: StateSchema, runtime: Runtime[ContextT] - ) -> LanguageModelLike: - """Resolve the model to use, handling both static and dynamic models.""" - if is_dynamic_model: - return _get_prompt_runnable(prompt) | model(state, runtime) # type: ignore[operator] - else: - return static_model - - async def _aresolve_model( - state: StateSchema, runtime: Runtime[ContextT] - ) -> LanguageModelLike: - """Async resolve the model to use, handling both static and dynamic models.""" - if is_async_dynamic_model: - resolved_model = await model(state, runtime) # type: ignore[misc,operator] - return _get_prompt_runnable(prompt) | resolved_model - elif is_dynamic_model: - return _get_prompt_runnable(prompt) | model(state, runtime) # type: ignore[operator] - else: - return static_model - - 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 - ) - remaining_steps = _get_state_value(state, "remaining_steps", None) - if remaining_steps is not None: - if remaining_steps < 1 and all_tools_return_direct: - return True - elif remaining_steps < 2 and has_tool_calls: - return True - - return False - - 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, runtime: Runtime[ContextT], config: RunnableConfig - ) -> StateSchema: - if is_async_dynamic_model: - msg = ( - "Async model callable provided but agent invoked synchronously. " - "Use agent.ainvoke() or agent.astream(), or " - "provide a sync model callable." - ) - raise RuntimeError(msg) - - model_input = _get_model_input_state(state) - - if is_dynamic_model: - # Resolve dynamic model at runtime and apply prompt - dynamic_model = _resolve_model(state, runtime) - response = cast(AIMessage, dynamic_model.invoke(model_input, config)) # type: ignore[arg-type] - else: - response = cast(AIMessage, static_model.invoke(model_input, config)) # type: ignore[union-attr] - - # add agent name to the AIMessage - response.name = name - - if _are_more_steps_needed(state, response): - return { - "messages": [ - AIMessage( - id=response.id, - content="Sorry, need more steps to process this request.", - ) - ] - } - # We return a list, because this will get added to the existing list - return {"messages": [response]} - - async def acall_model( - state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig - ) -> StateSchema: - model_input = _get_model_input_state(state) - - if is_dynamic_model: - # Resolve dynamic model at runtime and apply prompt - # (supports both sync and async) - dynamic_model = await _aresolve_model(state, runtime) - response = cast(AIMessage, await dynamic_model.ainvoke(model_input, config)) # type: ignore[arg-type] - else: - response = cast(AIMessage, await static_model.ainvoke(model_input, config)) # type: ignore[union-attr] - - # add agent name to the AIMessage - response.name = name - if _are_more_steps_needed(state, response): - return { - "messages": [ - AIMessage( - id=response.id, - content="Sorry, need more steps to process this request.", - ) - ] - } - # 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, runtime: Runtime[ContextT], config: RunnableConfig - ) -> StateSchema: - if is_async_dynamic_model: - msg = ( - "Async model callable provided but agent invoked synchronously. " - "Use agent.ainvoke() or agent.astream(), or provide a sync model callable." - ) - raise RuntimeError(msg) - - messages = _get_state_value(state, "messages") - structured_response_schema = response_format - if isinstance(response_format, tuple): - system_prompt, structured_response_schema = response_format - messages = [SystemMessage(content=system_prompt)] + list(messages) - - resolved_model = _resolve_model(state, runtime) - model_with_structured_output = _get_model( - resolved_model - ).with_structured_output( - cast(StructuredResponseSchema, structured_response_schema) - ) - response = model_with_structured_output.invoke(messages, config) - return {"structured_response": response} - - async def agenerate_structured_response( - state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig - ) -> StateSchema: - messages = _get_state_value(state, "messages") - structured_response_schema = response_format - if isinstance(response_format, tuple): - system_prompt, structured_response_schema = response_format - messages = [SystemMessage(content=system_prompt)] + list(messages) - - resolved_model = await _aresolve_model(state, runtime) - model_with_structured_output = _get_model( - resolved_model - ).with_structured_output( - cast(StructuredResponseSchema, structured_response_schema) - ) - response = await model_with_structured_output.ainvoke(messages, config) - return {"structured_response": response} - - if not tool_calling_enabled: - # Define a new graph - workflow = StateGraph(state_schema=state_schema, context_schema=context_schema) - workflow.add_node( - "agent", - RunnableCallable(call_model, acall_model), - input_schema=input_schema, - ) - if pre_model_hook is not None: - workflow.add_node("pre_model_hook", pre_model_hook) # type: ignore[arg-type] - workflow.add_edge("pre_model_hook", "agent") - entrypoint = "pre_model_hook" - else: - entrypoint = "agent" - - workflow.set_entry_point(entrypoint) - - if post_model_hook is not None: - workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type] - workflow.add_edge("agent", "post_model_hook") - - if response_format is not None: - workflow.add_node( - "generate_structured_response", - RunnableCallable( - generate_structured_response, - agenerate_structured_response, - ), - ) - if post_model_hook is not None: - workflow.add_edge("post_model_hook", "generate_structured_response") - else: - workflow.add_edge("agent", "generate_structured_response") - - return workflow.compile( - checkpointer=checkpointer, - store=store, - interrupt_before=interrupt_before, - interrupt_after=interrupt_after, - debug=debug, - name=name, - ) - - # Define the function that determines whether to continue or not - def should_continue(state: StateSchema) -> Union[str, list[Send]]: - 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: - if post_model_hook is not None: - return "post_model_hook" - elif response_format is not None: - return "generate_structured_response" - else: - return END - # Otherwise if there is, we continue - else: - if version == "v1": - return "tools" - elif version == "v2": - if post_model_hook is not None: - return "post_model_hook" - tool_calls = [ - tool_node.inject_tool_args(call, state, store) # type: ignore[arg-type] - for call in last_message.tool_calls - ] - return [Send("tools", [tool_call]) for tool_call in tool_calls] - - # Define a new graph - workflow = StateGraph( - state_schema=state_schema or AgentState, context_schema=context_schema + # Create and configure the agent builder + builder = _AgentBuilder( + model=model, + tools=tools, + prompt=prompt, + response_format=response_format, + pre_model_hook=pre_model_hook, + post_model_hook=post_model_hook, + state_schema=state_schema, + context_schema=context_schema, + version=version, + name=name, + store=store, ) - # Define the two nodes we will cycle between - workflow.add_node( - "agent", - RunnableCallable(call_model, acall_model), - input_schema=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) # type: ignore[arg-type] - 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(entrypoint) - - agent_paths = [] - post_model_hook_paths = [entrypoint, "tools"] - - # Add a post model hook node if post_model_hook is provided - if post_model_hook is not None: - workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type] - agent_paths.append("post_model_hook") - workflow.add_edge("agent", "post_model_hook") - else: - agent_paths.append("tools") - - # Add a structured output node if response_format is provided - if response_format is not None: - workflow.add_node( - "generate_structured_response", - RunnableCallable( - generate_structured_response, - agenerate_structured_response, - ), - ) - if post_model_hook is not None: - post_model_hook_paths.append("generate_structured_response") - else: - agent_paths.append("generate_structured_response") - else: - if post_model_hook is not None: - post_model_hook_paths.append(END) - else: - agent_paths.append(END) - - if post_model_hook is not None: - - def post_model_hook_router(state: StateSchema) -> Union[str, list[Send]]: - """Route to the next node after post_model_hook. - - Routes to one of: - * "tools": if there are pending tool calls without a corresponding message. - * "generate_structured_response": if no pending tool calls exist and response_format is specified. - * END: if no pending tool calls exist and no response_format is specified. - """ - - messages = _get_state_value(state, "messages") - tool_messages = [ - m.tool_call_id for m in messages if isinstance(m, ToolMessage) - ] - last_ai_message = next( - m for m in reversed(messages) if isinstance(m, AIMessage) - ) - pending_tool_calls = [ - c for c in last_ai_message.tool_calls if c["id"] not in tool_messages - ] - - if pending_tool_calls: - pending_tool_calls = [ - tool_node.inject_tool_args(call, state, store) # type: ignore[arg-type] - for call in pending_tool_calls - ] - return [Send("tools", [tool_call]) for tool_call in pending_tool_calls] - elif isinstance(messages[-1], ToolMessage): - return entrypoint - elif response_format is not None: - return "generate_structured_response" - else: - return END - - workflow.add_conditional_edges( - "post_model_hook", - post_model_hook_router, - path_map=post_model_hook_paths, - ) - - workflow.add_conditional_edges( - "agent", - should_continue, - path_map=agent_paths, - ) - - 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 - - # handle a case of parallel tool calls where - # the tool w/ `return_direct` was executed in a different `Send` - if isinstance(m, AIMessage) and m.tool_calls: - if any(call["name"] in should_return_direct for call in m.tool_calls): - return END - - return entrypoint - - if should_return_direct: - workflow.add_conditional_edges( - "tools", route_tool_responses, path_map=[entrypoint, END] - ) - else: - workflow.add_edge("tools", entrypoint) - - # Finally, we compile it! - # This compiles it into a LangChain Runnable, - # meaning you can use it as you would any other runnable + # Build and compile the workflow + workflow = builder.build() return workflow.compile( checkpointer=checkpointer, store=store, diff --git a/libs/prebuilt/tests/test_react_agent.py b/libs/prebuilt/tests/test_react_agent.py index cff4c13c3..419c77424 100644 --- a/libs/prebuilt/tests/test_react_agent.py +++ b/libs/prebuilt/tests/test_react_agent.py @@ -184,7 +184,7 @@ def test_runnable_prompt(): @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) -def test_prompt_with_store(version: str): +def test_prompt_with_store(version: Literal["v1", "v2"]): def add(a: int, b: int): """Adds a and b""" return a + b diff --git a/libs/prebuilt/tests/test_react_agent_graph.py b/libs/prebuilt/tests/test_react_agent_graph.py index 50d9f1846..24c827418 100644 --- a/libs/prebuilt/tests/test_react_agent_graph.py +++ b/libs/prebuilt/tests/test_react_agent_graph.py @@ -49,4 +49,14 @@ def test_react_agent_graph_structure( post_model_hook=post_model_hook, response_format=response_format, ) - assert agent.get_graph().draw_mermaid(with_styles=False) == snapshot + try: + assert agent.get_graph().draw_mermaid(with_styles=False) == snapshot + except Exception as e: + raise ValueError( + "The graph structure has changed. Please update the snapshot." + "Configuration used:\n" + f"tools: {tools}, " + f"pre_model_hook: {pre_model_hook}, " + f"post_model_hook: {post_model_hook}, " + f"response_format: {response_format}" + ) from e