mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-06 17:57:49 +02:00
prebuilt: switch to executing parallel tool calls via Send by default (#4438)
This commit is contained in:
@@ -39,6 +39,7 @@ from langgraph.types import (
|
||||
interrupt,
|
||||
)
|
||||
from tests.agents import AgentAction, AgentFinish
|
||||
from tests.any_int import AnyInt
|
||||
from tests.any_str import AnyDict, AnyStr, UnsortedSequence
|
||||
from tests.conftest import (
|
||||
ALL_CHECKPOINTERS_SYNC,
|
||||
@@ -2480,13 +2481,15 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
]
|
||||
}
|
||||
|
||||
assert [
|
||||
events = [
|
||||
c
|
||||
for c in app.stream(
|
||||
{"messages": [HumanMessage(content="what is weather in sf")]},
|
||||
stream_mode="messages",
|
||||
)
|
||||
] == [
|
||||
]
|
||||
|
||||
assert events[:3] == [
|
||||
(
|
||||
_AnyIdAIMessageChunk(
|
||||
content="",
|
||||
@@ -2528,8 +2531,8 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
{
|
||||
"langgraph_step": 2,
|
||||
"langgraph_node": "tools",
|
||||
"langgraph_triggers": ("branch:to:tools",),
|
||||
"langgraph_path": (PULL, "tools"),
|
||||
"langgraph_triggers": (PUSH,),
|
||||
"langgraph_path": (PUSH, AnyInt(), False),
|
||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||
},
|
||||
),
|
||||
@@ -2578,6 +2581,9 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
"ls_model_type": "chat",
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
assert events[3:5] == UnsortedSequence(
|
||||
(
|
||||
_AnyIdToolMessage(
|
||||
content="result for another",
|
||||
@@ -2587,8 +2593,8 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
{
|
||||
"langgraph_step": 4,
|
||||
"langgraph_node": "tools",
|
||||
"langgraph_triggers": ("branch:to:tools",),
|
||||
"langgraph_path": (PULL, "tools"),
|
||||
"langgraph_triggers": (PUSH,),
|
||||
"langgraph_path": (PUSH, AnyInt(), False),
|
||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||
},
|
||||
),
|
||||
@@ -2601,11 +2607,13 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
{
|
||||
"langgraph_step": 4,
|
||||
"langgraph_node": "tools",
|
||||
"langgraph_triggers": ("branch:to:tools",),
|
||||
"langgraph_path": (PULL, "tools"),
|
||||
"langgraph_triggers": (PUSH,),
|
||||
"langgraph_path": (PUSH, AnyInt(), False),
|
||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||
},
|
||||
),
|
||||
)
|
||||
assert events[5:] == [
|
||||
(
|
||||
_AnyIdAIMessageChunk(
|
||||
content="answer",
|
||||
@@ -2636,12 +2644,17 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
|
||||
model.i = 0 # reset the model
|
||||
|
||||
assert (
|
||||
app.invoke(
|
||||
{"messages": [HumanMessage(content="what is weather in sf")]},
|
||||
stream_mode="updates",
|
||||
)[0]["agent"]["messages"]
|
||||
== [
|
||||
invoke_updates_events = app.invoke(
|
||||
{"messages": [HumanMessage(content="what is weather in sf")]},
|
||||
stream_mode="updates",
|
||||
)
|
||||
|
||||
stream_updates_events = [
|
||||
*app.stream({"messages": [HumanMessage(content="what is weather in sf")]})
|
||||
]
|
||||
|
||||
for output in (invoke_updates_events, stream_updates_events):
|
||||
assert output[:3] == [
|
||||
{
|
||||
"agent": {
|
||||
"messages": [
|
||||
@@ -2690,6 +2703,8 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
]
|
||||
}
|
||||
},
|
||||
]
|
||||
assert output[3:5] == UnsortedSequence(
|
||||
{
|
||||
"tools": {
|
||||
"messages": [
|
||||
@@ -2698,6 +2713,12 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
name="search_api",
|
||||
tool_call_id="tool_call234",
|
||||
),
|
||||
]
|
||||
}
|
||||
},
|
||||
{
|
||||
"tools": {
|
||||
"messages": [
|
||||
_AnyIdToolMessage(
|
||||
content="result for a third one",
|
||||
name="search_api",
|
||||
@@ -2706,79 +2727,10 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
]
|
||||
}
|
||||
},
|
||||
{"agent": {"messages": [_AnyIdAIMessage(content="answer")]}},
|
||||
][0]["agent"]["messages"]
|
||||
)
|
||||
|
||||
assert [
|
||||
*app.stream({"messages": [HumanMessage(content="what is weather in sf")]})
|
||||
] == [
|
||||
{
|
||||
"agent": {
|
||||
"messages": [
|
||||
_AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call123",
|
||||
"name": "search_api",
|
||||
"args": {"query": "query"},
|
||||
},
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
{
|
||||
"tools": {
|
||||
"messages": [
|
||||
_AnyIdToolMessage(
|
||||
content="result for query",
|
||||
name="search_api",
|
||||
tool_call_id="tool_call123",
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
{
|
||||
"agent": {
|
||||
"messages": [
|
||||
_AnyIdAIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "tool_call234",
|
||||
"name": "search_api",
|
||||
"args": {"query": "another"},
|
||||
},
|
||||
{
|
||||
"id": "tool_call567",
|
||||
"name": "search_api",
|
||||
"args": {"query": "a third one"},
|
||||
},
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
{
|
||||
"tools": {
|
||||
"messages": [
|
||||
_AnyIdToolMessage(
|
||||
content="result for another",
|
||||
name="search_api",
|
||||
tool_call_id="tool_call234",
|
||||
),
|
||||
_AnyIdToolMessage(
|
||||
content="result for a third one",
|
||||
name="search_api",
|
||||
tool_call_id="tool_call567",
|
||||
),
|
||||
]
|
||||
}
|
||||
},
|
||||
{"agent": {"messages": [_AnyIdAIMessage(content="answer")]}},
|
||||
]
|
||||
)
|
||||
assert output[5:] == [
|
||||
{"agent": {"messages": [_AnyIdAIMessage(content="answer")]}}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
|
||||
@@ -35,7 +35,8 @@ from langgraph.prebuilt.tool_node import ToolNode
|
||||
from langgraph.pregel import Channel, Pregel
|
||||
from langgraph.store.memory import InMemoryStore
|
||||
from langgraph.types import PregelTask, Send, StateSnapshot, StreamWriter
|
||||
from tests.any_str import AnyDict, AnyStr
|
||||
from tests.any_int import AnyInt
|
||||
from tests.any_str import AnyDict, AnyStr, UnsortedSequence
|
||||
from tests.conftest import (
|
||||
ALL_CHECKPOINTERS_ASYNC,
|
||||
REGULAR_CHECKPOINTERS_ASYNC,
|
||||
@@ -2297,13 +2298,15 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
]
|
||||
}
|
||||
|
||||
assert [
|
||||
events = [
|
||||
c
|
||||
async for c in app.astream(
|
||||
{"messages": [HumanMessage(content="what is weather in sf")]},
|
||||
stream_mode="messages",
|
||||
)
|
||||
] == [
|
||||
]
|
||||
|
||||
assert events[:3] == [
|
||||
(
|
||||
_AnyIdAIMessageChunk(
|
||||
content="",
|
||||
@@ -2329,7 +2332,7 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
"langgraph_step": 1,
|
||||
"langgraph_node": "agent",
|
||||
"langgraph_triggers": ("branch:to:agent",),
|
||||
"langgraph_path": ("__pregel_pull", "agent"),
|
||||
"langgraph_path": (PULL, "agent"),
|
||||
"langgraph_checkpoint_ns": AnyStr("agent:"),
|
||||
"checkpoint_ns": AnyStr("agent:"),
|
||||
"ls_provider": "fakechatmodel",
|
||||
@@ -2345,8 +2348,8 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
{
|
||||
"langgraph_step": 2,
|
||||
"langgraph_node": "tools",
|
||||
"langgraph_triggers": ("branch:to:tools",),
|
||||
"langgraph_path": ("__pregel_pull", "tools"),
|
||||
"langgraph_triggers": (PUSH,),
|
||||
"langgraph_path": (PUSH, AnyInt(), False),
|
||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||
},
|
||||
),
|
||||
@@ -2388,13 +2391,16 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
"langgraph_step": 3,
|
||||
"langgraph_node": "agent",
|
||||
"langgraph_triggers": ("branch:to:agent",),
|
||||
"langgraph_path": ("__pregel_pull", "agent"),
|
||||
"langgraph_path": (PULL, "agent"),
|
||||
"langgraph_checkpoint_ns": AnyStr("agent:"),
|
||||
"checkpoint_ns": AnyStr("agent:"),
|
||||
"ls_provider": "fakechatmodel",
|
||||
"ls_model_type": "chat",
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
assert events[3:5] == UnsortedSequence(
|
||||
(
|
||||
_AnyIdToolMessage(
|
||||
content="result for another",
|
||||
@@ -2404,8 +2410,8 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
{
|
||||
"langgraph_step": 4,
|
||||
"langgraph_node": "tools",
|
||||
"langgraph_triggers": ("branch:to:tools",),
|
||||
"langgraph_path": ("__pregel_pull", "tools"),
|
||||
"langgraph_triggers": (PUSH,),
|
||||
"langgraph_path": (PUSH, AnyInt(), False),
|
||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||
},
|
||||
),
|
||||
@@ -2418,11 +2424,13 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
{
|
||||
"langgraph_step": 4,
|
||||
"langgraph_node": "tools",
|
||||
"langgraph_triggers": ("branch:to:tools",),
|
||||
"langgraph_path": ("__pregel_pull", "tools"),
|
||||
"langgraph_triggers": (PUSH,),
|
||||
"langgraph_path": (PUSH, AnyInt(), False),
|
||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||
},
|
||||
),
|
||||
)
|
||||
assert events[5:] == [
|
||||
(
|
||||
_AnyIdAIMessageChunk(
|
||||
content="answer",
|
||||
@@ -2431,7 +2439,7 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
"langgraph_step": 5,
|
||||
"langgraph_node": "agent",
|
||||
"langgraph_triggers": ("branch:to:agent",),
|
||||
"langgraph_path": ("__pregel_pull", "agent"),
|
||||
"langgraph_path": (PULL, "agent"),
|
||||
"langgraph_checkpoint_ns": AnyStr("agent:"),
|
||||
"checkpoint_ns": AnyStr("agent:"),
|
||||
"ls_provider": "fakechatmodel",
|
||||
@@ -2440,12 +2448,13 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
),
|
||||
]
|
||||
|
||||
assert [
|
||||
stream_updates_events = [
|
||||
c
|
||||
async for c in app.astream(
|
||||
{"messages": [HumanMessage(content="what is weather in sf")]}
|
||||
)
|
||||
] == [
|
||||
]
|
||||
assert stream_updates_events[:3] == [
|
||||
{
|
||||
"agent": {
|
||||
"messages": [
|
||||
@@ -2494,6 +2503,8 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
]
|
||||
}
|
||||
},
|
||||
]
|
||||
assert stream_updates_events[3:5] == UnsortedSequence(
|
||||
{
|
||||
"tools": {
|
||||
"messages": [
|
||||
@@ -2502,6 +2513,12 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
name="search_api",
|
||||
tool_call_id="tool_call234",
|
||||
),
|
||||
]
|
||||
}
|
||||
},
|
||||
{
|
||||
"tools": {
|
||||
"messages": [
|
||||
_AnyIdToolMessage(
|
||||
content="result for a third one",
|
||||
name="search_api",
|
||||
@@ -2510,7 +2527,9 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
]
|
||||
}
|
||||
},
|
||||
{"agent": {"messages": [_AnyIdAIMessage(content="answer")]}},
|
||||
)
|
||||
assert stream_updates_events[5:] == [
|
||||
{"agent": {"messages": [_AnyIdAIMessage(content="answer")]}}
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -277,7 +277,7 @@ def create_react_agent(
|
||||
interrupt_before: Optional[list[str]] = None,
|
||||
interrupt_after: Optional[list[str]] = None,
|
||||
debug: bool = False,
|
||||
version: Literal["v1", "v2"] = "v1",
|
||||
version: Literal["v1", "v2"] = "v2",
|
||||
name: Optional[str] = None,
|
||||
) -> CompiledGraph:
|
||||
"""Creates an agent graph that calls tools in a loop until a stopping condition is met.
|
||||
@@ -701,6 +701,13 @@ def create_react_agent(
|
||||
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:
|
||||
|
||||
@@ -1078,7 +1078,9 @@ async def test_return_direct(version: str) -> None:
|
||||
),
|
||||
]
|
||||
model = FakeToolCallingModel(tool_calls=[second_tool_call, []])
|
||||
agent = create_react_agent(model, [tool_return_direct, tool_normal])
|
||||
agent = create_react_agent(
|
||||
model, [tool_return_direct, tool_normal], version=version
|
||||
)
|
||||
result = agent.invoke(
|
||||
{"messages": [HumanMessage(content="Test normal", id="hum1")]}
|
||||
)
|
||||
@@ -1107,7 +1109,9 @@ async def test_return_direct(version: str) -> None:
|
||||
),
|
||||
]
|
||||
model = FakeToolCallingModel(tool_calls=[both_tool_calls, []])
|
||||
agent = create_react_agent(model, [tool_return_direct, tool_normal])
|
||||
agent = create_react_agent(
|
||||
model, [tool_return_direct, tool_normal], version=version
|
||||
)
|
||||
result = agent.invoke({"messages": [HumanMessage(content="Test both", id="hum2")]})
|
||||
assert result["messages"] == [
|
||||
HumanMessage(content="Test both", id="hum2"),
|
||||
|
||||
Reference in New Issue
Block a user