From dcd325e581e028cdb61f49c6d02ccd207695f8f7 Mon Sep 17 00:00:00 2001 From: Chester Curme Date: Wed, 16 Sep 2026 10:43:28 -0400 Subject: [PATCH] add available_tools to ToolCallRequest --- libs/prebuilt/langgraph/prebuilt/tool_node.py | 38 ++++++++++++++----- libs/prebuilt/tests/test_on_tool_call.py | 38 ++++++++++++++++++- 2 files changed, 64 insertions(+), 12 deletions(-) diff --git a/libs/prebuilt/langgraph/prebuilt/tool_node.py b/libs/prebuilt/langgraph/prebuilt/tool_node.py index ea58982ad..8e5118005 100644 --- a/libs/prebuilt/langgraph/prebuilt/tool_node.py +++ b/libs/prebuilt/langgraph/prebuilt/tool_node.py @@ -144,12 +144,16 @@ class ToolCallRequest: validation will occur and raise an error for unregistered tools. state: Agent state (`dict`, `list`, or `BaseModel`). 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: BaseTool | None state: Any runtime: ToolRuntime + available_tools: list[BaseTool] = field(default_factory=list) def __setattr__(self, name: str, value: Any) -> None: """Raise deprecation warning when setting attributes directly. @@ -339,6 +343,10 @@ def msg_content_output(output: Any) -> str | list[dict]: return str(output) +class ToolCallRequestMismatchError(ValueError): + """`tool_call["name"]` and `tool` disagree on a `ToolCallRequest`.""" + + class ToolInvocationError(ToolException): """An error occurred while invoking a tool due to invalid arguments. @@ -952,6 +960,13 @@ class ToolNode(RunnableCallable): msg = f"Tool {call['name']} is not registered with ToolNode" raise TypeError(msg) + if tool.name != call["name"]: + msg = ( + f"ToolCallRequest names tool {call['name']!r} but carries {tool.name!r}. " + f"An interceptor redirecting a call must set both `tool_call` and `tool`." + ) + raise ToolCallRequestMismatchError(msg) + # Inject state, store, and runtime right before invocation injected_call = self._inject_tool_args(call, request.runtime, tool) call_args = {**injected_call, "type": "tool_call"} @@ -1040,6 +1055,7 @@ class ToolNode(RunnableCallable): tool=tool, state=tool_runtime.state, runtime=tool_runtime, + available_tools=list(self.tools_by_name.values()), ) config = tool_runtime.config @@ -1049,18 +1065,15 @@ class ToolNode(RunnableCallable): return self._execute_tool_sync(tool_request, input_type, config) # Define execute callable that can be called multiple times - original_name = call["name"] - def execute(req: ToolCallRequest) -> ToolMessage | Command: """Execute tool with given request. Can be called multiple times.""" - if req.tool_call["name"] != original_name: - # An interceptor (e.g., HITL) redirected the tool call - req = replace(req, tool=self.tools_by_name.get(req.tool_call["name"])) return self._execute_tool_sync(req, input_type, config) # Call wrapper with request and execute callable try: return self._wrap_tool_call(tool_request, execute) + except ToolCallRequestMismatchError: + raise except Exception as e: # Wrapper threw an exception if not self._handle_tool_errors: @@ -1104,6 +1117,13 @@ class ToolNode(RunnableCallable): msg = f"Tool {call['name']} is not registered with ToolNode" raise TypeError(msg) + if tool.name != call["name"]: + msg = ( + f"ToolCallRequest names tool {call['name']!r} but carries {tool.name!r}. " + f"An interceptor redirecting a call must set both `tool_call` and `tool`." + ) + raise ToolCallRequestMismatchError(msg) + # Inject state, store, and runtime right before invocation injected_call = self._inject_tool_args(call, request.runtime, tool) call_args = {**injected_call, "type": "tool_call"} @@ -1192,6 +1212,7 @@ class ToolNode(RunnableCallable): tool=tool, state=tool_runtime.state, runtime=tool_runtime, + available_tools=list(self.tools_by_name.values()), ) config = tool_runtime.config @@ -1201,13 +1222,8 @@ class ToolNode(RunnableCallable): return await self._execute_tool_async(tool_request, input_type, config) # Define async execute callable that can be called multiple times - original_name = call["name"] - async def execute(req: ToolCallRequest) -> ToolMessage | Command: """Execute tool with given request. Can be called multiple times.""" - if req.tool_call["name"] != original_name: - # An interceptor (e.g., HITL) redirected the tool call - req = replace(req, tool=self.tools_by_name.get(req.tool_call["name"])) return await self._execute_tool_async(req, input_type, config) def _sync_execute(req: ToolCallRequest) -> ToolMessage | Command: @@ -1221,6 +1237,8 @@ class ToolNode(RunnableCallable): # None check was performed above already self._wrap_tool_call = cast("ToolCallWrapper", self._wrap_tool_call) return self._wrap_tool_call(tool_request, _sync_execute) + except ToolCallRequestMismatchError: + raise except Exception as e: # Wrapper threw an exception if not self._handle_tool_errors: diff --git a/libs/prebuilt/tests/test_on_tool_call.py b/libs/prebuilt/tests/test_on_tool_call.py index 5cd482466..db7c719c5 100644 --- a/libs/prebuilt/tests/test_on_tool_call.py +++ b/libs/prebuilt/tests/test_on_tool_call.py @@ -13,6 +13,7 @@ from langgraph.types import Command from langgraph.prebuilt.tool_node import ( ToolCallRequest, + ToolCallRequestMismatchError, ToolNode, ) @@ -1474,7 +1475,7 @@ def test_tool_call_request_is_frozen() -> None: async def test_interceptor_can_redirect_to_another_tool() -> None: - """Overriding `tool_call["name"]` routes execution to the newly named tool.""" + """Redirecting requires setting both `tool_call` and `tool`; routing follows them.""" @tool def subtract(a: int, b: int) -> int: @@ -1485,8 +1486,11 @@ async def test_interceptor_can_redirect_to_another_tool() -> None: 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"}) + request.override( + tool_call={**request.tool_call, "name": "subtract"}, tool=target + ) ) node = ToolNode([add, subtract], awrap_tool_call=redirect) @@ -1504,3 +1508,33 @@ async def test_interceptor_can_redirect_to_another_tool() -> None: # `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(), + )