mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-20 16:47:55 +02:00
Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f6d8e63e21 | ||
|
|
a4160c8f37 | ||
|
|
dcd325e581 | ||
|
|
ab478cb40f | ||
|
|
218887d7ce |
@@ -135,18 +135,25 @@ class ToolCallRequest:
|
|||||||
|
|
||||||
Attributes:
|
Attributes:
|
||||||
tool_call: Tool call dict with name, args, and id from model output.
|
tool_call: Tool call dict with name, args, and id from model output.
|
||||||
|
|
||||||
|
If an interceptor edits `tool_call["name"]` so it differs from `tool`,
|
||||||
|
`tool_call["name"]` is authoritative for what tool is executed.
|
||||||
tool: BaseTool instance to be invoked, or None if tool is not
|
tool: BaseTool instance to be invoked, or None if tool is not
|
||||||
registered with the `ToolNode`. When tool is `None`, interceptors can
|
registered with the `ToolNode`. When tool is `None`, interceptors can
|
||||||
handle the request without validation. If the interceptor calls `execute()`,
|
handle the request without validation. If the interceptor calls `execute()`,
|
||||||
validation will occur and raise an error for unregistered tools.
|
validation will occur and raise an error for unregistered tools.
|
||||||
state: Agent state (`dict`, `list`, or `BaseModel`).
|
state: Agent state (`dict`, `list`, or `BaseModel`).
|
||||||
runtime: LangGraph runtime context (optional, `None` if outside graph).
|
runtime: LangGraph runtime context (optional, `None` if outside graph).
|
||||||
|
available_tools: Client-side tools registered with the `ToolNode`. Provider
|
||||||
|
and built-in tools are not included. Use this to resolve a replacement
|
||||||
|
tool when redirecting a call, and set `tool` to the resolved instance.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
tool_call: ToolCall
|
tool_call: ToolCall
|
||||||
tool: BaseTool | None
|
tool: BaseTool | None
|
||||||
state: Any
|
state: Any
|
||||||
runtime: ToolRuntime
|
runtime: ToolRuntime
|
||||||
|
available_tools: list[BaseTool] = field(default_factory=list)
|
||||||
|
|
||||||
def __setattr__(self, name: str, value: Any) -> None:
|
def __setattr__(self, name: str, value: Any) -> None:
|
||||||
"""Raise deprecation warning when setting attributes directly.
|
"""Raise deprecation warning when setting attributes directly.
|
||||||
@@ -336,6 +343,27 @@ def msg_content_output(output: Any) -> str | list[dict]:
|
|||||||
return str(output)
|
return str(output)
|
||||||
|
|
||||||
|
|
||||||
|
class ToolCallRequestMismatchError(ValueError):
|
||||||
|
"""`tool_call["name"]` and `tool` disagree on a `ToolCallRequest`."""
|
||||||
|
|
||||||
|
|
||||||
|
def _check_not_redirected_without_tool(
|
||||||
|
request: ToolCallRequest, original_name: str, original_tool: BaseTool | None
|
||||||
|
) -> None:
|
||||||
|
"""Raise if an interceptor renamed the call but left `tool` as the resolved one."""
|
||||||
|
if (
|
||||||
|
original_tool is not None
|
||||||
|
and request.tool is original_tool
|
||||||
|
and request.tool_call["name"] != original_name
|
||||||
|
):
|
||||||
|
msg = (
|
||||||
|
f"Interceptor set tool_call name to {request.tool_call['name']!r} but left "
|
||||||
|
f"`tool` as {original_tool.name!r}. Redirecting a call requires setting both; "
|
||||||
|
f"resolve the replacement from `ToolCallRequest.available_tools`."
|
||||||
|
)
|
||||||
|
raise ToolCallRequestMismatchError(msg)
|
||||||
|
|
||||||
|
|
||||||
class ToolInvocationError(ToolException):
|
class ToolInvocationError(ToolException):
|
||||||
"""An error occurred while invoking a tool due to invalid arguments.
|
"""An error occurred while invoking a tool due to invalid arguments.
|
||||||
|
|
||||||
@@ -1037,6 +1065,7 @@ class ToolNode(RunnableCallable):
|
|||||||
tool=tool,
|
tool=tool,
|
||||||
state=tool_runtime.state,
|
state=tool_runtime.state,
|
||||||
runtime=tool_runtime,
|
runtime=tool_runtime,
|
||||||
|
available_tools=list(self.tools_by_name.values()),
|
||||||
)
|
)
|
||||||
|
|
||||||
config = tool_runtime.config
|
config = tool_runtime.config
|
||||||
@@ -1046,13 +1075,18 @@ class ToolNode(RunnableCallable):
|
|||||||
return self._execute_tool_sync(tool_request, input_type, config)
|
return self._execute_tool_sync(tool_request, input_type, config)
|
||||||
|
|
||||||
# Define execute callable that can be called multiple times
|
# Define execute callable that can be called multiple times
|
||||||
|
original_name, original_tool = call["name"], tool
|
||||||
|
|
||||||
def execute(req: ToolCallRequest) -> ToolMessage | Command:
|
def execute(req: ToolCallRequest) -> ToolMessage | Command:
|
||||||
"""Execute tool with given request. Can be called multiple times."""
|
"""Execute tool with given request. Can be called multiple times."""
|
||||||
|
_check_not_redirected_without_tool(req, original_name, original_tool)
|
||||||
return self._execute_tool_sync(req, input_type, config)
|
return self._execute_tool_sync(req, input_type, config)
|
||||||
|
|
||||||
# Call wrapper with request and execute callable
|
# Call wrapper with request and execute callable
|
||||||
try:
|
try:
|
||||||
return self._wrap_tool_call(tool_request, execute)
|
return self._wrap_tool_call(tool_request, execute)
|
||||||
|
except ToolCallRequestMismatchError:
|
||||||
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# Wrapper threw an exception
|
# Wrapper threw an exception
|
||||||
if not self._handle_tool_errors:
|
if not self._handle_tool_errors:
|
||||||
@@ -1184,6 +1218,7 @@ class ToolNode(RunnableCallable):
|
|||||||
tool=tool,
|
tool=tool,
|
||||||
state=tool_runtime.state,
|
state=tool_runtime.state,
|
||||||
runtime=tool_runtime,
|
runtime=tool_runtime,
|
||||||
|
available_tools=list(self.tools_by_name.values()),
|
||||||
)
|
)
|
||||||
|
|
||||||
config = tool_runtime.config
|
config = tool_runtime.config
|
||||||
@@ -1193,12 +1228,16 @@ class ToolNode(RunnableCallable):
|
|||||||
return await self._execute_tool_async(tool_request, input_type, config)
|
return await self._execute_tool_async(tool_request, input_type, config)
|
||||||
|
|
||||||
# Define async execute callable that can be called multiple times
|
# Define async execute callable that can be called multiple times
|
||||||
|
original_name, original_tool = call["name"], tool
|
||||||
|
|
||||||
async def execute(req: ToolCallRequest) -> ToolMessage | Command:
|
async def execute(req: ToolCallRequest) -> ToolMessage | Command:
|
||||||
"""Execute tool with given request. Can be called multiple times."""
|
"""Execute tool with given request. Can be called multiple times."""
|
||||||
|
_check_not_redirected_without_tool(req, original_name, original_tool)
|
||||||
return await self._execute_tool_async(req, input_type, config)
|
return await self._execute_tool_async(req, input_type, config)
|
||||||
|
|
||||||
def _sync_execute(req: ToolCallRequest) -> ToolMessage | Command:
|
def _sync_execute(req: ToolCallRequest) -> ToolMessage | Command:
|
||||||
"""Sync execute fallback for sync wrapper."""
|
"""Sync execute fallback for sync wrapper."""
|
||||||
|
_check_not_redirected_without_tool(req, original_name, original_tool)
|
||||||
return self._execute_tool_sync(req, input_type, config)
|
return self._execute_tool_sync(req, input_type, config)
|
||||||
|
|
||||||
# Call wrapper with request and execute callable
|
# Call wrapper with request and execute callable
|
||||||
@@ -1208,6 +1247,8 @@ class ToolNode(RunnableCallable):
|
|||||||
# None check was performed above already
|
# None check was performed above already
|
||||||
self._wrap_tool_call = cast("ToolCallWrapper", self._wrap_tool_call)
|
self._wrap_tool_call = cast("ToolCallWrapper", self._wrap_tool_call)
|
||||||
return self._wrap_tool_call(tool_request, _sync_execute)
|
return self._wrap_tool_call(tool_request, _sync_execute)
|
||||||
|
except ToolCallRequestMismatchError:
|
||||||
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# Wrapper threw an exception
|
# Wrapper threw an exception
|
||||||
if not self._handle_tool_errors:
|
if not self._handle_tool_errors:
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
"""Unit tests for tool call interceptor in ToolNode."""
|
"""Unit tests for tool call interceptor in ToolNode."""
|
||||||
|
|
||||||
import functools
|
import functools
|
||||||
from collections.abc import Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from unittest.mock import Mock
|
from unittest.mock import Mock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -13,6 +13,7 @@ from langgraph.types import Command
|
|||||||
|
|
||||||
from langgraph.prebuilt.tool_node import (
|
from langgraph.prebuilt.tool_node import (
|
||||||
ToolCallRequest,
|
ToolCallRequest,
|
||||||
|
ToolCallRequestMismatchError,
|
||||||
ToolNode,
|
ToolNode,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1471,3 +1472,118 @@ def test_tool_call_request_is_frozen() -> None:
|
|||||||
assert fresh_new_request.tool == add # Other fields should remain the same
|
assert fresh_new_request.tool == add # Other fields should remain the same
|
||||||
assert fresh_new_request.state == state
|
assert fresh_new_request.state == state
|
||||||
assert fresh_new_request.runtime is None
|
assert fresh_new_request.runtime is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_interceptor_can_redirect_to_another_tool() -> None:
|
||||||
|
"""Redirecting requires setting both `tool_call` and `tool`; routing follows them."""
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def subtract(a: int, b: int) -> int:
|
||||||
|
"""Subtract two numbers."""
|
||||||
|
return a - b
|
||||||
|
|
||||||
|
async def redirect(
|
||||||
|
request: ToolCallRequest,
|
||||||
|
execute: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]],
|
||||||
|
) -> ToolMessage | Command:
|
||||||
|
target = next(t for t in request.available_tools if t.name == "subtract")
|
||||||
|
return await execute(
|
||||||
|
request.override(
|
||||||
|
tool_call={**request.tool_call, "name": "subtract"}, tool=target
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
node = ToolNode([add, subtract], awrap_tool_call=redirect)
|
||||||
|
result = await node.ainvoke(
|
||||||
|
[
|
||||||
|
AIMessage(
|
||||||
|
"",
|
||||||
|
tool_calls=[
|
||||||
|
{
|
||||||
|
"name": "add",
|
||||||
|
"args": {"a": 5, "b": 3},
|
||||||
|
"id": "1",
|
||||||
|
"type": "tool_call",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
_create_config_with_runtime(),
|
||||||
|
)
|
||||||
|
|
||||||
|
# `add` would return 8; `subtract` returns 2.
|
||||||
|
assert result[0].content == "2"
|
||||||
|
|
||||||
|
|
||||||
|
def test_interceptor_tool_call_name_and_tool_must_agree() -> None:
|
||||||
|
"""Renaming `tool_call` without `tool` raises rather than running the wrong tool."""
|
||||||
|
|
||||||
|
def rename_only(
|
||||||
|
request: ToolCallRequest,
|
||||||
|
execute: Callable[[ToolCallRequest], ToolMessage | Command],
|
||||||
|
) -> ToolMessage | Command:
|
||||||
|
return execute(
|
||||||
|
request.override(tool_call={**request.tool_call, "name": "other"})
|
||||||
|
)
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def other(a: int, b: int) -> int:
|
||||||
|
"""Another tool."""
|
||||||
|
return 0
|
||||||
|
|
||||||
|
# handle_tool_errors is on by default; the mismatch must not become a ToolMessage
|
||||||
|
node = ToolNode([add, other], wrap_tool_call=rename_only)
|
||||||
|
with pytest.raises(ToolCallRequestMismatchError, match="other"):
|
||||||
|
node.invoke(
|
||||||
|
[
|
||||||
|
AIMessage(
|
||||||
|
"",
|
||||||
|
tool_calls=[
|
||||||
|
{
|
||||||
|
"name": "add",
|
||||||
|
"args": {"a": 1, "b": 2},
|
||||||
|
"id": "1",
|
||||||
|
"type": "tool_call",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
_create_config_with_runtime(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_sync_interceptor_under_ainvoke_also_validates_redirect() -> None:
|
||||||
|
"""The sync-wrapper fallback used by `ainvoke` must validate too, not just `invoke`."""
|
||||||
|
|
||||||
|
def rename_only(
|
||||||
|
request: ToolCallRequest,
|
||||||
|
execute: Callable[[ToolCallRequest], ToolMessage | Command],
|
||||||
|
) -> ToolMessage | Command:
|
||||||
|
return execute(
|
||||||
|
request.override(tool_call={**request.tool_call, "name": "other"})
|
||||||
|
)
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def other(a: int, b: int) -> int:
|
||||||
|
"""Another tool."""
|
||||||
|
return 0
|
||||||
|
|
||||||
|
# Only a sync wrapper is configured, so `ainvoke` routes through `_sync_execute`.
|
||||||
|
node = ToolNode([add, other], wrap_tool_call=rename_only)
|
||||||
|
with pytest.raises(ToolCallRequestMismatchError, match="other"):
|
||||||
|
await node.ainvoke(
|
||||||
|
[
|
||||||
|
AIMessage(
|
||||||
|
"",
|
||||||
|
tool_calls=[
|
||||||
|
{
|
||||||
|
"name": "add",
|
||||||
|
"args": {"a": 1, "b": 2},
|
||||||
|
"id": "1",
|
||||||
|
"type": "tool_call",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
],
|
||||||
|
_create_config_with_runtime(),
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user