feat(prebuilt): Add structured output tools to ToolNode (#5899)

* Add structured output tools to ToolNode
* Fix default tool node name to match the actual default ('tools')
* Update doc-strings to explain what inputs/outputs are for the ToolNode.
* Mark internal attributes as private (potentially breaking -- although hopefully users aren't accessing these)


## Decisions points

* OK with two properties? Done since users may be relying on
`tools_by_name` and expanding the return type will break user code.

## Changes in public/private interface

### Marked as public

* Make `tools_by_name` an official public property
* Make `structured_output_tools` a public property

### Marked as private

There should be no reason why users are accessing these attributes

```python
_tool_to_state_args
_tool_to_store_arg
_handle_tool_errors
_messages_key
```


### Usage

```python

    class OutputSchema(BaseModel):
        name: str
        age: int
        location: str

    tool_node = ToolNode([OutputSchema])

    # Test that the structured output tool is registered correctly
    assert "OutputSchema" in tool_node.structured_output_tools

    # Create a tool call that matches the schema
    tool_call = {
        "name": "OutputSchema",
        "args": {"name": "Alice", "age": 30, "location": "NYC"},
        "id": "call_123",
        "type": "tool_call",
    }

    # Test sync execution
    result = tool_node.invoke(
        {"messages": [AIMessage(content="", tool_calls=[tool_call])]}
    )

    # Should return a Command with structured response
    assert isinstance(result, list)
    assert len(result) == 1
    command = result[0]
    assert isinstance(command, Command)

    # Check the update structure
    assert "messages" in command.update
    assert "structured_response" in command.update

    # Check the tool message
    tool_message = command.update["messages"][0]
    assert isinstance(tool_message, ToolMessage)
    assert tool_message.name == "OutputSchema"
    assert tool_message.tool_call_id == "call_123"

    # Check the structured response
    structured_response = command.update["structured_response"]
    assert isinstance(structured_response, OutputSchema)
    assert structured_response.name == "Alice"
    assert structured_response.age == 30
    assert structured_response.location == "NYC"
```
This commit is contained in:
Eugene Yurtsev
2025-08-13 15:16:53 -04:00
committed by GitHub
parent 9e174e7e8b
commit 50601dc02c
2 changed files with 273 additions and 85 deletions
+177 -81
View File
@@ -31,6 +31,8 @@ Typical Usage:
```
"""
from __future__ import annotations
import asyncio
import inspect
import json
@@ -237,17 +239,50 @@ def _infer_handled_types(handler: Callable[..., str]) -> tuple[type[Exception],
class ToolNode(RunnableCallable):
"""A node that runs the tools called in the last AIMessage.
"""A node for executing tools in LangGraph workflows.
It can be used either in StateGraph with a "messages" state key (or a custom key passed via ToolNode's 'messages_key').
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.
Handles tool execution patterns including function calls, state injection,
persistent storage, and control flow. Manages parallel execution,
error handling.
Tool calls can also be passed directly as a list of `ToolCall` dicts.
Input Formats:
1. Graph state with `messages` key that has a list of messages:
- Common representation for agentic workflows
- Supports custom messages key via ``messages_key`` parameter
2. **Message List**: ``[AIMessage(..., tool_calls=[...])]``
- List of messages with tool calls in the last AIMessage
3. **Direct Tool Calls**: ``[{"name": "tool", "args": {...}, "id": "1", "type": "tool_call"}]``
- Bypasses message parsing for direct tool execution
- For programmatic tool invocation and testing
Tool Types:
1. **Regular tools**: Functions or BaseTool instances that return values or
Commands.
2. **Structured output tools**: Pydantic model classes for schema-validated
responses
Output Formats:
Output format depends on input type and tool behavior:
**For Regular tools**:
- Dict input → ``{"messages": [ToolMessage(...)]}``
- List input → ``[ToolMessage(...)]``
**For Command tools**:
- Returns ``[Command(...)]`` or mixed list with regular tool outputs
- Commands can update state, trigger navigation, or send messages
**For Structured output tools**:
- Returns ``[Command(update={"messages": [...], "structured_response": schema_instance})]``
- Includes both message and structured data in the graph state
Args:
tools: A sequence of tools that can be invoked by this node. Tools can be
BaseTool instances or plain functions that will be converted to tools.
tools: A sequence of tools that can be invoked by this node. Supports:
- **BaseTool instances**: Tools with schemas and metadata
- **Plain functions**: Automatically converted to tools with inferred schemas
- **Pydantic model classes**: Treated as structured output tools
name: The name identifier for this node in the graph. Used for debugging
and visualization. Defaults to "tools".
tags: Optional metadata tags to associate with the node for filtering
@@ -255,21 +290,24 @@ class ToolNode(RunnableCallable):
handle_tool_errors: Configuration for error handling during tool execution.
Defaults to True. Supports multiple strategies:
- True: Catch all errors and return a ToolMessage with the default
- **True**: Catch all errors and return a ToolMessage with the default
error template containing the exception details.
- str: Catch all errors and return a ToolMessage with this custom
- **str**: Catch all errors and return a ToolMessage with this custom
error message string.
- tuple[type[Exception], ...]: Only catch exceptions of the specified
- **tuple[type[Exception], ...]**: Only catch exceptions with the specified
types and return default error messages for them.
- Callable[..., str]: Catch exceptions matching the callable's signature
- **Callable[..., str]**: Catch exceptions matching the callable's signature
and return the string result of calling it with the exception.
- False: Disable error handling entirely, allowing exceptions to propagate.
- **False**: Disable error handling entirely, allowing exceptions to
propagate.
messages_key: The key in the state dictionary that contains the message list.
This same key will be used for the output ToolMessages. Defaults to "messages".
This same key will be used for the output ToolMessages.
Defaults to "messages".
Allows custom state schemas with different message field names.
Example:
Basic usage with simple tools:
Examples:
Basic usage:
```python
from langgraph.prebuilt import ToolNode
@@ -283,42 +321,35 @@ class ToolNode(RunnableCallable):
tool_node = ToolNode([calculator])
```
Custom error handling:
State injection:
```python
def handle_math_errors(e: ZeroDivisionError) -> str:
return "Cannot divide by zero!"
from typing_extensions import Annotated
from langgraph.prebuilt import InjectedState
tool_node = ToolNode([calculator], handle_tool_errors=handle_math_errors)
@tool
def context_tool(query: str, state: Annotated[dict, InjectedState]) -> str:
\"\"\"Some tool that uses state.\"\"\"
return f"Query: {query}, Messages: {len(state['messages'])}"
tool_node = ToolNode([context_tool])
```
Direct tool call execution:
Error handling:
```python
tool_calls = [{"name": "calculator", "args": {"a": 5, "b": 3}, "id": "1", "type": "tool_call"}]
result = tool_node.invoke(tool_calls)
def handle_errors(e: ValueError) -> str:
return "Invalid input provided"
tool_node = ToolNode([my_tool], handle_tool_errors=handle_errors)
```
Note:
The ToolNode expects input in one of three formats:
1. A dictionary with a messages key containing a list of messages
2. A list of messages directly
3. A list of tool call dictionaries
When using message formats, the last message must be an AIMessage with
tool_calls populated. The node automatically extracts and processes these
tool calls concurrently.
For advanced use cases involving state injection or store access, tools
can be annotated with InjectedState or InjectedStore to receive graph
context automatically.
"""
name: str = "ToolNode"
name: str = "tools"
def __init__(
self,
tools: Sequence[Union[BaseTool, Callable]],
tools: Sequence[Union[BaseTool, BaseModel, Callable]],
*,
name: str = "tools",
tags: Optional[list[str]] = None,
@@ -337,17 +368,36 @@ class ToolNode(RunnableCallable):
messages_key: State key containing messages.
"""
super().__init__(self._func, self._afunc, name=name, tags=tags, trace=False)
self.tools_by_name: dict[str, BaseTool] = {}
self.tool_to_state_args: dict[str, dict[str, Optional[str]]] = {}
self.tool_to_store_arg: dict[str, Optional[str]] = {}
self.handle_tool_errors = handle_tool_errors
self.messages_key = messages_key
for tool_ in tools:
if not isinstance(tool_, BaseTool):
tool_ = create_tool(tool_)
self.tools_by_name[tool_.name] = tool_
self.tool_to_state_args[tool_.name] = _get_state_args(tool_)
self.tool_to_store_arg[tool_.name] = _get_store_arg(tool_)
self._tools_by_name: dict[str, BaseTool] = {}
self._structured_output_tools_by_name: dict[str, type[BaseModel]] = {}
self._tool_to_state_args: dict[str, dict[str, Optional[str]]] = {}
self._tool_to_store_arg: dict[str, Optional[str]] = {}
self._handle_tool_errors = handle_tool_errors
self._messages_key = messages_key
for tool in tools:
if inspect.isclass(tool) and issubclass(tool, BaseModel):
# Handle Pydantic model classes as structured output tools
self._structured_output_tools_by_name[tool.__name__] = tool
self._tool_to_state_args[tool.__name__] = {}
self._tool_to_store_arg[tool.__name__] = None
else:
if not isinstance(tool, BaseTool):
tool_ = create_tool(cast(Type[BaseTool], tool))
else:
tool_ = tool
self._tools_by_name[tool_.name] = tool_
self._tool_to_state_args[tool_.name] = _get_state_args(tool_)
self._tool_to_store_arg[tool_.name] = _get_store_arg(tool_)
@property
def tools_by_name(self) -> dict[str, BaseTool]:
"""Mapping from tool name to BaseTool instance."""
return self._tools_by_name
@property
def structured_output_tools(self) -> dict[str, type[BaseModel]]:
"""Mapping from structured output tool name to Pydantic model class."""
return self._structured_output_tools_by_name
def _func(
self,
@@ -390,14 +440,14 @@ class ToolNode(RunnableCallable):
def _combine_tool_outputs(
self,
outputs: list[ToolMessage],
outputs: list[Union[ToolMessage, Command]],
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}
return outputs if input_type == "list" else {self._messages_key: outputs}
# LangGraph will automatically handle list of Command and non-command node
# updates
@@ -425,7 +475,7 @@ class ToolNode(RunnableCallable):
combined_outputs.append(output)
else:
combined_outputs.append(
[output] if input_type == "list" else {self.messages_key: [output]}
[output] if input_type == "list" else {self._messages_key: [output]}
)
if parent_command:
@@ -437,13 +487,31 @@ class ToolNode(RunnableCallable):
call: ToolCall,
input_type: Literal["list", "dict", "tool_calls"],
config: RunnableConfig,
) -> ToolMessage:
) -> Union[ToolMessage, Command]:
"""Run a single tool call synchronously."""
if invalid_tool_message := self._validate_tool_call(call):
return invalid_tool_message
# Handle structured output tools
if call["name"] in self.structured_output_tools:
response_schema = self._structured_output_tools_by_name[call["name"]]
return Command(
update={
"messages": [
ToolMessage(
content="ok!",
name=call["name"],
tool_call_id=call["id"],
)
],
"structured_response": response_schema(**call["args"]),
}
)
try:
call_args = {**call, **{"type": "tool_call"}}
response = self.tools_by_name[call["name"]].invoke(call_args, config)
tool = self.tools_by_name[call["name"]]
response = tool.invoke(call_args, config)
# GraphInterrupt is a special exception that will always be raised.
# It can be triggered in the following scenarios,
@@ -455,20 +523,20 @@ class ToolNode(RunnableCallable):
except GraphBubbleUp as e:
raise e
except Exception as e:
if isinstance(self.handle_tool_errors, tuple):
handled_types: tuple = self.handle_tool_errors
elif callable(self.handle_tool_errors):
handled_types = _infer_handled_types(self.handle_tool_errors)
if isinstance(self._handle_tool_errors, tuple):
handled_types: tuple = self._handle_tool_errors
elif callable(self._handle_tool_errors):
handled_types = _infer_handled_types(self._handle_tool_errors)
else:
# default behavior is catching all exceptions
handled_types = (Exception,)
# Unhandled
if not self.handle_tool_errors or not isinstance(e, handled_types):
if not self._handle_tool_errors or not isinstance(e, handled_types):
raise e
# Handled
else:
content = _handle_tool_error(e, flag=self.handle_tool_errors)
content = _handle_tool_error(e, flag=self._handle_tool_errors)
return ToolMessage(
content=content,
name=call["name"],
@@ -493,15 +561,31 @@ class ToolNode(RunnableCallable):
call: ToolCall,
input_type: Literal["list", "dict", "tool_calls"],
config: RunnableConfig,
) -> ToolMessage:
) -> Union[ToolMessage, Command]:
"""Run a single tool call asynchronously."""
if invalid_tool_message := self._validate_tool_call(call):
return invalid_tool_message
# Handle structured output tools
if call["name"] in self.structured_output_tools:
response_schema = self._structured_output_tools_by_name[call["name"]]
return Command(
update={
"messages": [
ToolMessage(
content="ok!",
name=call["name"],
tool_call_id=call["id"],
)
],
"structured_response": response_schema(**call["args"]),
}
)
try:
call_args = {**call, **{"type": "tool_call"}}
response = await self.tools_by_name[call["name"]].ainvoke(call_args, config)
tool = self.tools_by_name[call["name"]]
response = await tool.ainvoke(call_args, config)
# GraphInterrupt is a special exception that will always be raised.
# It can be triggered in the following scenarios,
# Where GraphInterrupt(GraphBubbleUp) is raised from an `interrupt` invocation most commonly:
@@ -512,20 +596,20 @@ class ToolNode(RunnableCallable):
except GraphBubbleUp as e:
raise e
except Exception as e:
if isinstance(self.handle_tool_errors, tuple):
handled_types: tuple = self.handle_tool_errors
elif callable(self.handle_tool_errors):
handled_types = _infer_handled_types(self.handle_tool_errors)
if isinstance(self._handle_tool_errors, tuple):
handled_types: tuple = self._handle_tool_errors
elif callable(self._handle_tool_errors):
handled_types = _infer_handled_types(self._handle_tool_errors)
else:
# default behavior is catching all exceptions
handled_types = (Exception,)
# Unhandled
if not self.handle_tool_errors or not isinstance(e, handled_types):
if not self._handle_tool_errors or not isinstance(e, handled_types):
raise e
# Handled
else:
content = _handle_tool_error(e, flag=self.handle_tool_errors)
content = _handle_tool_error(e, flag=self._handle_tool_errors)
return ToolMessage(
content=content,
@@ -564,9 +648,11 @@ class ToolNode(RunnableCallable):
else:
input_type = "list"
messages = input
elif isinstance(input, dict) and (messages := input.get(self.messages_key, [])):
elif isinstance(input, dict) and (
messages := input.get(self._messages_key, [])
):
input_type = "dict"
elif messages := getattr(input, self.messages_key, []):
elif messages := getattr(input, self._messages_key, []):
# Assume dataclass-like state that can coerce from dict
input_type = "dict"
else:
@@ -586,10 +672,17 @@ class ToolNode(RunnableCallable):
return tool_calls, input_type
def _validate_tool_call(self, call: ToolCall) -> Optional[ToolMessage]:
if (requested_tool := call["name"]) not in self.tools_by_name:
requested_tool = call["name"]
if (
requested_tool not in self.tools_by_name
and requested_tool not in self._structured_output_tools_by_name
):
all_tool_names = list(self.tools_by_name.keys()) + list(
self._structured_output_tools_by_name.keys()
)
content = INVALID_TOOL_NAME_ERROR_TEMPLATE.format(
requested_tool=requested_tool,
available_tools=", ".join(self.tools_by_name.keys()),
available_tools=", ".join(all_tool_names),
)
return ToolMessage(
content, name=requested_tool, tool_call_id=call["id"], status="error"
@@ -606,15 +699,15 @@ class ToolNode(RunnableCallable):
BaseModel,
],
) -> ToolCall:
state_args = self.tool_to_state_args[tool_call["name"]]
state_args = self._tool_to_state_args[tool_call["name"]]
if state_args and isinstance(input, list):
required_fields = list(state_args.values())
if (
len(required_fields) == 1
and required_fields[0] == self.messages_key
and required_fields[0] == self._messages_key
or required_fields[0] is None
):
input = {self.messages_key: input}
input = {self._messages_key: input}
else:
err_msg = (
f"Invalid input to ToolNode. Tool {tool_call['name']} requires "
@@ -645,7 +738,7 @@ class ToolNode(RunnableCallable):
def _inject_store(
self, tool_call: ToolCall, store: Optional[BaseStore]
) -> ToolCall:
store_arg = self.tool_to_store_arg[tool_call["name"]]
store_arg = self._tool_to_store_arg[tool_call["name"]]
if not store_arg:
return tool_call
@@ -704,7 +797,10 @@ class ToolNode(RunnableCallable):
The injection is performed on a copy of the tool call to avoid mutating
the original.
"""
if tool_call["name"] not in self.tools_by_name:
if (
tool_call["name"] not in self.tools_by_name
and tool_call["name"] not in self._structured_output_tools_by_name
):
return tool_call
tool_call_copy: ToolCall = copy(tool_call)
@@ -722,15 +818,15 @@ class ToolNode(RunnableCallable):
# input type is dict when ToolNode is invoked with a dict input (e.g. {"messages": [AIMessage(..., tool_calls=[...])]})
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"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']}'"
)
updated_command = deepcopy(command)
state_update = cast(dict[str, Any], updated_command.update) or {}
messages_update = state_update.get(self.messages_key, [])
messages_update = state_update.get(self._messages_key, [])
elif isinstance(command.update, list):
# input type is list when ToolNode is invoked with a list input (e.g. [AIMessage(..., tool_calls=[...])])
# Input type is list when ToolNode is invoked with a list input (e.g. [AIMessage(..., tool_calls=[...])])
if input_type != "list":
raise ValueError(
f"Tools can provide a list of messages in Command.update only when using list of messages as ToolNode input, "
+96 -4
View File
@@ -86,7 +86,8 @@ def tool5(some_val: int):
tool5.handle_tool_error = "foo"
async def test_tool_node():
async def test_tool_node() -> None:
"""Test tool node."""
result = ToolNode([tool1]).invoke(
{
"messages": [
@@ -178,7 +179,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",
@@ -219,7 +220,7 @@ async def test_tool_node_tool_call_input():
]
async def test_tool_node_error_handling():
async def test_tool_node_error_handling() -> None:
def handle_all(e: Union[ValueError, ToolException, ValidationError]):
return TOOL_CALL_ERROR_TEMPLATE.format(error=repr(e))
@@ -281,7 +282,7 @@ async def test_tool_node_error_handling():
assert result_error["messages"][2].tool_call_id == "another id"
async def test_tool_node_error_handling_callable():
async def test_tool_node_error_handling_callable() -> None:
def handle_value_error(e: ValueError):
return "Value error"
@@ -1481,3 +1482,94 @@ def test_tool_node_stream_writer() -> None:
},
),
]
def test_structured_output_tools_sync() -> None:
"""Test that ToolNode handles Pydantic model classes as structured output tools."""
class OutputSchema(BaseModel):
name: str
age: int
location: str
tool_node = ToolNode([OutputSchema])
# Test that the structured output tool is registered correctly
assert "OutputSchema" in tool_node.structured_output_tools
# Create a tool call that matches the schema
tool_call = {
"name": "OutputSchema",
"args": {"name": "Alice", "age": 30, "location": "NYC"},
"id": "call_123",
"type": "tool_call",
}
# Test sync execution
result = tool_node.invoke(
{"messages": [AIMessage(content="", tool_calls=[tool_call])]}
)
# Should return a Command with structured response
assert isinstance(result, list)
assert len(result) == 1
command = result[0]
assert isinstance(command, Command)
# Check the update structure
assert "messages" in command.update
assert "structured_response" in command.update
# Check the tool message
tool_message = command.update["messages"][0]
assert isinstance(tool_message, ToolMessage)
assert tool_message.name == "OutputSchema"
assert tool_message.tool_call_id == "call_123"
# Check the structured response
structured_response = command.update["structured_response"]
assert isinstance(structured_response, OutputSchema)
assert structured_response.name == "Alice"
assert structured_response.age == 30
assert structured_response.location == "NYC"
async def test_structured_output_tools_async() -> None:
"""Test that ToolNode handles Pydantic model classes as structured output tools."""
class OutputSchema(BaseModel):
name: str
age: int
location: str
tool_node = ToolNode([OutputSchema])
# Test that the structured output tool is registered correctly
assert "OutputSchema" not in tool_node.tools_by_name
assert "OutputSchema" in tool_node.structured_output_tools
# Create a tool call that matches the schema
tool_call = {
"name": "OutputSchema",
"args": {"name": "Alice", "age": 30, "location": "NYC"},
"id": "call_123",
"type": "tool_call",
}
# Test async execution
result_async = await tool_node.ainvoke(
{"messages": [AIMessage(content="", tool_calls=[tool_call])]}
)
# Should produce the same result
assert isinstance(result_async, list)
assert len(result_async) == 1
command_async = result_async[0]
assert isinstance(command_async, Command)
assert "structured_response" in command_async.update
structured_response_async = command_async.update["structured_response"]
assert isinstance(structured_response_async, OutputSchema)
assert structured_response_async.name == "Alice"
assert structured_response_async.age == 30
assert structured_response_async.location == "NYC"