From 9fac9da5f7947565908e96109c8627555f68a99c Mon Sep 17 00:00:00 2001 From: Will Fu-Hinthorn Date: Tue, 21 Apr 2026 15:43:24 -0700 Subject: [PATCH] update --- libs/langgraph/langgraph/channels/delta.py | 27 ++++++-- libs/langgraph/tests/test_channels.py | 62 ++++++++++++++++++- libs/prebuilt/langgraph/prebuilt/tool_node.py | 26 +++++--- .../tests/test_injected_state_not_required.py | 43 ++++++++++++- 4 files changed, 141 insertions(+), 17 deletions(-) diff --git a/libs/langgraph/langgraph/channels/delta.py b/libs/langgraph/langgraph/channels/delta.py index 99aad844f..77e37dc00 100644 --- a/libs/langgraph/langgraph/channels/delta.py +++ b/libs/langgraph/langgraph/channels/delta.py @@ -2,6 +2,7 @@ from __future__ import annotations import collections.abc from collections.abc import Callable, Sequence +from copy import copy from typing import Any, Generic, Literal from langgraph.checkpoint.base import DeltaChainValue, DeltaValue @@ -15,6 +16,15 @@ from langgraph.errors import EmptyChannelError __all__ = ("DeltaChannel",) +def _copy_value(value: Any) -> Any: + if value is MISSING: + return value + try: + return value.copy() + except AttributeError: + return copy(value) + + class DeltaChannel(Generic[Value], BaseChannel[list[Value], Value, DeltaValue]): """A channel that stores only per-step write deltas in checkpoints. @@ -105,7 +115,7 @@ class DeltaChannel(Generic[Value], BaseChannel[list[Value], Value, DeltaValue]): def copy(self) -> Self: new = DeltaChannel(self.operator, self.typ, snapshot_every=self.snapshot_every) new.key = self.key - new.value = self.value if self.value is MISSING else self.value.copy() + new.value = _copy_value(self.value) new._pending = self._pending[:] new._base_version = self._base_version new._last_checkpoint_id = self._last_checkpoint_id @@ -138,8 +148,9 @@ class DeltaChannel(Generic[Value], BaseChannel[list[Value], Value, DeltaValue]): "have occurred before from_checkpoint was called." ) else: - # Backwards compat: plain list from old BinaryOperatorAggregate checkpoint. - new.value = list(checkpoint) + # Backwards compat: plain value from old BinaryOperatorAggregate checkpoint + # or a full snapshot emitted by DeltaChannel. + new.value = _copy_value(checkpoint) new._pending = [] new._base_version = None # set by the subsequent after_checkpoint() call new._overwritten = False @@ -165,9 +176,13 @@ class DeltaChannel(Generic[Value], BaseChannel[list[Value], Value, DeltaValue]): ) raise InvalidUpdateError(msg) self.value = ( - list(overwrite_value) if overwrite_value is not None else self.typ() + _copy_value(overwrite_value) + if overwrite_value is not None + else self.typ() + ) + self._pending = ( + [] if overwrite_value is None else [_copy_value(self.value)] ) - self._pending = list(self.value) self._overwritten = True seen_overwrite = True elif not seen_overwrite: @@ -192,7 +207,7 @@ class DeltaChannel(Generic[Value], BaseChannel[list[Value], Value, DeltaValue]): # Emit a full snapshot to cap chain depth at snapshot_every. # The saver stores this as a plain (non-diff) blob, so future # deltas will chain back to it and traversal depth resets to 1. - return list(self.value) + return _copy_value(self.value) return DeltaValue( delta=self._pending[:], prev_checkpoint_id=None if self._overwritten else self._last_checkpoint_id, diff --git a/libs/langgraph/tests/test_channels.py b/libs/langgraph/tests/test_channels.py index 4a8ad238f..57839d905 100644 --- a/libs/langgraph/tests/test_channels.py +++ b/libs/langgraph/tests/test_channels.py @@ -206,7 +206,7 @@ def test_delta_channel_from_checkpoint_backwards_compat() -> None: def test_delta_channel_overwrite_resets_chain() -> None: from langchain_core.messages import HumanMessage - from langgraph.checkpoint.base import DeltaValue + from langgraph.checkpoint.base import DeltaChainValue, DeltaValue from langgraph.channels.delta import DeltaChannel from langgraph.graph.message import add_messages @@ -223,7 +223,12 @@ def test_delta_channel_overwrite_resets_chain() -> None: assert isinstance(d, DeltaValue) assert d.prev_checkpoint_id is None # chain root assert len(d.delta) == 1 - assert d.delta[0].content == "new" + assert len(d.delta[0]) == 1 + assert d.delta[0][0].content == "new" + + spec = DeltaChannel(add_messages) + replayed = spec.from_checkpoint(DeltaChainValue(base=None, deltas=[d.delta])) + assert replayed.get()[0].content == "new" def test_delta_channel_assembly_fallback_via_get_tuple() -> None: @@ -449,6 +454,59 @@ def test_delta_channel_snapshot_every_end_to_end() -> None: assert len(msgs) == 10, f"expected 10 messages, got {len(msgs)}: {msgs}" +def test_delta_channel_dict_reducer_overwrite_preserves_mapping() -> None: + """Overwrite should preserve dict values instead of coercing them to keys.""" + from langgraph.checkpoint.base import DeltaChainValue, DeltaValue + + from langgraph.channels.delta import DeltaChannel + from langgraph.types import Overwrite + + def merge_dicts(left: dict, right: dict) -> dict: + return {**left, **right} + + ch = DeltaChannel(merge_dicts, dict).from_checkpoint(MISSING) + ch.after_checkpoint(None) + ch.update([{"a": 1}]) + ch.after_checkpoint("v1", checkpoint_id="cid1") + + ch.update([Overwrite({"b": 2})]) + d = ch.checkpoint() + assert isinstance(d, DeltaValue) + assert d.prev_checkpoint_id is None + assert d.delta == [{"b": 2}] + assert ch.get() == {"b": 2} + + spec = DeltaChannel(merge_dicts, dict) + replayed = spec.from_checkpoint(DeltaChainValue(base=None, deltas=[d.delta])) + assert replayed.get() == {"b": 2} + + +def test_delta_channel_dict_snapshot_every_round_trip() -> None: + """Full snapshots should preserve non-list reducers across reload.""" + from langgraph.checkpoint.base import DeltaValue + + from langgraph.channels.delta import DeltaChannel + + def merge_dicts(left: dict, right: dict) -> dict: + return {**left, **right} + + ch = DeltaChannel(merge_dicts, dict, snapshot_every=1).from_checkpoint(MISSING) + ch.after_checkpoint("v0", checkpoint_id="cid0") + + ch.update([{"a": 1}]) + first = ch.checkpoint() + assert isinstance(first, DeltaValue) + ch.after_checkpoint("v1", checkpoint_id="cid1") + + ch.update([{"b": 2}]) + snap = ch.checkpoint() + assert isinstance(snap, dict) + assert snap == {"a": 1, "b": 2} + + rehydrated = DeltaChannel(merge_dicts, dict, snapshot_every=1).from_checkpoint(snap) + assert rehydrated.get() == {"a": 1, "b": 2} + + def test_delta_channel_assembly_fast_path_returns_delta_value() -> None: """get_channel_blob returning a DeltaValue continues chain traversal (fast-path).""" from langgraph.checkpoint.base import ( diff --git a/libs/prebuilt/langgraph/prebuilt/tool_node.py b/libs/prebuilt/langgraph/prebuilt/tool_node.py index 32d293248..cd5ee2ae1 100644 --- a/libs/prebuilt/langgraph/prebuilt/tool_node.py +++ b/libs/prebuilt/langgraph/prebuilt/tool_node.py @@ -614,6 +614,7 @@ class _InjectedArgs: store: str | None runtime: str | None all_injected_keys: set[str] + _optional_state_args: set[str] class ToolNode(RunnableCallable): @@ -1333,7 +1334,7 @@ class ToolNode(RunnableCallable): return tool_call tool_call_copy: ToolCall = copy(tool_call) - injected_args = {} + injected_args: dict[str, Any] = {} # Inject state if injected.state: @@ -1361,14 +1362,20 @@ class ToolNode(RunnableCallable): # Extract state values if isinstance(state, dict): for tool_arg, state_field in injected.state.items(): - injected_args[tool_arg] = ( - state[state_field] if state_field else state - ) + if not state_field: + injected_args[tool_arg] = state + elif state_field in state: + injected_args[tool_arg] = state[state_field] + elif tool_arg not in injected._optional_state_args: + raise KeyError(state_field) else: for tool_arg, state_field in injected.state.items(): - injected_args[tool_arg] = ( - getattr(state, state_field) if state_field else state - ) + if not state_field: + injected_args[tool_arg] = state + elif hasattr(state, state_field): + injected_args[tool_arg] = getattr(state, state_field) + elif tool_arg not in injected._optional_state_args: + raise AttributeError(state_field) # Inject store if injected.store: @@ -1859,6 +1866,7 @@ def _get_all_injected_args(tool: BaseTool) -> _InjectedArgs: store_arg: str | None = None runtime_arg: str | None = None all_injected_keys: set[str] = set() + _optional_state_args: set[str] = set() for name, type_ in all_annotations.items(): # Track all InjectedToolArg-annotated params (including custom subclasses) @@ -1873,6 +1881,9 @@ def _get_all_injected_args(tool: BaseTool) -> _InjectedArgs: if state_inj := _get_injection_from_type(type_, InjectedState): if isinstance(state_inj, InjectedState) and state_inj.field: state_args[name] = state_inj.field + field_info = full_schema.model_fields.get(name) + if field_info and not field_info.is_required(): + _optional_state_args.add(name) else: state_args[name] = None @@ -1889,4 +1900,5 @@ def _get_all_injected_args(tool: BaseTool) -> _InjectedArgs: store=store_arg, runtime=runtime_arg, all_injected_keys=all_injected_keys, + _optional_state_args=_optional_state_args, ) diff --git a/libs/prebuilt/tests/test_injected_state_not_required.py b/libs/prebuilt/tests/test_injected_state_not_required.py index 6c3900416..fac7b624a 100644 --- a/libs/prebuilt/tests/test_injected_state_not_required.py +++ b/libs/prebuilt/tests/test_injected_state_not_required.py @@ -4,7 +4,7 @@ This tests the fix for https://github.com/langchain-ai/langchain/issues/35585 When using InjectedState() on a tool parameter, and the referenced field is declared as NotRequired in the custom state schema, the ToolNode should gracefully -handle missing fields by injecting None instead of raising KeyError. +handle missing fields without raising KeyError so the tool's default can apply. """ import sys @@ -45,6 +45,14 @@ def get_weather(city: Annotated[str | None, InjectedState("city")] = None) -> st return f"It's always sunny in {city}!" +@tool +def get_weather_with_default( + city: Annotated[str, InjectedState("city")] = "Boston", +) -> str: + """Get weather for a given city, defaulting when state omits the field.""" + return f"It's always sunny in {city}!" + + def _create_mock_runtime( state: dict | None = None, store=None, @@ -84,7 +92,7 @@ def _create_config_with_runtime(store=None, state=None): reason="InjectedState field extraction from Optional[Annotated[...]] not supported on Python <3.11", ) def test_injected_state_not_required_field_missing_injects_none(): - """Test that InjectedState with NotRequired field injects None when field is missing. + """Test that missing optional InjectedState leaves the tool default in place. This verifies the fix for https://github.com/langchain-ai/langchain/issues/35585 """ @@ -114,6 +122,37 @@ def test_injected_state_not_required_field_missing_injects_none(): assert "No city provided" in tool_msg.content +@pytest.mark.skipif( + sys.version_info < (3, 11), + reason="InjectedState field extraction from Optional[Annotated[...]] not supported on Python <3.11", +) +def test_injected_state_not_required_field_missing_preserves_tool_default(): + """Test that missing optional InjectedState preserves a non-None tool default.""" + tool_node = ToolNode([get_weather_with_default]) + + tool_call = { + "name": "get_weather_with_default", + "args": {}, + "id": "call_1", + "type": "tool_call", + } + ai_msg = AIMessage("Let me check the weather", tool_calls=[tool_call]) + + state_without_city: CustomAgentStateWithNotRequired = { + "messages": [HumanMessage("What's the weather?"), ai_msg], + } + + result = tool_node.invoke( + state_without_city, + config=_create_config_with_runtime(state=state_without_city), + ) + + assert len(result["messages"]) == 1 + tool_msg = result["messages"][0] + assert isinstance(tool_msg, ToolMessage) + assert "Boston" in tool_msg.content + + @pytest.mark.skipif( sys.version_info < (3, 11), reason="InjectedState field extraction from Optional[Annotated[...]] not supported on Python <3.11",