From b97cda4290664bb03a5d73fae986ded8a0ffb244 Mon Sep 17 00:00:00 2001 From: Vadym Barda Date: Tue, 25 Mar 2025 13:51:05 -0400 Subject: [PATCH] prebuilt: add support for multiple Command(graph=Command.PARENT) returned by tools (#4003) --- libs/prebuilt/langgraph/prebuilt/tool_node.py | 54 +++-- libs/prebuilt/tests/test_tool_node.py | 212 ++++++++++-------- 2 files changed, 152 insertions(+), 114 deletions(-) diff --git a/libs/prebuilt/langgraph/prebuilt/tool_node.py b/libs/prebuilt/langgraph/prebuilt/tool_node.py index 8175c9cda..952500163 100644 --- a/libs/prebuilt/langgraph/prebuilt/tool_node.py +++ b/libs/prebuilt/langgraph/prebuilt/tool_node.py @@ -2,6 +2,7 @@ import asyncio import inspect import json from copy import copy, deepcopy +from dataclasses import replace from typing import ( Any, Callable, @@ -35,7 +36,7 @@ from typing_extensions import Annotated, get_args, get_origin from langgraph.errors import GraphBubbleUp from langgraph.store.base import BaseStore -from langgraph.types import Command +from langgraph.types import Command, Send from langgraph.utils.runnable import RunnableCallable INVALID_TOOL_NAME_ERROR_TEMPLATE = ( @@ -239,25 +240,7 @@ class ToolNode(RunnableCallable): *executor.map(self._run_one, tool_calls, input_types, config_list) ] - # preserve existing behavior for non-command tool outputs for backwards - # compatibility - if not any(isinstance(output, Command) for output in outputs): - # TypedDict, pydantic, dataclass, etc. should all be able to load from dict - return outputs if input_type == "list" else {self.messages_key: outputs} - - # LangGraph will automatically handle list of Command and non-command node - # updates - combined_outputs: list[ - Command | list[ToolMessage] | dict[str, list[ToolMessage]] - ] = [] - for output in outputs: - if isinstance(output, Command): - combined_outputs.append(output) - else: - combined_outputs.append( - [output] if input_type == "list" else {self.messages_key: [output]} - ) - return combined_outputs + return self._combine_tool_outputs(outputs, input_type) async def _afunc( self, @@ -275,22 +258,47 @@ class ToolNode(RunnableCallable): *(self._arun_one(call, input_type, config) for call in tool_calls) ) - # preserve existing behavior for non-command tool outputs for backwards compatibility + return self._combine_tool_outputs(outputs, input_type) + + def _combine_tool_outputs( + self, + outputs: list[ToolMessage], + input_type: Literal["list", "dict", "tool_calls"], + ) -> list[Union[Command, list[ToolMessage], dict[str, list[ToolMessage]]]]: + # preserve existing behavior for non-command tool outputs for backwards + # compatibility if not any(isinstance(output, Command) for output in outputs): # TypedDict, pydantic, dataclass, etc. should all be able to load from dict return outputs if input_type == "list" else {self.messages_key: outputs} - # LangGraph will automatically handle list of Command and non-command node updates + # LangGraph will automatically handle list of Command and non-command node + # updates combined_outputs: list[ Command | list[ToolMessage] | dict[str, list[ToolMessage]] ] = [] + + # combine all parent commands with goto into a single parent command + parent_command: Optional[Command] = None for output in outputs: if isinstance(output, Command): - combined_outputs.append(output) + if output.graph is Command.PARENT and isinstance(output.goto, str): + parent_send = [Send(output.goto, output.update)] + if parent_command: + parent_command = replace( + parent_command, + goto=cast(list[Send], parent_command.goto) + parent_send, + ) + else: + parent_command = Command(graph=Command.PARENT, goto=parent_send) + else: + combined_outputs.append(output) else: combined_outputs.append( [output] if input_type == "list" else {self.messages_key: [output]} ) + + if parent_command: + combined_outputs.append(parent_command) return combined_outputs def _run_one( diff --git a/libs/prebuilt/tests/test_tool_node.py b/libs/prebuilt/tests/test_tool_node.py index 2dd5ef635..f2ca725c1 100644 --- a/libs/prebuilt/tests/test_tool_node.py +++ b/libs/prebuilt/tests/test_tool_node.py @@ -17,7 +17,7 @@ from pydantic.v1 import ValidationError as ValidationErrorV1 from langgraph.errors import NodeInterrupt from langgraph.prebuilt import ToolNode from langgraph.prebuilt.tool_node import TOOL_CALL_ERROR_TEMPLATE -from langgraph.types import Command +from langgraph.types import Command, Send from tests.conftest import IS_LANGCHAIN_CORE_030_OR_GREATER pytestmark = pytest.mark.anyio @@ -595,16 +595,20 @@ async def test_tool_node_command(input_type: str): ] }, Command( - update={ - "messages": [ - ToolMessage( - content="Transferred to Bob", - tool_call_id="2", - name="transfer_to_bob", - ) - ] - }, - goto="bob", + goto=[ + Send( + "bob", + { + "messages": [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="2", + name="transfer_to_bob", + ) + ] + }, + ) + ], graph=Command.PARENT, ), ] @@ -624,16 +628,20 @@ async def test_tool_node_command(input_type: str): ) assert result == [ Command( - update={ - "messages": [ - ToolMessage( - content="Transferred to Bob", - tool_call_id="1", - name=tool.name, - ) - ] - }, - goto="bob", + goto=[ + Send( + "bob", + { + "messages": [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="1", + name=tool.name, + ) + ] + }, + ) + ], graph=Command.PARENT, ) ] @@ -651,16 +659,20 @@ async def test_tool_node_command(input_type: str): ) assert result == [ Command( - update={ - "messages": [ - ToolMessage( - content="Transferred to Bob", - tool_call_id="1", - name=tool.name, - ) - ] - }, - goto="bob", + goto=[ + Send( + "bob", + { + "messages": [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="1", + name=tool.name, + ) + ] + }, + ) + ], graph=Command.PARENT, ) ] @@ -681,30 +693,33 @@ async def test_tool_node_command(input_type: str): ) assert result == [ Command( - update={ - "messages": [ - ToolMessage( - content="Transferred to Bob", - tool_call_id="1", - name="transfer_to_bob", - ) - ] - }, - goto="bob", - graph=Command.PARENT, - ), - Command( - update={ - "messages": [ - ToolMessage( - content="Transferred to Bob", - tool_call_id="2", - name="custom_transfer_to_bob", - ) - ] - }, - goto="bob", graph=Command.PARENT, + goto=[ + Send( + "bob", + { + "messages": [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="1", + name="transfer_to_bob", + ) + ] + }, + ), + Send( + "bob", + { + "messages": [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="2", + name="custom_transfer_to_bob", + ) + ] + }, + ), + ], ), ] @@ -891,15 +906,19 @@ async def test_tool_node_command_list_input(): ) ], Command( - update=[ - ToolMessage( - content="Transferred to Bob", - tool_call_id="2", - name="transfer_to_bob", + graph=Command.PARENT, + goto=[ + Send( + "bob", + [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="2", + name="transfer_to_bob", + ) + ], ) ], - goto="bob", - graph=Command.PARENT, ), ] @@ -912,14 +931,18 @@ async def test_tool_node_command_list_input(): ) assert result == [ Command( - update=[ - ToolMessage( - content="Transferred to Bob", - tool_call_id="1", - name=tool.name, + goto=[ + Send( + "bob", + [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="1", + name=tool.name, + ) + ], ) ], - goto="bob", graph=Command.PARENT, ) ] @@ -931,14 +954,18 @@ async def test_tool_node_command_list_input(): ) assert result == [ Command( - update=[ - ToolMessage( - content="Transferred to Bob", - tool_call_id="1", - name=tool.name, + goto=[ + Send( + "bob", + [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="1", + name=tool.name, + ) + ], ) ], - goto="bob", graph=Command.PARENT, ) ] @@ -957,25 +984,28 @@ async def test_tool_node_command_list_input(): ) assert result == [ Command( - update=[ - ToolMessage( - content="Transferred to Bob", - tool_call_id="1", - name="transfer_to_bob", - ) + goto=[ + Send( + "bob", + [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="1", + name="transfer_to_bob", + ) + ], + ), + Send( + "bob", + [ + ToolMessage( + content="Transferred to Bob", + tool_call_id="2", + name="custom_transfer_to_bob", + ) + ], + ), ], - goto="bob", - graph=Command.PARENT, - ), - Command( - update=[ - ToolMessage( - content="Transferred to Bob", - tool_call_id="2", - name="custom_transfer_to_bob", - ) - ], - goto="bob", graph=Command.PARENT, ), ]