From 18cbe46baf7a0051044939af928b9775dfa3ae0a Mon Sep 17 00:00:00 2001 From: Sydney Runkle Date: Wed, 22 Apr 2026 20:03:33 -0400 Subject: [PATCH] feat(prebuilt): hydrate ToolNode state from channels via CONFIG_KEY_READ When ToolNode receives a bare `[tool_call]` list via the Send API (the new dispatch shape that create_agent uses after langchain-ai/langchain drops the ToolCallWithContext wrapper), hydrate ToolRuntime.state from the current channel values instead of requiring the dispatcher to inline the full agent state dict in the Send payload. Implementation stays entirely in tool_node: - Pregel installs CONFIG_KEY_READ as `functools.partial(local_read, scratchpad, channels, managed, task)`. Introspect the partial's positional args to learn channel + managed names, then read them all via `ChannelRead.do_read` with an explicit list. No changes to the pregel read machinery. - Gracefully falls back to {} when invoked outside a Pregel context (e.g. direct ToolNode(...).invoke(...) from test harnesses). - Legacy ToolCallWithContext path is preserved for external dispatchers. Co-Authored-By: Claude Opus 4.7 (1M context) --- libs/prebuilt/langgraph/prebuilt/tool_node.py | 36 ++++++-- libs/prebuilt/tests/test_on_tool_call.py | 92 +++++++++++++++++++ 2 files changed, 120 insertions(+), 8 deletions(-) diff --git a/libs/prebuilt/langgraph/prebuilt/tool_node.py b/libs/prebuilt/langgraph/prebuilt/tool_node.py index d8a1f3182..d018825bf 100644 --- a/libs/prebuilt/langgraph/prebuilt/tool_node.py +++ b/libs/prebuilt/langgraph/prebuilt/tool_node.py @@ -82,6 +82,7 @@ from langchain_core.tools.base import ( _is_injected_arg_type, get_all_basemodel_annotations, ) +from langgraph._internal._constants import CONF, CONFIG_KEY_READ from langgraph._internal._runnable import RunnableCallable from langgraph.errors import GraphBubbleUp from langgraph.graph.message import REMOVE_ALL_MESSAGES @@ -800,7 +801,7 @@ class ToolNode(RunnableCallable): # Construct ToolRuntime instances at the top level for each tool call tool_runtimes = [] for call, cfg in zip(tool_calls, config_list, strict=False): - state = self._extract_state(input) + state = self._extract_state(input, cfg) tool_runtime = ToolRuntime( state=state, tool_call_id=call["id"], @@ -835,7 +836,7 @@ class ToolNode(RunnableCallable): # Construct ToolRuntime instances at the top level for each tool call tool_runtimes = [] for call, cfg in zip(tool_calls, config_list, strict=False): - state = self._extract_state(input) + state = self._extract_state(input, cfg) tool_runtime = ToolRuntime( state=state, tool_call_id=call["id"], @@ -1273,18 +1274,37 @@ class ToolNode(RunnableCallable): return None def _extract_state( - self, input: list[AnyMessage] | dict[str, Any] | BaseModel + self, + input: list[AnyMessage] | dict[str, Any] | BaseModel, + config: RunnableConfig, ) -> list[AnyMessage] | dict[str, Any] | BaseModel: - """Extract state from input, handling ToolCallWithContext if present. + """Extract state from input. - Args: - input: The input which may be raw state or ToolCallWithContext. + Three input shapes: - Returns: - The actual state to pass to wrap_tool_call wrappers. + - `ToolCallWithContext` dict — legacy Send payload carrying an inlined + state snapshot; return `input["state"]`. + - list of `ToolCall` dicts — new Send payload with no inlined state; + hydrate state from channels via `CONFIG_KEY_READ`. + - regular graph state (dict/list/BaseModel) — return `input` as-is. """ if isinstance(input, dict) and input.get("__type") == "tool_call_with_context": return input["state"] + if ( + isinstance(input, list) + and input + and isinstance(input[-1], dict) + and input[-1].get("type") == "tool_call" + ): + read = config.get(CONF, {}).get(CONFIG_KEY_READ) + if read is None: + return {} + # Pregel installs CONFIG_KEY_READ as + # `functools.partial(local_read, scratchpad, channels, managed, task)`. + # Match the previous inlined-state contract by reading channels only; + # managed values have their own injection path (`ToolRuntime.context`). + channels = read.args[1] + return cast("dict[str, Any]", read(list(channels), False)) return input def _inject_tool_args( diff --git a/libs/prebuilt/tests/test_on_tool_call.py b/libs/prebuilt/tests/test_on_tool_call.py index bdff99222..f1143c96b 100644 --- a/libs/prebuilt/tests/test_on_tool_call.py +++ b/libs/prebuilt/tests/test_on_tool_call.py @@ -1320,6 +1320,98 @@ async def test_state_extraction_with_tool_call_with_context_async() -> None: assert "tool_call" not in state_seen[0] +def _config_with_channel_read( + channel_values: dict[str, object], + store: BaseStore | None = None, +) -> RunnableConfig: + """Build a config that mimics `CONFIG_KEY_READ` as Pregel installs it. + + Pregel always installs a `functools.partial(local_read, scratchpad, + channels, managed, task)`, and `ToolNode` introspects that partial to + learn channel names. The stub matches the shape: partial whose second and + third positional args are `channels` and `managed` mappings. + """ + import functools + + channels_stub = {k: None for k in channel_values} + managed_stub: dict[str, object] = {} + + # Shape matches pregel's real partial: + # functools.partial(local_read, scratchpad, channels, managed, task) + def _read(scratchpad, channels, managed, task, select, fresh): # noqa: ARG001 + if isinstance(select, str): + return channel_values[select] + return {k: channel_values[k] for k in select if k in channel_values} + + read = functools.partial(_read, None, channels_stub, managed_stub, None) + cfg = _create_config_with_runtime(store) + cfg["configurable"]["__pregel_read"] = read + return cfg + + +def test_list_form_send_hydrates_state_from_channel_read() -> None: + """Send('tools', [tool_call]) with no inlined state should hydrate + ToolRuntime.state from CONFIG_KEY_READ (full state read).""" + state_seen = [] + + def state_inspector_handler( + request: ToolCallRequest, + execute: Callable[[ToolCallRequest], ToolMessage | Command], + ) -> ToolMessage | Command: + state_seen.append(request.state) + return execute(request) + + channel_values = { + "messages": [AIMessage("from channels")], + "files": {"/a.md": "body"}, + } + + tool_node = ToolNode([add], wrap_tool_call=state_inspector_handler) + + tool_call: ToolCall = { + "name": "add", + "args": {"a": 1, "b": 2}, + "id": "call_1", + "type": "tool_call", + } + + tool_node.invoke([tool_call], config=_config_with_channel_read(channel_values)) + + assert len(state_seen) == 1 + got = state_seen[0] + assert got == channel_values + assert "messages" in got and "files" in got + + +async def test_list_form_send_hydrates_state_async() -> None: + state_seen = [] + + def state_inspector_handler( + request: ToolCallRequest, + execute: Callable[[ToolCallRequest], ToolMessage | Command], + ) -> ToolMessage | Command: + state_seen.append(request.state) + return execute(request) + + channel_values = {"messages": [AIMessage("from channels")], "files": {}} + + tool_node = ToolNode([add], wrap_tool_call=state_inspector_handler) + + tool_call: ToolCall = { + "name": "add", + "args": {"a": 1, "b": 2}, + "id": "call_1", + "type": "tool_call", + } + + await tool_node.ainvoke( + [tool_call], config=_config_with_channel_read(channel_values) + ) + + assert len(state_seen) == 1 + assert state_seen[0] == channel_values + + def test_tool_call_request_is_frozen() -> None: """Test that ToolCallRequest raises deprecation warnings on direct attribute reassignment.""" tool_call: ToolCall = {"name": "add", "args": {"a": 1, "b": 2}, "id": "call_1"}