This commit is contained in:
Will Fu-Hinthorn
2026-04-21 15:43:24 -07:00
parent def55a5ac5
commit 9fac9da5f7
4 changed files with 141 additions and 17 deletions
+21 -6
View File
@@ -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,
+60 -2
View File
@@ -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 (
+19 -7
View File
@@ -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,
)
@@ -4,7 +4,7 @@ This tests the fix for https://github.com/langchain-ai/langchain/issues/35585
When using InjectedState(<field>) 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",