From 9e9a5d24985ea281df9bf7498de7edf257677d50 Mon Sep 17 00:00:00 2001 From: Eugene Yurtsev Date: Wed, 13 Aug 2025 09:49:27 -0400 Subject: [PATCH] feat(prebuilt): Split tool node to individual tool nodes (#5888) Add option to split tool node to individual nodes. Summary: * User code (specifically streaming) may break if it's relying on the name of the `tools` node * The boolean flag in the interface is likely **temporary** (especially if there are no major breaking changes) * We'll need to decide if we can get rid of the version in create react agent. "v1" is not consistent conceptually with a node per tool. --- libs/prebuilt/Makefile | 3 + .../langgraph/prebuilt/chat_agent_executor.py | 128 ++++++++++-- .../__snapshots__/test_react_agent_graph.ambr | 188 ++++++++++++++++++ libs/prebuilt/tests/test_react_agent.py | 55 ++++- libs/prebuilt/tests/test_react_agent_graph.py | 35 ++++ 5 files changed, 382 insertions(+), 27 deletions(-) diff --git a/libs/prebuilt/Makefile b/libs/prebuilt/Makefile index d8c0abc58..7bf374b75 100644 --- a/libs/prebuilt/Makefile +++ b/libs/prebuilt/Makefile @@ -30,6 +30,9 @@ test_watch: make stop-services; \ exit $$EXIT_CODE +snapshot_upate: + LANGGRAPH_TEST_FAST=1 uv run pytest --snapshot-update $(TEST) + ###################### # LINTING AND FORMATTING ###################### diff --git a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py index 4f6cb6e9a..f7df559ee 100644 --- a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py @@ -277,7 +277,16 @@ class _AgentBuilder: version: Literal["v1", "v2"] = "v2", name: Optional[str] = None, store: Optional[BaseStore] = None, - ): + use_individual_tool_nodes: bool = False, + ) -> None: + """Initialize the agent builder.""" + if version == "v1" and use_individual_tool_nodes: + # This edge case is ill-defined. "v1" refers specifically to a single + # tools node that handles all tool calls. + raise AssertionError( + "The 'use_individual_tool_nodes' option is only supported " + "in version 'v2' agents." + ) self.model = model self.tools = tools self.prompt = prompt @@ -289,6 +298,7 @@ class _AgentBuilder: self.version = version self.name = name self.store = store + self._use_individual_tool_nodes = use_individual_tool_nodes self._setup_tools() self._setup_state_schema() @@ -591,11 +601,24 @@ class _AgentBuilder: 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] + + if self._use_individual_tool_nodes: + # Route to individual tool nodes + 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(tool_call["name"], [tool_call]) + for tool_call in tool_calls + ] + else: + # Use the original combined tools node + 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 @@ -621,7 +644,18 @@ class _AgentBuilder: 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] + + if self._use_individual_tool_nodes: + return [ + # TODO: Add validation for tool name being a valid node name + # and one that matches a tool. + Send(tool_call["name"], [tool_call]) + for tool_call in pending_tool_calls + ] + else: + 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: @@ -654,6 +688,21 @@ class _AgentBuilder: return route_tool_responses + def add_tool_node(self, tool: BaseTool) -> RunnableCallable: + """Create a node that executes a specific tool. + + This method creates a node that wraps a single tool in a ToolNode + and executes it, returning the result as {"messages": [message]}. + + Args: + tool: The tool to wrap in a node. + + Returns: + A RunnableCallable node that can be added to the graph. + """ + tool_node = ToolNode([tool]) + return tool_node + def _get_entry_point(self) -> str: """Get the workflow entry point.""" return "pre_model_hook" if self.pre_model_hook else "agent" @@ -662,7 +711,10 @@ class _AgentBuilder: """Get possible edge destinations from model node.""" paths = [] if self._tool_calling_enabled: - paths.append("tools") + if self._use_individual_tool_nodes: + paths.extend([tool.name for tool in self._tool_classes]) + else: + paths.append("tools") if self.response_format: paths.append("generate_structured_response") else: @@ -674,7 +726,12 @@ class _AgentBuilder: """Get possible edge destinations from post_model_hook node.""" paths = [] if self._tool_calling_enabled: - paths = [self._get_entry_point(), "tools"] + if self._use_individual_tool_nodes: + paths = [self._get_entry_point()] + [ + tool.name for tool in self._tool_classes + ] + else: + paths = [self._get_entry_point(), "tools"] if self.response_format is not None: paths.append("generate_structured_response") else: @@ -682,7 +739,7 @@ class _AgentBuilder: return paths def build(self) -> StateGraph: - """Build the agent workflow graph (uncompiled).""" + """Build the agent workflow graph.""" workflow = StateGraph( state_schema=self._final_state_schema, context_schema=self.context_schema, @@ -697,7 +754,14 @@ class _AgentBuilder: ) if self._tool_calling_enabled: - workflow.add_node("tools", self._tool_node) + if self._use_individual_tool_nodes: + # Add individual tool nodes + for tool in self._tool_classes: + tool_node = self.add_tool_node(tool) + workflow.add_node(tool.name, tool_node) + else: + # Add the combined tools node + 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] @@ -738,18 +802,32 @@ class _AgentBuilder: ) 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], - ) + if self._use_individual_tool_nodes: + # Add edges for individual tool nodes + tools_router = self.create_tools_router() + for tool in self._tool_classes: + tool_node_name = tool.name + if tools_router: + workflow.add_conditional_edges( + tool_node_name, + tools_router, + path_map=[self._get_entry_point(), END], + ) + else: + workflow.add_edge(tool_node_name, self._get_entry_point()) else: - workflow.add_edge("tools", self._get_entry_point()) + # 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 @@ -785,6 +863,7 @@ def create_react_agent( debug: bool = False, version: Literal["v1", "v2"] = "v2", name: Optional[str] = None, + use_individual_tool_nodes: bool = False, **deprecated_kwargs: Any, ) -> CompiledStateGraph: """Creates an agent graph that calls tools in a loop until a stopping condition is met. @@ -923,6 +1002,10 @@ def create_react_agent( name: An optional name for the CompiledStateGraph. This name will be automatically used when adding ReAct agent graph to another graph as a subgraph node - particularly useful for building multi-agent systems. + use_individual_tool_nodes: A flag indicating whether to use individual tool nodes for each tool. + If set to `True`, each tool will have its own node in the graph. + This has been added for the beta period. The default behavior will change + in v1.0.0 to use individual tool nodes. !!! warning "`config_schema` Deprecated" The `config_schema` parameter is deprecated in v0.6.0 and support will be removed in v2.0.0. @@ -1005,6 +1088,7 @@ def create_react_agent( version=version, name=name, store=store, + use_individual_tool_nodes=use_individual_tool_nodes, ) # Build and compile the workflow diff --git a/libs/prebuilt/tests/__snapshots__/test_react_agent_graph.ambr b/libs/prebuilt/tests/__snapshots__/test_react_agent_graph.ambr index dca7fc751..2323534da 100644 --- a/libs/prebuilt/tests/__snapshots__/test_react_agent_graph.ambr +++ b/libs/prebuilt/tests/__snapshots__/test_react_agent_graph.ambr @@ -171,3 +171,191 @@ ''' # --- +# name: test_react_agent_graph_structure_with_individual_nodes[no_response_format-no_post_hook-no_pre_hook-no_tools] + ''' + graph TD; + __start__ --> agent; + agent --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[no_response_format-no_post_hook-no_pre_hook-two_tools] + ''' + graph TD; + __start__ --> agent; + agent -.-> __end__; + agent -.-> tool; + agent -.-> tool2; + tool --> agent; + tool2 --> agent; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[no_response_format-no_post_hook-with_pre_hook-no_tools] + ''' + graph TD; + __start__ --> pre_model_hook; + pre_model_hook --> agent; + agent --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[no_response_format-no_post_hook-with_pre_hook-two_tools] + ''' + graph TD; + __start__ --> pre_model_hook; + agent -.-> __end__; + agent -.-> tool; + agent -.-> tool2; + pre_model_hook --> agent; + tool --> pre_model_hook; + tool2 --> pre_model_hook; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[no_response_format-with_post_hook-no_pre_hook-no_tools] + ''' + graph TD; + __start__ --> agent; + agent --> post_model_hook; + post_model_hook --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[no_response_format-with_post_hook-no_pre_hook-two_tools] + ''' + graph TD; + __start__ --> agent; + agent --> post_model_hook; + post_model_hook -.-> __end__; + post_model_hook -.-> agent; + post_model_hook -.-> tool; + post_model_hook -.-> tool2; + tool --> agent; + tool2 --> agent; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[no_response_format-with_post_hook-with_pre_hook-no_tools] + ''' + graph TD; + __start__ --> pre_model_hook; + agent --> post_model_hook; + pre_model_hook --> agent; + post_model_hook --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[no_response_format-with_post_hook-with_pre_hook-two_tools] + ''' + graph TD; + __start__ --> pre_model_hook; + agent --> post_model_hook; + post_model_hook -.-> __end__; + post_model_hook -.-> pre_model_hook; + post_model_hook -.-> tool; + post_model_hook -.-> tool2; + pre_model_hook --> agent; + tool --> pre_model_hook; + tool2 --> pre_model_hook; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[with_response_format-no_post_hook-no_pre_hook-no_tools] + ''' + graph TD; + __start__ --> agent; + agent --> generate_structured_response; + generate_structured_response --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[with_response_format-no_post_hook-no_pre_hook-two_tools] + ''' + graph TD; + __start__ --> agent; + agent -.-> generate_structured_response; + agent -.-> tool; + agent -.-> tool2; + tool --> agent; + tool2 --> agent; + generate_structured_response --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[with_response_format-no_post_hook-with_pre_hook-no_tools] + ''' + graph TD; + __start__ --> pre_model_hook; + agent --> generate_structured_response; + pre_model_hook --> agent; + generate_structured_response --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[with_response_format-no_post_hook-with_pre_hook-two_tools] + ''' + graph TD; + __start__ --> pre_model_hook; + agent -.-> generate_structured_response; + agent -.-> tool; + agent -.-> tool2; + pre_model_hook --> agent; + tool --> pre_model_hook; + tool2 --> pre_model_hook; + generate_structured_response --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[with_response_format-with_post_hook-no_pre_hook-no_tools] + ''' + graph TD; + __start__ --> agent; + agent --> post_model_hook; + post_model_hook --> generate_structured_response; + generate_structured_response --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[with_response_format-with_post_hook-no_pre_hook-two_tools] + ''' + graph TD; + __start__ --> agent; + agent --> post_model_hook; + post_model_hook -.-> agent; + post_model_hook -.-> generate_structured_response; + post_model_hook -.-> tool; + post_model_hook -.-> tool2; + tool --> agent; + tool2 --> agent; + generate_structured_response --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[with_response_format-with_post_hook-with_pre_hook-no_tools] + ''' + graph TD; + __start__ --> pre_model_hook; + agent --> post_model_hook; + post_model_hook --> generate_structured_response; + pre_model_hook --> agent; + generate_structured_response --> __end__; + + ''' +# --- +# name: test_react_agent_graph_structure_with_individual_nodes[with_response_format-with_post_hook-with_pre_hook-two_tools] + ''' + graph TD; + __start__ --> pre_model_hook; + agent --> post_model_hook; + post_model_hook -.-> generate_structured_response; + post_model_hook -.-> pre_model_hook; + post_model_hook -.-> tool; + post_model_hook -.-> tool2; + pre_model_hook --> agent; + tool --> pre_model_hook; + tool2 --> pre_model_hook; + generate_structured_response --> __end__; + + ''' +# --- diff --git a/libs/prebuilt/tests/test_react_agent.py b/libs/prebuilt/tests/test_react_agent.py index 419c77424..23b26a715 100644 --- a/libs/prebuilt/tests/test_react_agent.py +++ b/libs/prebuilt/tests/test_react_agent.py @@ -780,14 +780,24 @@ class AgentStateExtraKeyPydantic(AgentStatePydantic): foo: int -@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) +@pytest.mark.parametrize("version", ["v1", "v2"]) @pytest.mark.parametrize( "state_schema", [AgentStateExtraKey, AgentStateExtraKeyPydantic] ) +@pytest.mark.parametrize( + "use_individual_tool_nodes", + [False, True], + ids=["single_tool_node", "node_per_tool"], +) def test_create_react_agent_inject_vars( - version: Literal["v1", "v2"], state_schema: StateSchemaType + version: Literal["v1", "v2"], + state_schema: StateSchemaType, + use_individual_tool_nodes: bool, ) -> None: """Test that the agent can inject state and store into tool functions.""" + if version == "v1" and use_individual_tool_nodes: + pytest.skip("v1 does not support individual tool nodes") + store = InMemoryStore() namespace = ("test",) store.put(namespace, "test_key", {"bar": 3}) @@ -826,6 +836,7 @@ def test_create_react_agent_inject_vars( state_schema=state_schema, store=store, version=version, + use_individual_tool_nodes=use_individual_tool_nodes, ) result = agent.invoke({"messages": [{"role": "user", "content": "hi"}], "foo": 2}) assert result["messages"] == [ @@ -967,7 +978,17 @@ def test_tool_node_messages_key() -> None: @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) -async def test_return_direct(version: str) -> None: +@pytest.mark.parametrize( + "use_individual_tool_nodes", + [False, True], + ids=["single_tool_node", "node_per_tool"], +) +async def test_return_direct( + version: Literal["v1", "v2"], use_individual_tool_nodes: bool +) -> None: + if version == "v1" and use_individual_tool_nodes: + pytest.skip("v1 does not support individual tool nodes") + @dec_tool(return_direct=True) def tool_return_direct(input: str) -> str: """A tool that returns directly.""" @@ -995,6 +1016,7 @@ async def test_return_direct(version: str) -> None: model, [tool_return_direct, tool_normal], version=version, + use_individual_tool_nodes=use_individual_tool_nodes, ) # Test direct return for tool_return_direct @@ -1088,15 +1110,27 @@ def test__get_state_args() -> None: def test_inspect_react() -> None: + """Test that we can inspect the agent and its nodes.""" model = FakeToolCallingModel(tool_calls=[]) agent = create_react_agent(model, []) inspect.getclosurevars(agent.nodes["agent"].bound.func) @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) +@pytest.mark.parametrize( + "use_individual_tool_nodes", + [False, True], + ids=["single_tool_node", "node_per_tool"], +) def test_react_with_subgraph_tools( - sync_checkpointer: BaseCheckpointSaver, version: Literal["v1", "v2"] + sync_checkpointer: BaseCheckpointSaver, + version: Literal["v1", "v2"], + use_individual_tool_nodes: bool, ) -> None: + """Test React agent with subgraph tools.""" + if version == "v1" and use_individual_tool_nodes: + pytest.skip("v1 does not support individual tool nodes") + class State(TypedDict): a: int b: int @@ -1152,6 +1186,7 @@ def test_react_with_subgraph_tools( tool_node, checkpointer=sync_checkpointer, version=version, + use_individual_tool_nodes=use_individual_tool_nodes, ) result = agent.invoke( {"messages": [HumanMessage(content="What's 2 + 3 and 2 * 3?")]}, @@ -1237,8 +1272,17 @@ def test_tool_node_stream_writer() -> None: @pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS) -def test_react_agent_subgraph_streaming_sync(version: Literal["v1", "v2"]) -> None: +@pytest.mark.parametrize( + "use_individual_tool_nodes", + [False, True], + ids=["single_tool_node", "node_per_tool"], +) +def test_react_agent_subgraph_streaming_sync( + version: Literal["v1", "v2"], use_individual_tool_nodes: bool +) -> None: """Test React agent streaming when used as a subgraph node sync version""" + if version == "v1" and use_individual_tool_nodes: + pytest.skip("v1 does not support individual tool nodes") @dec_tool def get_weather(city: str) -> str: @@ -1258,6 +1302,7 @@ def test_react_agent_subgraph_streaming_sync(version: Literal["v1", "v2"]) -> No tools=[get_weather], prompt="You are a helpful travel assistant.", version=version, + use_individual_tool_nodes=use_individual_tool_nodes, ) # Create a subgraph that uses the React agent as a node diff --git a/libs/prebuilt/tests/test_react_agent_graph.py b/libs/prebuilt/tests/test_react_agent_graph.py index 24c827418..6f76fe9e2 100644 --- a/libs/prebuilt/tests/test_react_agent_graph.py +++ b/libs/prebuilt/tests/test_react_agent_graph.py @@ -15,6 +15,11 @@ def tool() -> None: ... +def tool2() -> None: + """Another testing tool.""" + ... + + def pre_model_hook() -> None: """Pre-model hook.""" ... @@ -60,3 +65,33 @@ def test_react_agent_graph_structure( f"post_model_hook: {post_model_hook}, " f"response_format: {response_format}" ) from e + + +@pytest.mark.parametrize("tools", [[], [tool, tool2]], ids=["no_tools", "two_tools"]) +@pytest.mark.parametrize( + "pre_model_hook", [None, pre_model_hook], ids=["no_pre_hook", "with_pre_hook"] +) +@pytest.mark.parametrize( + "post_model_hook", [None, post_model_hook], ids=["no_post_hook", "with_post_hook"] +) +@pytest.mark.parametrize( + "response_format", + [None, ResponseFormat], + ids=["no_response_format", "with_response_format"], +) +def test_react_agent_graph_structure_with_individual_nodes( + snapshot: SnapshotAssertion, + tools: list[Callable], + pre_model_hook: Union[Callable, None], + post_model_hook: Union[Callable, None], + response_format: Union[type[BaseModel], None], +) -> None: + agent = create_react_agent( + model, + tools=tools, + pre_model_hook=pre_model_hook, + post_model_hook=post_model_hook, + response_format=response_format, + use_individual_tool_nodes=True, + ) + assert agent.get_graph().draw_mermaid(with_styles=False) == snapshot