langgraph[patch]: allow ToolNode to accept ToolCalls (#3126)

Alternative to https://github.com/langchain-ai/langgraph/pull/3124

Currently if a tool interrupts, the entire tool node executes again
after resuming. So tools can get executed twice if parallel tool calls
are generated. Here we allow ToolNode to accept tool calls, so we can
use the `Send` API to distribute the tool calls to multiple instances of
the tool node.

```python
from langchain_anthropic import ChatAnthropic
from langchain_core.tools import tool
from langgraph.checkpoint.memory import MemorySaver
from langgraph.prebuilt import create_react_agent
from langgraph.types import Command, Send, interrupt


@tool
def human_assistance(query: str) -> str:
    """Request assistance from a human."""
    human_response = interrupt({"query": query})
    return human_response["data"]


@tool
def get_weather(location: str) -> str:
    """Use this tool to get the weather."""
    return "It's sunny!"


tools = [get_weather, human_assistance]
llm = ChatAnthropic(model="claude-3-5-sonnet-20240620")

agent = create_react_agent(
    llm,
    tools,
    checkpointer=MemorySaver(),
    tool_call_parallelism="parallel_tool_nodes",
)


user_input = (
    "Could you please (1) request assistance for building an AI agent "
    "from a human, and (2) search for the weather in Boston, MA? "
    "Generate two tool calls at once."
)

config = {"configurable": {"thread_id": "1"}}

for event in agent.stream(
    {"messages": [{"role": "user", "content": user_input}]},
    config,
    stream_mode="values",
):
    event["messages"][-1].pretty_print()
```
```
...
```
```python
human_response = "You should check out LangGraph to build your agent."
human_command = Command(resume={"data": human_response})

for event in agent.stream(human_command, config, stream_mode="values"):
    event["messages"][-1].pretty_print()
```

---------

Co-authored-by: Vadym Barda <vadym@langchain.dev>
This commit is contained in:
ccurme
2025-01-31 17:20:59 +00:00
committed by GitHub
co-authored by Vadym Barda
parent 4b3e07b67a
commit a37c4d6f49
3 changed files with 331 additions and 50 deletions
@@ -31,7 +31,7 @@ from langgraph.managed import IsLastStep, RemainingSteps
from langgraph.prebuilt.tool_executor import ToolExecutor
from langgraph.prebuilt.tool_node import ToolNode
from langgraph.store.base import BaseStore
from langgraph.types import Checkpointer
from langgraph.types import Checkpointer, Send
from langgraph.utils.runnable import RunnableCallable
StructuredResponse = Union[dict, BaseModel]
@@ -248,6 +248,9 @@ def create_react_agent(
interrupt_before: Optional[list[str]] = None,
interrupt_after: Optional[list[str]] = None,
debug: bool = False,
tool_call_parallelism: Literal[
"single_tool_node", "parallel_tool_nodes"
] = "single_tool_node",
) -> CompiledGraph:
"""Creates a graph that works with a chat model that utilizes tool calling.
@@ -297,6 +300,15 @@ def create_react_agent(
Should be one of the following: "agent", "tools".
This is useful if you want to return directly or run additional processing on an output.
debug: A flag indicating whether to enable debug mode.
tool_call_parallelism: Determines what state is sent to the tool node.
Can be one of:
- `"single_tool_node"`: The tool node processes a single message. All tool
calls in the message are executed in parallel within the tool node.
- `"parallel_tool_nodes"`: The tool node processes a tool call.
Tool calls are distributed across multiple instances of the tool
node using the [Send](https://langchain-ai.github.io/langgraph/concepts/low_level/#send)
API.
Returns:
A compiled LangChain runnable that can be used for chat interactions.
@@ -572,6 +584,11 @@ def create_react_agent(
TimeoutError: Timed out at step 2
```
"""
if tool_call_parallelism not in ("single_tool_node", "parallel_tool_nodes"):
raise ValueError(
f"Invalid version {tool_call_parallelism}. Supported versions are "
"'single_tool_node' and 'parallel_tool_nodes'."
)
if state_schema is not None:
required_keys = {"messages", "remaining_steps"}
@@ -747,7 +764,7 @@ def create_react_agent(
)
# Define the function that determines whether to continue or not
def should_continue(state: AgentState) -> str:
def should_continue(state: AgentState) -> Union[str, list]:
messages = state["messages"]
last_message = messages[-1]
# If there is no function call, then we finish
@@ -755,7 +772,14 @@ def create_react_agent(
return END if response_format is None else "generate_structured_response"
# Otherwise if there is, we continue
else:
return "tools"
if tool_call_parallelism == "single_tool_node":
return "tools"
elif tool_call_parallelism == "parallel_tool_nodes":
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 or AgentState)
+55 -12
View File
@@ -135,6 +135,8 @@ class ToolNode(RunnableCallable):
If multiple tool calls are requested, they will be run in parallel. The output will be
a list of ToolMessages, one for each tool call.
Tool calls can also be passed directly as a list of `ToolCall` dicts.
Args:
tools: A sequence of tools that can be invoked by the ToolNode.
name: The name of the ToolNode in the graph. Defaults to "tools".
@@ -168,10 +170,28 @@ class ToolNode(RunnableCallable):
return {"messages": result}
```
Tool calls can also be passed directly to a ToolNode. This can be useful when using
the Send API, e.g., in a conditional edge:
```python
def example_conditional_edge(state: dict) -> List[Send]:
tool_calls = state["messages"][-1].tool_calls
# If tools rely on state or store variables (whose values are not generated
# directly by a model), you can inject them into the tool calls.
tool_calls = [
tool_node.inject_tool_args(call, state, store)
for call in last_message.tool_calls
]
return [Send("tools", [tool_call]) for tool_call in tool_calls]
```
Important:
- The state MUST contain a list of messages.
- The last message MUST be an `AIMessage`.
- The `AIMessage` MUST have `tool_calls` populated.
- The input state can be one of the following:
- A dict with a messages key containing a list of messages.
- A list of messages.
- A list of tool calls.
- If operating on a message list, the last message must be an `AIMessage` with
`tool_calls` populated.
"""
name: str = "ToolNode"
@@ -276,7 +296,7 @@ class ToolNode(RunnableCallable):
def _run_one(
self,
call: ToolCall,
input_type: Literal["list", "dict"],
input_type: Literal["list", "dict", "tool_calls"],
config: RunnableConfig,
) -> ToolMessage:
if invalid_tool_message := self._validate_tool_call(call):
@@ -331,7 +351,7 @@ class ToolNode(RunnableCallable):
async def _arun_one(
self,
call: ToolCall,
input_type: Literal["list", "dict"],
input_type: Literal["list", "dict", "tool_calls"],
config: RunnableConfig,
) -> ToolMessage:
if invalid_tool_message := self._validate_tool_call(call):
@@ -392,10 +412,15 @@ class ToolNode(RunnableCallable):
BaseModel,
],
store: Optional[BaseStore],
) -> Tuple[list[ToolCall], Literal["list", "dict"]]:
) -> Tuple[list[ToolCall], Literal["list", "dict", "tool_calls"]]:
if isinstance(input, list):
input_type = "list"
message: AnyMessage = input[-1]
if isinstance(input[-1], dict) and input[-1].get("type") == "tool_call":
input_type = "tool_calls"
tool_calls = input
return tool_calls, input_type
else:
input_type = "list"
message: AnyMessage = input[-1]
elif isinstance(input, dict) and (messages := input.get(self.messages_key, [])):
input_type = "dict"
message = messages[-1]
@@ -410,7 +435,7 @@ class ToolNode(RunnableCallable):
raise ValueError("Last message is not an AIMessage")
tool_calls = [
self._inject_tool_args(call, input, store) for call in message.tool_calls
self.inject_tool_args(call, input, store) for call in message.tool_calls
]
return tool_calls, input_type
@@ -490,7 +515,7 @@ class ToolNode(RunnableCallable):
}
return tool_call
def _inject_tool_args(
def inject_tool_args(
self,
tool_call: ToolCall,
input: Union[
@@ -500,6 +525,21 @@ class ToolNode(RunnableCallable):
],
store: Optional[BaseStore],
) -> ToolCall:
"""Injects the state and store into the tool call.
Tool arguments with types annotated as `InjectedState` and `InjectedStore` are
ignored in tool schemas for generation purposes. This method injects them into
tool calls for tool invocation.
Args:
tool_call (ToolCall): The tool call to inject state and store into.
input (Union[list[AnyMessage], dict[str, Any], BaseModel]): The input state
to inject.
store (Optional[BaseStore]): The store to inject.
Returns:
ToolCall: The tool call with injected state and store.
"""
if tool_call["name"] not in self.tools_by_name:
return tool_call
@@ -509,11 +549,14 @@ class ToolNode(RunnableCallable):
return tool_call_with_store
def _validate_tool_command(
self, command: Command, call: ToolCall, input_type: Literal["list", "dict"]
self,
command: Command,
call: ToolCall,
input_type: Literal["list", "dict", "tool_calls"],
) -> Command:
if isinstance(command.update, dict):
# input type is dict when ToolNode is invoked with a dict input (e.g. {"messages": [AIMessage(..., tool_calls=[...])]})
if input_type != "dict":
if input_type not in ("dict", "tool_calls"):
raise ValueError(
f"Tools can provide a dict in Command.update only when using dict with '{self.messages_key}' key as ToolNode input, "
f"got: {command.update} for tool '{call['name']}'"
+249 -35
View File
@@ -73,6 +73,8 @@ from tests.messages import _AnyIdHumanMessage, _AnyIdToolMessage
pytestmark = pytest.mark.anyio
REACT_TOOL_CALL_PARALLELISM = ["single_tool_node", "parallel_tool_nodes"]
class FakeToolCallingModel(BaseChatModel):
tool_calls: Optional[list[list[ToolCall]]] = None
@@ -148,13 +150,21 @@ class FakeToolCallingModel(BaseChatModel):
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
def test_no_prompt(request: pytest.FixtureRequest, checkpointer_name: str) -> None:
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
def test_no_prompt(
request: pytest.FixtureRequest, checkpointer_name: str, tool_call_parallelism: str
) -> None:
checkpointer: BaseCheckpointSaver = request.getfixturevalue(
"checkpointer_" + checkpointer_name
)
model = FakeToolCallingModel()
agent = create_react_agent(model, [], checkpointer=checkpointer)
agent = create_react_agent(
model,
[],
checkpointer=checkpointer,
tool_call_parallelism=tool_call_parallelism,
)
inputs = [HumanMessage("hi?")]
thread = {"configurable": {"thread_id": "123"}}
response = agent.invoke({"messages": inputs}, thread, debug=True)
@@ -316,7 +326,8 @@ def test_runnable_prompt():
assert response == expected_response
def test_prompt_with_store():
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
def test_prompt_with_store(tool_call_parallelism: str):
def add(a: int, b: int):
"""Adds a and b"""
return a + b
@@ -336,7 +347,13 @@ def test_prompt_with_store():
model = FakeToolCallingModel()
# test state modifier that uses store works
agent = create_react_agent(model, [add], prompt=prompt, store=in_memory_store)
agent = create_react_agent(
model,
[add],
prompt=prompt,
store=in_memory_store,
tool_call_parallelism=tool_call_parallelism,
)
response = agent.invoke(
{"messages": [("user", "hi")]}, {"configurable": {"user_id": "1"}}
)
@@ -344,7 +361,11 @@ def test_prompt_with_store():
# test state modifier that doesn't use store works
agent = create_react_agent(
model, [add], prompt=prompt_no_store, store=in_memory_store
model,
[add],
prompt=prompt_no_store,
store=in_memory_store,
tool_call_parallelism=tool_call_parallelism,
)
response = agent.invoke(
{"messages": [("user", "hi")]}, {"configurable": {"user_id": "2"}}
@@ -395,7 +416,8 @@ async def test_prompt_with_store_async():
@pytest.mark.parametrize("tool_style", ["openai", "anthropic"])
def test_model_with_tools(tool_style: str):
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
def test_model_with_tools(tool_style: str, tool_call_parallelism: str):
model = FakeToolCallingModel(tool_style=tool_style)
@dec_tool
@@ -409,7 +431,11 @@ def test_model_with_tools(tool_style: str):
return f"Tool 2: {some_val}"
# check valid agent constructor
agent = create_react_agent(model.bind_tools([tool1, tool2]), [tool1, tool2])
agent = create_react_agent(
model.bind_tools([tool1, tool2]),
[tool1, tool2],
tool_call_parallelism=tool_call_parallelism,
)
result = agent.nodes["tools"].invoke(
{
"messages": [
@@ -576,7 +602,8 @@ def test__infer_handled_types() -> None:
not IS_LANGCHAIN_CORE_030_OR_GREATER,
reason="Pydantic v1 is required for this test to pass in langchain-core < 0.3",
)
def test_react_agent_with_structured_response() -> None:
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
def test_react_agent_with_structured_response(tool_call_parallelism: str) -> None:
class WeatherResponse(BaseModel):
temperature: float = Field(description="The temperature in fahrenheit")
@@ -592,7 +619,10 @@ def test_react_agent_with_structured_response() -> None:
)
for response_format in (WeatherResponse, ("Meow", WeatherResponse)):
agent = create_react_agent(
model, [get_weather], response_format=response_format
model,
[get_weather],
response_format=response_format,
tool_call_parallelism=tool_call_parallelism,
)
response = agent.invoke({"messages": [HumanMessage("What's the weather?")]})
assert response["structured_response"] == expected_structured_response
@@ -600,7 +630,7 @@ def test_react_agent_with_structured_response() -> None:
assert response["messages"][-2].content == "The weather is sunny and 75°F."
# tools for testing Too
# tools for testing ToolNode
def tool1(some_val: int, some_other_val: str) -> str:
"""Tool 1 docstring."""
if some_val == 0:
@@ -731,6 +761,47 @@ async def test_tool_node():
assert tool_message.tool_call_id == "some 3"
async def test_tool_node_tool_call_input():
# Single tool call
tool_call_1 = {
"name": "tool1",
"args": {"some_val": 1, "some_other_val": "foo"},
"id": "some 0",
"type": "tool_call",
}
result = ToolNode([tool1]).invoke([tool_call_1])
assert result["messages"] == [
ToolMessage(content="1 - foo", tool_call_id="some 0", name="tool1"),
]
# Multiple tool calls
tool_call_2 = {
"name": "tool1",
"args": {"some_val": 2, "some_other_val": "bar"},
"id": "some 1",
"type": "tool_call",
}
result = ToolNode([tool1]).invoke([tool_call_1, tool_call_2])
assert result["messages"] == [
ToolMessage(content="1 - foo", tool_call_id="some 0", name="tool1"),
ToolMessage(content="2 - bar", tool_call_id="some 1", name="tool1"),
]
# Test with unknown tool
tool_call_3 = tool_call_1.copy()
tool_call_3["name"] = "tool2"
result = ToolNode([tool1]).invoke([tool_call_1, tool_call_3])
assert result["messages"] == [
ToolMessage(content="1 - foo", tool_call_id="some 0", name="tool1"),
ToolMessage(
content="Error: tool2 is not a valid tool, try one of [tool1].",
name="tool2",
tool_call_id="some 0",
status="error",
),
]
async def test_tool_node_error_handling():
def handle_all(e: Union[ValueError, ToolException, ValidationError]):
return TOOL_CALL_ERROR_TEMPLATE.format(error=repr(e))
@@ -999,7 +1070,8 @@ def test_tool_node_incorrect_tool_name():
assert tool_message.tool_call_id == "some 0"
def test_tool_node_node_interrupt():
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
def test_tool_node_node_interrupt(tool_call_parallelism: str):
def tool_normal(some_val: int) -> str:
"""Tool docstring."""
return "normal"
@@ -1045,13 +1117,14 @@ def test_tool_node_node_interrupt():
checkpointer = MemorySaver()
config = {"configurable": {"thread_id": "1"}}
agent = create_react_agent(
model, [tool_interrupt, tool_normal], checkpointer=checkpointer
model,
[tool_interrupt, tool_normal],
checkpointer=checkpointer,
tool_call_parallelism=tool_call_parallelism,
)
result = agent.invoke({"messages": [HumanMessage("hi?")]}, config)
assert result["messages"] == [
_AnyIdHumanMessage(
content="hi?",
),
expected_messages = [
_AnyIdHumanMessage(content="hi?"),
AIMessage(
content="hi?",
id="0",
@@ -1070,7 +1143,14 @@ def test_tool_node_node_interrupt():
},
],
),
_AnyIdToolMessage(content="normal", name="tool_normal", tool_call_id="2"),
]
if tool_call_parallelism == "single_tool_node":
# Interrupt blocks second tool result
assert result["messages"] == expected_messages[:-1]
elif tool_call_parallelism == "parallel_tool_nodes":
assert result["messages"] == expected_messages
state = agent.get_state(config)
assert state.next == ("tools",)
task = state.tasks[0]
@@ -1082,7 +1162,8 @@ def test_tool_node_node_interrupt():
not IS_LANGCHAIN_CORE_030_OR_GREATER,
reason="Langchain core 0.3.0 or greater is required",
)
async def test_tool_node_command():
@pytest.mark.parametrize("input_type", ["dict", "tool_calls"])
async def test_tool_node_command(input_type: str):
from langchain_core.tools.base import InjectedToolCallId
@dec_tool
@@ -1159,19 +1240,15 @@ async def test_tool_node_command():
"""Add two numbers"""
return a + b
result = ToolNode([add, transfer_to_bob]).invoke(
{
"messages": [
AIMessage(
"",
tool_calls=[
{"args": {"a": 1, "b": 2}, "id": "1", "name": "add"},
{"args": {}, "id": "2", "name": "transfer_to_bob"},
],
)
]
}
)
tool_calls = [
{"args": {"a": 1, "b": 2}, "id": "1", "name": "add", "type": "tool_call"},
{"args": {}, "id": "2", "name": "transfer_to_bob", "type": "tool_call"},
]
if input_type == "dict":
input_ = {"messages": [AIMessage("", tool_calls=tool_calls)]}
elif input_type == "tool_calls":
input_ = tool_calls
result = ToolNode([add, transfer_to_bob]).invoke(input_)
assert result == [
{
@@ -1646,7 +1723,8 @@ async def test_tool_node_command_list_input():
not IS_LANGCHAIN_CORE_030_OR_GREATER,
reason="Langchain core 0.3.0 or greater is required",
)
def test_react_agent_update_state():
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
def test_react_agent_update_state(tool_call_parallelism: str):
from langchain_core.tools.base import InjectedToolCallId
class State(AgentState):
@@ -1684,6 +1762,7 @@ def test_react_agent_update_state():
state_schema=State,
prompt=prompt,
checkpointer=checkpointer,
tool_call_parallelism=tool_call_parallelism,
)
config = {"configurable": {"thread_id": "1"}}
# run until interrpupted
@@ -1699,6 +1778,87 @@ def test_react_agent_update_state():
assert tool_message.name == "get_user_name"
@pytest.mark.skipif(
not IS_LANGCHAIN_CORE_030_OR_GREATER,
reason="Langchain core 0.3.0 or greater is required",
)
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
def test_react_agent_parallel_tool_calls(tool_call_parallelism: str):
human_assistance_execution_count = 0
@dec_tool
def human_assistance(query: str) -> str:
"""Request assistance from a human."""
nonlocal human_assistance_execution_count
human_response = interrupt({"query": query})
human_assistance_execution_count += 1
return human_response["data"]
get_weather_execution_count = 0
@dec_tool
def get_weather(location: str) -> str:
"""Use this tool to get the weather."""
nonlocal get_weather_execution_count
get_weather_execution_count += 1
return "It's sunny!"
checkpointer = MemorySaver()
tool_calls = [
[
{"args": {"location": "sf"}, "id": "1", "name": "get_weather"},
{"args": {"query": "request help"}, "id": "2", "name": "human_assistance"},
],
[],
]
model = FakeToolCallingModel(tool_calls=tool_calls)
agent = create_react_agent(
model,
[human_assistance, get_weather],
checkpointer=checkpointer,
tool_call_parallelism=tool_call_parallelism,
)
config = {"configurable": {"thread_id": "1"}}
query = "Get user assistance and also check the weather"
message_types = []
for event in agent.stream(
{"messages": [("user", query)]}, config, stream_mode="values"
):
message_types.append([message.type for message in event["messages"]])
if tool_call_parallelism == "single_tool_node":
assert message_types == [
["human"],
["human", "ai"],
]
elif tool_call_parallelism == "parallel_tool_nodes":
assert message_types == [
["human"],
["human", "ai"],
["human", "ai", "tool"],
]
# Resume
message_types = []
for event in agent.stream(
Command(resume={"data": "Hello"}), config, stream_mode="values"
):
message_types.append([message.type for message in event["messages"]])
assert message_types == [
["human", "ai"],
["human", "ai", "tool", "tool"],
["human", "ai", "tool", "tool", "ai"],
]
if tool_call_parallelism == "single_tool_node":
assert human_assistance_execution_count == 1
assert get_weather_execution_count == 2
elif tool_call_parallelism == "parallel_tool_nodes":
assert human_assistance_execution_count == 1
assert get_weather_execution_count == 1
def my_function(some_val: int, some_other_val: str) -> str:
return f"{some_val} - {some_other_val}"
@@ -1887,6 +2047,49 @@ def test_tool_node_inject_state(schema_: Type[T]) -> None:
assert tool_message.content == "hi?"
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
def test_create_react_agent_inject_vars(tool_call_parallelism: str) -> None:
class AgentStateExtraKey(AgentState):
foo: int
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
tool_call = {
"name": "tool1",
"args": {"some_val": 1},
"id": "some 0",
"type": "tool_call",
}
model = FakeToolCallingModel(tool_calls=[[tool_call], []])
agent = create_react_agent(
model,
[tool1],
state_schema=AgentStateExtraKey,
store=store,
tool_call_parallelism=tool_call_parallelism,
)
input_message = HumanMessage("hi")
result = agent.invoke({"messages": [input_message], "foo": 2})
assert result["messages"] == [
input_message,
AIMessage(content="hi", tool_calls=[tool_call], id="0"),
_AnyIdToolMessage(content="6", name="tool1", tool_call_id="some 0"),
AIMessage("hi-hi-6", id="1"),
]
assert result["foo"] == 2
@pytest.mark.skipif(
not IS_LANGCHAIN_CORE_030_OR_GREATER,
reason="Langchain core 0.3.0 or greater is required",
@@ -2020,7 +2223,8 @@ def test_tool_node_messages_key() -> None:
]
async def test_return_direct() -> None:
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
async def test_return_direct(tool_call_parallelism: str) -> None:
@dec_tool(return_direct=True)
def tool_return_direct(input: str) -> str:
"""A tool that returns directly."""
@@ -2044,7 +2248,11 @@ async def test_return_direct() -> None:
tool_calls=first_tool_call,
)
model = FakeToolCallingModel(tool_calls=[first_tool_call, []])
agent = create_react_agent(model, [tool_return_direct, tool_normal])
agent = create_react_agent(
model,
[tool_return_direct, tool_normal],
tool_call_parallelism=tool_call_parallelism,
)
# Test direct return for tool_return_direct
result = agent.invoke(
@@ -2138,7 +2346,8 @@ def test_inspect_react() -> None:
inspect.getclosurevars(agent.nodes["agent"].bound.func)
def test_react_with_subgraph_tools() -> None:
@pytest.mark.parametrize("tool_call_parallelism", REACT_TOOL_CALL_PARALLELISM)
def test_react_with_subgraph_tools(tool_call_parallelism: str) -> None:
class State(TypedDict):
a: int
b: int
@@ -2187,7 +2396,12 @@ def test_react_with_subgraph_tools() -> None:
)
checkpointer = MemorySaver()
tool_node = ToolNode([addition, multiplication], handle_tool_errors=False)
agent = create_react_agent(model, tool_node, checkpointer=checkpointer)
agent = create_react_agent(
model,
tool_node,
checkpointer=checkpointer,
tool_call_parallelism=tool_call_parallelism,
)
result = agent.invoke(
{"messages": [HumanMessage(content="What's 2 + 3 and 2 * 3?")]},
config={"configurable": {"thread_id": "1"}},