diff --git a/libs/prebuilt/langgraph/prebuilt/tool_node.py b/libs/prebuilt/langgraph/prebuilt/tool_node.py index d97dba2c3..cd0a5e01b 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"], @@ -1277,18 +1278,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), True)) 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"}