mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-12 12:47:53 +02:00
chore: port tool node improvements back to langgraph (#6321)
namespace decisions ``` langgraph.prebuilt ├── ToolRuntime # new # all of the other stuff that was already there langgraph.prebuilt.tool_node ├── ToolNode ├── ToolCallRequest # new ├── ToolRuntime # new ├── InjectedState ├── InjectedStore ├── ToolCallWrapper ├── AsyncToolCallWrapper ├── tools_condition ``` ``` langchain.tools ├── ToolRuntime # now from langgraph.prebuilt ├── InjectedState # now from langgraph.prebuilt ├── InjectedStore # now from langgraph.prebuilt ├── ToolException ├── tool ├── BaseTool ├── InjectedToolArg ├── InjectedToolCallId ```
This commit is contained in:
@@ -1,39 +1,92 @@
|
||||
import contextlib
|
||||
import dataclasses
|
||||
import json
|
||||
import sys
|
||||
from functools import partial
|
||||
from typing import (
|
||||
Annotated,
|
||||
Any,
|
||||
NoReturn,
|
||||
TypeVar,
|
||||
)
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import (
|
||||
AIMessage,
|
||||
AnyMessage,
|
||||
HumanMessage,
|
||||
RemoveMessage,
|
||||
ToolCall,
|
||||
ToolMessage,
|
||||
)
|
||||
from langchain_core.runnables.config import RunnableConfig
|
||||
from langchain_core.tools import BaseTool, ToolException
|
||||
from langchain_core.tools import tool as dec_tool
|
||||
from langgraph.config import get_stream_writer
|
||||
from langgraph.errors import GraphBubbleUp, GraphInterrupt
|
||||
from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
||||
from langgraph.graph import START, MessagesState, StateGraph
|
||||
from langgraph.graph.message import REMOVE_ALL_MESSAGES, add_messages
|
||||
from langgraph.store.base import BaseStore
|
||||
from langgraph.store.memory import InMemoryStore
|
||||
from langgraph.types import Command, Send
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from pydantic.v1 import ValidationError as ValidationErrorV1
|
||||
from pydantic import BaseModel
|
||||
from pydantic.v1 import BaseModel as BaseModelV1
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.prebuilt import ToolNode
|
||||
from langgraph.prebuilt.tool_node import TOOL_CALL_ERROR_TEMPLATE
|
||||
from langgraph.prebuilt import (
|
||||
InjectedState,
|
||||
InjectedStore,
|
||||
ToolNode,
|
||||
)
|
||||
from langgraph.prebuilt.tool_node import (
|
||||
TOOL_CALL_ERROR_TEMPLATE,
|
||||
ToolInvocationError,
|
||||
tools_condition,
|
||||
)
|
||||
|
||||
from .messages import _AnyIdHumanMessage, _AnyIdToolMessage
|
||||
from .model import FakeToolCallingModel
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
def _create_mock_runtime(store: BaseStore | None = None) -> Mock:
|
||||
"""Create a mock Runtime object for testing ToolNode outside of graph context.
|
||||
|
||||
This helper is needed because ToolNode._func expects a Runtime parameter
|
||||
which is injected by RunnableCallable from config["configurable"]["__pregel_runtime"].
|
||||
When testing ToolNode directly (outside a graph), we need to provide this manually.
|
||||
"""
|
||||
mock_runtime = Mock()
|
||||
mock_runtime.store = store
|
||||
mock_runtime.context = None
|
||||
mock_runtime.stream_writer = lambda *args, **kwargs: None
|
||||
return mock_runtime
|
||||
|
||||
|
||||
def _create_config_with_runtime(store: BaseStore | None = None) -> RunnableConfig:
|
||||
"""Create a RunnableConfig with mock Runtime for testing ToolNode.
|
||||
|
||||
Returns:
|
||||
RunnableConfig with __pregel_runtime in configurable dict.
|
||||
"""
|
||||
return {"configurable": {"__pregel_runtime": _create_mock_runtime(store)}}
|
||||
|
||||
|
||||
def tool1(some_val: int, some_other_val: str) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
if some_val == 0:
|
||||
raise ValueError("Test error")
|
||||
msg = "Test error"
|
||||
raise ValueError(msg)
|
||||
return f"{some_val} - {some_other_val}"
|
||||
|
||||
|
||||
async def tool2(some_val: int, some_other_val: str) -> str:
|
||||
"""Tool 2 docstring."""
|
||||
if some_val == 0:
|
||||
raise ToolException("Test error")
|
||||
msg = "Test error"
|
||||
raise ToolException(msg)
|
||||
return f"tool2: {some_val} - {some_other_val}"
|
||||
|
||||
|
||||
@@ -53,15 +106,17 @@ async def tool4(some_val: int, some_other_val: str) -> str:
|
||||
|
||||
|
||||
@dec_tool
|
||||
def tool5(some_val: int):
|
||||
def tool5(some_val: int) -> NoReturn:
|
||||
"""Tool 5 docstring."""
|
||||
raise ToolException("Test error")
|
||||
msg = "Test error"
|
||||
raise ToolException(msg)
|
||||
|
||||
|
||||
tool5.handle_tool_error = "foo"
|
||||
|
||||
|
||||
async def test_tool_node():
|
||||
async def test_tool_node() -> None:
|
||||
"""Test tool node."""
|
||||
result = ToolNode([tool1]).invoke(
|
||||
{
|
||||
"messages": [
|
||||
@@ -76,7 +131,8 @@ async def test_tool_node():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
tool_message: ToolMessage = result["messages"][-1]
|
||||
@@ -98,7 +154,8 @@ async def test_tool_node():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
tool_message: ToolMessage = result2["messages"][-1]
|
||||
@@ -120,7 +177,8 @@ async def test_tool_node():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
tool_message: ToolMessage = result3["messages"][-1]
|
||||
assert tool_message.type == "tool"
|
||||
@@ -145,7 +203,8 @@ async def test_tool_node():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
tool_message: ToolMessage = result4["messages"][-1]
|
||||
assert tool_message.type == "tool"
|
||||
@@ -153,7 +212,7 @@ async def test_tool_node():
|
||||
assert tool_message.tool_call_id == "some 3"
|
||||
|
||||
|
||||
async def test_tool_node_tool_call_input():
|
||||
async def test_tool_node_tool_call_input() -> None:
|
||||
# Single tool call
|
||||
tool_call_1 = {
|
||||
"name": "tool1",
|
||||
@@ -161,7 +220,9 @@ async def test_tool_node_tool_call_input():
|
||||
"id": "some 0",
|
||||
"type": "tool_call",
|
||||
}
|
||||
result = ToolNode([tool1]).invoke([tool_call_1])
|
||||
result = ToolNode([tool1]).invoke(
|
||||
[tool_call_1], config=_create_config_with_runtime()
|
||||
)
|
||||
assert result["messages"] == [
|
||||
ToolMessage(content="1 - foo", tool_call_id="some 0", name="tool1"),
|
||||
]
|
||||
@@ -173,7 +234,9 @@ async def test_tool_node_tool_call_input():
|
||||
"id": "some 1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
result = ToolNode([tool1]).invoke([tool_call_1, tool_call_2])
|
||||
result = ToolNode([tool1]).invoke(
|
||||
[tool_call_1, tool_call_2], config=_create_config_with_runtime()
|
||||
)
|
||||
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"),
|
||||
@@ -182,7 +245,9 @@ async def test_tool_node_tool_call_input():
|
||||
# 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])
|
||||
result = ToolNode([tool1]).invoke(
|
||||
[tool_call_1, tool_call_3], config=_create_config_with_runtime()
|
||||
)
|
||||
assert result["messages"] == [
|
||||
ToolMessage(content="1 - foo", tool_call_id="some 0", name="tool1"),
|
||||
ToolMessage(
|
||||
@@ -194,8 +259,58 @@ async def test_tool_node_tool_call_input():
|
||||
]
|
||||
|
||||
|
||||
async def test_tool_node_error_handling():
|
||||
def handle_all(e: ValueError | ToolException | ValidationError):
|
||||
def test_tool_node_error_handling_default_invocation() -> None:
|
||||
tn = ToolNode([tool1])
|
||||
result = tn.invoke(
|
||||
{
|
||||
"messages": [
|
||||
AIMessage(
|
||||
"hi?",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "tool1",
|
||||
"args": {"invalid": 0, "args": "foo"},
|
||||
"id": "some id",
|
||||
},
|
||||
],
|
||||
)
|
||||
]
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
assert all(m.type == "tool" for m in result["messages"])
|
||||
assert all(m.status == "error" for m in result["messages"])
|
||||
assert (
|
||||
"Error invoking tool 'tool1' with kwargs {'invalid': 0, 'args': 'foo'} with error:\n"
|
||||
in result["messages"][0].content
|
||||
)
|
||||
|
||||
|
||||
def test_tool_node_error_handling_default_exception() -> None:
|
||||
tn = ToolNode([tool1])
|
||||
with pytest.raises(ValueError):
|
||||
tn.invoke(
|
||||
{
|
||||
"messages": [
|
||||
AIMessage(
|
||||
"hi?",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "tool1",
|
||||
"args": {"some_val": 0, "some_other_val": "foo"},
|
||||
"id": "some id",
|
||||
},
|
||||
],
|
||||
)
|
||||
]
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
|
||||
async def test_tool_node_error_handling() -> None:
|
||||
def handle_all(e: ValueError | ToolException | ToolInvocationError):
|
||||
return TOOL_CALL_ERROR_TEMPLATE.format(error=repr(e))
|
||||
|
||||
# test catching all exceptions, via:
|
||||
@@ -204,7 +319,7 @@ async def test_tool_node_error_handling():
|
||||
# - passing a callable with all exceptions in the signature
|
||||
for handle_tool_errors in (
|
||||
True,
|
||||
(ValueError, ToolException, ValidationError),
|
||||
(ValueError, ToolException, ToolInvocationError),
|
||||
handle_all,
|
||||
):
|
||||
result_error = await ToolNode(
|
||||
@@ -233,34 +348,33 @@ async def test_tool_node_error_handling():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
assert all(m.type == "tool" for m in result_error["messages"])
|
||||
assert all(m.status == "error" for m in result_error["messages"])
|
||||
assert (
|
||||
result_error["messages"][0].content
|
||||
== f"Error: {repr(ValueError('Test error'))}\n Please fix your mistakes."
|
||||
== f"Error: {ValueError('Test error')!r}\n Please fix your mistakes."
|
||||
)
|
||||
assert (
|
||||
result_error["messages"][1].content
|
||||
== f"Error: {repr(ToolException('Test error'))}\n Please fix your mistakes."
|
||||
)
|
||||
assert (
|
||||
"ValidationError" in result_error["messages"][2].content
|
||||
or "validation error" in result_error["messages"][2].content
|
||||
== f"Error: {ToolException('Test error')!r}\n Please fix your mistakes."
|
||||
)
|
||||
# Check that the validation error contains the field name
|
||||
assert "some_other_val" in result_error["messages"][2].content
|
||||
|
||||
assert result_error["messages"][0].tool_call_id == "some id"
|
||||
assert result_error["messages"][1].tool_call_id == "some other id"
|
||||
assert result_error["messages"][2].tool_call_id == "another id"
|
||||
|
||||
|
||||
async def test_tool_node_error_handling_callable():
|
||||
def handle_value_error(e: ValueError):
|
||||
async def test_tool_node_error_handling_callable() -> None:
|
||||
def handle_value_error(e: ValueError) -> str:
|
||||
return "Value error"
|
||||
|
||||
def handle_tool_exception(e: ToolException):
|
||||
def handle_tool_exception(e: ToolException) -> str:
|
||||
return "Tool exception"
|
||||
|
||||
for handle_tool_errors in ("Value error", handle_value_error):
|
||||
@@ -280,7 +394,8 @@ async def test_tool_node_error_handling_callable():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
tool_message: ToolMessage = result_error["messages"][-1]
|
||||
assert tool_message.type == "tool"
|
||||
@@ -313,7 +428,8 @@ async def test_tool_node_error_handling_callable():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert str(exc_info.value) == "Test error"
|
||||
|
||||
@@ -340,12 +456,13 @@ async def test_tool_node_error_handling_callable():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert str(exc_info.value) == "Test error"
|
||||
|
||||
|
||||
async def test_tool_node_handle_tool_errors_false():
|
||||
async def test_tool_node_handle_tool_errors_false() -> None:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
ToolNode([tool1], handle_tool_errors=False).invoke(
|
||||
{
|
||||
@@ -361,7 +478,8 @@ async def test_tool_node_handle_tool_errors_false():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
assert str(exc_info.value) == "Test error"
|
||||
@@ -381,13 +499,14 @@ async def test_tool_node_handle_tool_errors_false():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
assert str(exc_info.value) == "Test error"
|
||||
|
||||
# test validation errors get raised if handle_tool_errors is False
|
||||
with pytest.raises((ValidationError, ValidationErrorV1)):
|
||||
with pytest.raises(ToolInvocationError):
|
||||
ToolNode([tool1], handle_tool_errors=False).invoke(
|
||||
{
|
||||
"messages": [
|
||||
@@ -402,11 +521,12 @@ async def test_tool_node_handle_tool_errors_false():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
|
||||
def test_tool_node_individual_tool_error_handling():
|
||||
def test_tool_node_individual_tool_error_handling() -> None:
|
||||
# test error handling on individual tools (and that it overrides overall error handling!)
|
||||
result_individual_tool_error_handler = ToolNode(
|
||||
[tool5], handle_tool_errors="bar"
|
||||
@@ -424,7 +544,8 @@ def test_tool_node_individual_tool_error_handling():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
tool_message: ToolMessage = result_individual_tool_error_handler["messages"][-1]
|
||||
@@ -434,7 +555,7 @@ def test_tool_node_individual_tool_error_handling():
|
||||
assert tool_message.tool_call_id == "some 0"
|
||||
|
||||
|
||||
def test_tool_node_incorrect_tool_name():
|
||||
def test_tool_node_incorrect_tool_name() -> None:
|
||||
result_incorrect_name = ToolNode([tool1, tool2]).invoke(
|
||||
{
|
||||
"messages": [
|
||||
@@ -449,7 +570,8 @@ def test_tool_node_incorrect_tool_name():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
tool_message: ToolMessage = result_incorrect_name["messages"][-1]
|
||||
@@ -462,12 +584,13 @@ def test_tool_node_incorrect_tool_name():
|
||||
assert tool_message.tool_call_id == "some 0"
|
||||
|
||||
|
||||
def test_tool_node_node_interrupt():
|
||||
def test_tool_node_node_interrupt() -> None:
|
||||
def tool_interrupt(some_val: int) -> None:
|
||||
"""Tool docstring."""
|
||||
raise GraphBubbleUp("foo")
|
||||
msg = "foo"
|
||||
raise GraphBubbleUp(msg)
|
||||
|
||||
def handle(e: GraphInterrupt):
|
||||
def handle(e: GraphInterrupt) -> str:
|
||||
return "handled"
|
||||
|
||||
for handle_tool_errors in (True, (GraphBubbleUp,), "handled", handle, False):
|
||||
@@ -487,13 +610,14 @@ def test_tool_node_node_interrupt():
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert exc_info.value == "foo"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("input_type", ["dict", "tool_calls"])
|
||||
async def test_tool_node_command(input_type: str):
|
||||
async def test_tool_node_command(input_type: str) -> None:
|
||||
from langchain_core.tools.base import InjectedToolCallId
|
||||
|
||||
@dec_tool
|
||||
@@ -578,7 +702,9 @@ async def test_tool_node_command(input_type: str):
|
||||
input_ = {"messages": [AIMessage("", tool_calls=tool_calls)]}
|
||||
elif input_type == "tool_calls":
|
||||
input_ = tool_calls
|
||||
result = ToolNode([add, transfer_to_bob]).invoke(input_)
|
||||
result = ToolNode([add, transfer_to_bob]).invoke(
|
||||
input_, config=_create_config_with_runtime()
|
||||
)
|
||||
|
||||
assert result == [
|
||||
{
|
||||
@@ -616,7 +742,8 @@ async def test_tool_node_command(input_type: str):
|
||||
"", tool_calls=[{"args": {}, "id": "1", "name": tool.name}]
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert result == [
|
||||
Command(
|
||||
@@ -643,7 +770,8 @@ async def test_tool_node_command(input_type: str):
|
||||
"", tool_calls=[{"args": {}, "id": "1", "name": tool.name}]
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert result == [
|
||||
Command(
|
||||
@@ -673,7 +801,8 @@ async def test_tool_node_command(input_type: str):
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert result == [
|
||||
Command(
|
||||
@@ -724,7 +853,8 @@ async def test_tool_node_command(input_type: str):
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
# test validation (missing tool message in the update for current graph)
|
||||
@@ -743,7 +873,8 @@ async def test_tool_node_command(input_type: str):
|
||||
tool_calls=[{"args": {}, "id": "1", "name": "no_update_tool"}],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
# test validation (tool message with a wrong tool call ID)
|
||||
@@ -770,7 +901,8 @@ async def test_tool_node_command(input_type: str):
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
# test validation (missing tool message in the update for parent graph is OK)
|
||||
@@ -789,11 +921,12 @@ async def test_tool_node_command(input_type: str):
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
},
|
||||
config=_create_config_with_runtime(),
|
||||
) == [Command(update={"messages": []}, graph=Command.PARENT)]
|
||||
|
||||
|
||||
async def test_tool_node_command_list_input():
|
||||
async def test_tool_node_command_list_input() -> None:
|
||||
from langchain_core.tools.base import InjectedToolCallId
|
||||
|
||||
@dec_tool
|
||||
@@ -871,7 +1004,8 @@ async def test_tool_node_command_list_input():
|
||||
{"args": {}, "id": "2", "name": "transfer_to_bob"},
|
||||
],
|
||||
)
|
||||
]
|
||||
],
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
assert result == [
|
||||
@@ -900,7 +1034,8 @@ async def test_tool_node_command_list_input():
|
||||
# test sync tools
|
||||
for tool in [transfer_to_bob, custom_tool]:
|
||||
result = ToolNode([tool]).invoke(
|
||||
[AIMessage("", tool_calls=[{"args": {}, "id": "1", "name": tool.name}])]
|
||||
[AIMessage("", tool_calls=[{"args": {}, "id": "1", "name": tool.name}])],
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert result == [
|
||||
Command(
|
||||
@@ -919,7 +1054,8 @@ async def test_tool_node_command_list_input():
|
||||
# test async tools
|
||||
for tool in [async_transfer_to_bob, async_custom_tool]:
|
||||
result = await ToolNode([tool]).ainvoke(
|
||||
[AIMessage("", tool_calls=[{"args": {}, "id": "1", "name": tool.name}])]
|
||||
[AIMessage("", tool_calls=[{"args": {}, "id": "1", "name": tool.name}])],
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert result == [
|
||||
Command(
|
||||
@@ -945,7 +1081,8 @@ async def test_tool_node_command_list_input():
|
||||
{"args": {}, "id": "2", "name": "custom_transfer_to_bob"},
|
||||
],
|
||||
)
|
||||
]
|
||||
],
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert result == [
|
||||
Command(
|
||||
@@ -990,7 +1127,8 @@ async def test_tool_node_command_list_input():
|
||||
"",
|
||||
tool_calls=[{"args": {}, "id": "1", "name": "list_update_tool"}],
|
||||
)
|
||||
]
|
||||
],
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
# test validation (missing tool message in the update for current graph)
|
||||
@@ -1007,7 +1145,8 @@ async def test_tool_node_command_list_input():
|
||||
"",
|
||||
tool_calls=[{"args": {}, "id": "1", "name": "no_update_tool"}],
|
||||
)
|
||||
]
|
||||
],
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
# test validation (tool message with a wrong tool call ID)
|
||||
@@ -1026,7 +1165,8 @@ async def test_tool_node_command_list_input():
|
||||
{"args": {}, "id": "1", "name": "mismatching_tool_call_id_tool"}
|
||||
],
|
||||
)
|
||||
]
|
||||
],
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
# test validation (missing tool message in the update for parent graph is OK)
|
||||
@@ -1041,11 +1181,12 @@ async def test_tool_node_command_list_input():
|
||||
"",
|
||||
tool_calls=[{"args": {}, "id": "1", "name": "node_update_parent_tool"}],
|
||||
)
|
||||
]
|
||||
],
|
||||
config=_create_config_with_runtime(),
|
||||
) == [Command(update=[], graph=Command.PARENT)]
|
||||
|
||||
|
||||
def test_tool_node_parent_command_with_send():
|
||||
def test_tool_node_parent_command_with_send() -> None:
|
||||
from langchain_core.tools.base import InjectedToolCallId
|
||||
|
||||
@dec_tool
|
||||
@@ -1096,7 +1237,8 @@ def test_tool_node_parent_command_with_send():
|
||||
]
|
||||
|
||||
result = ToolNode([transfer_to_alice, transfer_to_bob]).invoke(
|
||||
[AIMessage("", tool_calls=tool_calls)]
|
||||
[AIMessage("", tool_calls=tool_calls)],
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
assert result == [
|
||||
@@ -1132,7 +1274,7 @@ def test_tool_node_parent_command_with_send():
|
||||
]
|
||||
|
||||
|
||||
async def test_tool_node_command_remove_all_messages():
|
||||
async def test_tool_node_command_remove_all_messages() -> None:
|
||||
from langchain_core.tools.base import InjectedToolCallId
|
||||
|
||||
@dec_tool
|
||||
@@ -1147,7 +1289,8 @@ async def test_tool_node_command_remove_all_messages():
|
||||
"id": "tool_call_123",
|
||||
}
|
||||
result = await tool_node.ainvoke(
|
||||
{"messages": [AIMessage(content="", tool_calls=[tool_call])]}
|
||||
{"messages": [AIMessage(content="", tool_calls=[tool_call])]},
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
|
||||
assert isinstance(result, list)
|
||||
@@ -1155,3 +1298,315 @@ async def test_tool_node_command_remove_all_messages():
|
||||
command = result[0]
|
||||
assert isinstance(command, Command)
|
||||
assert command.update == {"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES)]}
|
||||
|
||||
|
||||
class _InjectStateSchema(TypedDict):
|
||||
messages: list
|
||||
foo: str
|
||||
|
||||
|
||||
class _InjectedStatePydanticV2Schema(BaseModel):
|
||||
messages: list
|
||||
foo: str
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class _InjectedStateDataclassSchema:
|
||||
messages: list
|
||||
foo: str
|
||||
|
||||
|
||||
_INJECTED_STATE_SCHEMAS = [
|
||||
_InjectStateSchema,
|
||||
_InjectedStatePydanticV2Schema,
|
||||
_InjectedStateDataclassSchema,
|
||||
]
|
||||
|
||||
if sys.version_info < (3, 14):
|
||||
|
||||
class _InjectedStatePydanticSchema(BaseModelV1):
|
||||
messages: list
|
||||
foo: str
|
||||
|
||||
_INJECTED_STATE_SCHEMAS.append(_InjectedStatePydanticSchema)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("schema_", _INJECTED_STATE_SCHEMAS)
|
||||
def test_tool_node_inject_state(schema_: type[T]) -> None:
|
||||
def tool1(some_val: int, state: Annotated[T, InjectedState]) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
if isinstance(state, dict):
|
||||
return state["foo"]
|
||||
return state.foo
|
||||
|
||||
def tool2(some_val: int, state: Annotated[T, InjectedState()]) -> str:
|
||||
"""Tool 2 docstring."""
|
||||
if isinstance(state, dict):
|
||||
return state["foo"]
|
||||
return state.foo
|
||||
|
||||
def tool3(
|
||||
some_val: int,
|
||||
foo: Annotated[str, InjectedState("foo")],
|
||||
msgs: Annotated[list[AnyMessage], InjectedState("messages")],
|
||||
) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
return foo
|
||||
|
||||
def tool4(
|
||||
some_val: int, msgs: Annotated[list[AnyMessage], InjectedState("messages")]
|
||||
) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
return msgs[0].content
|
||||
|
||||
node = ToolNode([tool1, tool2, tool3, tool4], handle_tool_errors=True)
|
||||
for tool_name in ("tool1", "tool2", "tool3"):
|
||||
tool_call = {
|
||||
"name": tool_name,
|
||||
"args": {"some_val": 1},
|
||||
"id": "some 0",
|
||||
"type": "tool_call",
|
||||
}
|
||||
msg = AIMessage("hi?", tool_calls=[tool_call])
|
||||
result = node.invoke(
|
||||
schema_(messages=[msg], foo="bar"), config=_create_config_with_runtime()
|
||||
)
|
||||
tool_message = result["messages"][-1]
|
||||
assert tool_message.content == "bar", f"Failed for tool={tool_name}"
|
||||
|
||||
if tool_name == "tool3":
|
||||
failure_input = None
|
||||
with contextlib.suppress(Exception):
|
||||
failure_input = schema_(messages=[msg], notfoo="bar")
|
||||
if failure_input is not None:
|
||||
with pytest.raises(KeyError):
|
||||
node.invoke(failure_input, config=_create_config_with_runtime())
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
node.invoke([msg], config=_create_config_with_runtime())
|
||||
else:
|
||||
failure_input = None
|
||||
try:
|
||||
failure_input = schema_(messages=[msg], notfoo="bar")
|
||||
except Exception:
|
||||
# We'd get a validation error from pydantic state and wouldn't make it to the node
|
||||
# anyway
|
||||
pass
|
||||
if failure_input is not None:
|
||||
messages_ = node.invoke(
|
||||
failure_input, config=_create_config_with_runtime()
|
||||
)
|
||||
tool_message = messages_["messages"][-1]
|
||||
assert "KeyError" in tool_message.content
|
||||
tool_message = node.invoke([msg], config=_create_config_with_runtime())[
|
||||
-1
|
||||
]
|
||||
assert "KeyError" in tool_message.content
|
||||
|
||||
tool_call = {
|
||||
"name": "tool4",
|
||||
"args": {"some_val": 1},
|
||||
"id": "some 0",
|
||||
"type": "tool_call",
|
||||
}
|
||||
msg = AIMessage("hi?", tool_calls=[tool_call])
|
||||
result = node.invoke(
|
||||
schema_(messages=[msg], foo=""), config=_create_config_with_runtime()
|
||||
)
|
||||
tool_message = result["messages"][-1]
|
||||
assert tool_message.content == "hi?"
|
||||
|
||||
result = node.invoke([msg], config=_create_config_with_runtime())
|
||||
tool_message = result[-1]
|
||||
assert tool_message.content == "hi?"
|
||||
|
||||
|
||||
def test_tool_node_inject_store() -> None:
|
||||
store = InMemoryStore()
|
||||
namespace = ("test",)
|
||||
|
||||
def tool1(some_val: int, store: Annotated[BaseStore, InjectedStore()]) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
store_val = store.get(namespace, "test_key").value["foo"]
|
||||
return f"Some val: {some_val}, store val: {store_val}"
|
||||
|
||||
def tool2(some_val: int, store: Annotated[BaseStore, InjectedStore()]) -> str:
|
||||
"""Tool 2 docstring."""
|
||||
store_val = store.get(namespace, "test_key").value["foo"]
|
||||
return f"Some val: {some_val}, store val: {store_val}"
|
||||
|
||||
def tool3(
|
||||
some_val: int,
|
||||
bar: Annotated[str, InjectedState("bar")],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
) -> str:
|
||||
"""Tool 3 docstring."""
|
||||
store_val = store.get(namespace, "test_key").value["foo"]
|
||||
return f"Some val: {some_val}, store val: {store_val}, state val: {bar}"
|
||||
|
||||
node = ToolNode([tool1, tool2, tool3], handle_tool_errors=True)
|
||||
store.put(namespace, "test_key", {"foo": "bar"})
|
||||
|
||||
class State(MessagesState):
|
||||
bar: str
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("tools", node)
|
||||
builder.add_edge(START, "tools")
|
||||
graph = builder.compile(store=store)
|
||||
|
||||
for tool_name in ("tool1", "tool2"):
|
||||
tool_call = {
|
||||
"name": tool_name,
|
||||
"args": {"some_val": 1},
|
||||
"id": "some 0",
|
||||
"type": "tool_call",
|
||||
}
|
||||
msg = AIMessage("hi?", tool_calls=[tool_call])
|
||||
node_result = node.invoke(
|
||||
{"messages": [msg]}, config=_create_config_with_runtime(store=store)
|
||||
)
|
||||
graph_result = graph.invoke({"messages": [msg]})
|
||||
for result in (node_result, graph_result):
|
||||
result["messages"][-1]
|
||||
tool_message = result["messages"][-1]
|
||||
assert tool_message.content == "Some val: 1, store val: bar", (
|
||||
f"Failed for tool={tool_name}"
|
||||
)
|
||||
|
||||
tool_call = {
|
||||
"name": "tool3",
|
||||
"args": {"some_val": 1},
|
||||
"id": "some 0",
|
||||
"type": "tool_call",
|
||||
}
|
||||
msg = AIMessage("hi?", tool_calls=[tool_call])
|
||||
node_result = node.invoke(
|
||||
{"messages": [msg], "bar": "baz"},
|
||||
config=_create_config_with_runtime(store=store),
|
||||
)
|
||||
graph_result = graph.invoke({"messages": [msg], "bar": "baz"})
|
||||
for result in (node_result, graph_result):
|
||||
result["messages"][-1]
|
||||
tool_message = result["messages"][-1]
|
||||
assert tool_message.content == "Some val: 1, store val: bar, state val: baz", (
|
||||
f"Failed for tool={tool_name}"
|
||||
)
|
||||
|
||||
# test injected store without passing store to compiled graph
|
||||
failing_graph = builder.compile()
|
||||
with pytest.raises(ValueError):
|
||||
failing_graph.invoke({"messages": [msg], "bar": "baz"})
|
||||
|
||||
|
||||
def test_tool_node_ensure_utf8() -> None:
|
||||
@dec_tool
|
||||
def get_day_list(days: list[str]) -> list[str]:
|
||||
"""choose days"""
|
||||
return days
|
||||
|
||||
data = ["星期一", "水曜日", "목요일", "Friday"]
|
||||
tools = [get_day_list]
|
||||
tool_calls = [ToolCall(name=get_day_list.name, args={"days": data}, id="test_id")]
|
||||
outputs: list[ToolMessage] = ToolNode(tools).invoke(
|
||||
[AIMessage(content="", tool_calls=tool_calls)],
|
||||
config=_create_config_with_runtime(),
|
||||
)
|
||||
assert outputs[0].content == json.dumps(data, ensure_ascii=False)
|
||||
|
||||
|
||||
def test_tool_node_messages_key() -> None:
|
||||
@dec_tool
|
||||
def add(a: int, b: int) -> int:
|
||||
"""Adds a and b."""
|
||||
return a + b
|
||||
|
||||
model = FakeToolCallingModel(
|
||||
tool_calls=[[ToolCall(name=add.name, args={"a": 1, "b": 2}, id="test_id")]]
|
||||
)
|
||||
|
||||
class State(TypedDict):
|
||||
subgraph_messages: Annotated[list[AnyMessage], add_messages]
|
||||
|
||||
def call_model(state: State) -> dict[str, Any]:
|
||||
response = model.invoke(state["subgraph_messages"])
|
||||
model.tool_calls = []
|
||||
return {"subgraph_messages": response}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("agent", call_model)
|
||||
builder.add_node("tools", ToolNode([add], messages_key="subgraph_messages"))
|
||||
builder.add_conditional_edges(
|
||||
"agent", partial(tools_condition, messages_key="subgraph_messages")
|
||||
)
|
||||
builder.add_edge(START, "agent")
|
||||
builder.add_edge("tools", "agent")
|
||||
|
||||
graph = builder.compile()
|
||||
result = graph.invoke({"subgraph_messages": [HumanMessage(content="hi")]})
|
||||
assert result["subgraph_messages"] == [
|
||||
_AnyIdHumanMessage(content="hi"),
|
||||
AIMessage(
|
||||
content="hi",
|
||||
id="0",
|
||||
tool_calls=[ToolCall(name=add.name, args={"a": 1, "b": 2}, id="test_id")],
|
||||
),
|
||||
_AnyIdToolMessage(content="3", name=add.name, tool_call_id="test_id"),
|
||||
AIMessage(content="hi-hi-3", id="1"),
|
||||
]
|
||||
|
||||
|
||||
def test_tool_node_stream_writer() -> None:
|
||||
@dec_tool
|
||||
def streaming_tool(x: int) -> str:
|
||||
"""Do something with writer."""
|
||||
my_writer = get_stream_writer()
|
||||
for value in ["foo", "bar", "baz"]:
|
||||
my_writer({"custom_tool_value": value})
|
||||
|
||||
return x
|
||||
|
||||
tool_node = ToolNode([streaming_tool])
|
||||
graph = (
|
||||
StateGraph(MessagesState)
|
||||
.add_node("tools", tool_node)
|
||||
.add_edge(START, "tools")
|
||||
.compile()
|
||||
)
|
||||
|
||||
tool_call = {
|
||||
"name": "streaming_tool",
|
||||
"args": {"x": 1},
|
||||
"id": "1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
inputs = {
|
||||
"messages": [AIMessage("", tool_calls=[tool_call])],
|
||||
}
|
||||
|
||||
assert list(graph.stream(inputs, stream_mode="custom")) == [
|
||||
{"custom_tool_value": "foo"},
|
||||
{"custom_tool_value": "bar"},
|
||||
{"custom_tool_value": "baz"},
|
||||
]
|
||||
assert list(graph.stream(inputs, stream_mode=["custom", "updates"])) == [
|
||||
("custom", {"custom_tool_value": "foo"}),
|
||||
("custom", {"custom_tool_value": "bar"}),
|
||||
("custom", {"custom_tool_value": "baz"}),
|
||||
(
|
||||
"updates",
|
||||
{
|
||||
"tools": {
|
||||
"messages": [
|
||||
_AnyIdToolMessage(
|
||||
content="1",
|
||||
name="streaming_tool",
|
||||
tool_call_id="1",
|
||||
),
|
||||
],
|
||||
},
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user