mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 12:35:08 +02:00
update
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user