Files
langgraph/libs/langgraph/tests/test_channels.py
T
0a53c385b2 feat: public get_writes_history saver API + delta cadence rework (#7699)
## Summary

- Promotes the private K-channel batched ancestor-walk to a stable
public `get_delta_channel_history` / `aget_delta_channel_history` API on
`BaseCheckpointSaver` (returns `Mapping[str, DeltaChannelHistory]`, a
TypedDict with `writes` always present and `seed` `NotRequired`)
- Removes `DELTA_SENTINEL` / `_DeltaSentinel` entirely — the saver layer
is now delta-agnostic on both write and read paths
- Reworks `DeltaChannel` snapshot cadence from "every Nth superstep" to
"every N updates to this channel," persisted in
`CheckpointMetadata.delta_updates_since_snapshot`
- Adds Postgres optimizations: paged stage-1 with cursor (1024-row
pages) and per-channel UNION ALL stage-2 (no over-fetch when channels
have different chain depths)
- Default `snapshot_frequency` becomes a positive int (default `1000`);
the previous `None` opt-out is removed

## Public API

```python
class DeltaChannelHistory(TypedDict):
    writes: list[PendingWrite]            # always present, possibly empty
    seed: NotRequired[Any]                # absent if walk reached root

def get_delta_channel_history(
    self, *, config: RunnableConfig, channels: Sequence[str]
) -> Mapping[str, DeltaChannelHistory]: ...

async def aget_delta_channel_history(
    self, *, config: RunnableConfig, channels: Sequence[str]
) -> Mapping[str, DeltaChannelHistory]: ...
```

`config` and `channels` are keyword-only so later additions (e.g.
`page_size`) don't shift the positional API.

The TypedDict-with-`NotRequired[seed]` shape matches the existing
checkpoint-package convention (`CheckpointMetadata` is
`TypedDict(total=False)`) — absence-via-key-omission rather than
introducing a new sentinel. Pregel translates `"seed" not in hist` to
`MISSING` on its side at consume time.

The default impl walks `get_tuple` + `parent_config` correctly but is
slow on long chains; savers that care override (`InMemorySaver`,
`PostgresSaver`).

## Sentinel removal

`DELTA_SENTINEL` and `_DeltaSentinel` are deleted entirely. The saver
layer becomes delta-agnostic:

- `DeltaChannel.checkpoint()` returns `MISSING` for non-snapshot steps;
pregel's `create_checkpoint` skips MISSING so delta channels without a
snapshot simply don't appear in `channel_values`
- `InMemorySaver.put` and Postgres `put` no longer filter sentinels
(they have nothing to filter)
- `_needs_replay` becomes `stored is MISSING`
- `DeltaChannel.from_checkpoint` accepts: `MISSING` → empty,
`_DeltaSnapshot(value)` → snapshot value, plain value → pre-migration
legacy

## Snapshot cadence

`DeltaChannel.snapshot_frequency: int` (default `1000`, positive). The
previous `None` opt-out is gone.

```python
def should_snapshot(ch_name, ch):
    if force_delta_snapshot:                                  # durability="exit"
        return True
    return updates_since_snapshot.get(ch_name, 0) >= ch.snapshot_frequency
```

Per-channel update counters are persisted in
`CheckpointMetadata.delta_updates_since_snapshot` (`NotRequired`,
`total=False`). The counter is incremented by `_put_checkpoint` for any
delta channel in `updated_channels` and reset to `0` by
`create_checkpoint` for channels that fire a snapshot this step.
Version-format-independent — works for `int`, `float`, and `str`
versioning schemes alike.

## Postgres optimization

Two improvements internal to the override:

**Stage-1 paged with cursor** (`LIMIT 1024` internal const, `AND
checkpoint_id < ?` for subsequent pages). The previous unpaged form
scanned every checkpoint in `(thread_id, ns)` and was pathological at
high thread depths.

**Stage-2 per-channel UNION ALL**: one `WHERE channel='X' AND
checkpoint_id = ANY(chain_X)` branch per channel plus one seed-blob
branch per channel with a seed. The previous form filtered by `channel =
ANY(channels) AND checkpoint_id = ANY(union_chain_cids)`, over-fetching
writes when channels had different chain depths (`K ×
max(chain_lengths)` vs the correct `sum(chain_lengths)`).

Both improvements stay internal to `PostgresSaver`/`AsyncPostgresSaver`;
the public contract returns a single `Mapping`.

## Benchmarks

`libs/langgraph/tests/test_delta_channel_benchmark.py`. Run via `python
libs/langgraph/tests/test_delta_channel_benchmark.py`. Postgres against
local pg:5441.

Results below trimmed to the high-signal cells. Sub-millisecond /
sub-100-turn rows omitted as warmup-bound; freq=1 omitted (chain depth =
1, nothing to optimize); peak read-time memory and Postgres storage are
flat between branches and omitted. Deep-thread reads and the
cadence-rework storage win are the load-bearing numbers.

### Postgres reads, 500 turns

| Scenario | main | branch | Δ |
|---|---:|---:|---:|
| Single-channel deep read | 17.7 ms | **6.1 ms** | **-66%** |
| Single-channel, 1000 turns | 35.0 ms | **14.3 ms** | **-59%** |
| K=3 channels, freq=50 uniform | 70.5 ms | **41.4 ms** | **-41%** |
| K=8 channels, freq=50 uniform | 214.2 ms | **139.4 ms** | **-35%** |
| K=8 channels, mixed freq (25/50/100/.../1000) | 295.6 ms | **214.4
ms** | **-27%** |

K-channel batching + paged stage-1 + per-channel UNION ALL stage-2 doing
exactly what they should at depth.

### InMemory reads, 500 turns

| Scenario | main | branch | Δ |
|---|---:|---:|---:|
| Single-channel deep read | 7.9 ms | **3.8 ms** | **-52%** |
| Single-channel, 1000 turns | 15.6 ms | **7.2 ms** | **-54%** |
| K=8 channels, freq=50 uniform | 112.3 ms | 94.6 ms | -16% |
| K=8 channels, mixed freq | 184.9 ms | **134.5 ms** | **-27%** |

### InMemory storage, 500 turns (cadence-rework win)

| Scenario | main | branch | Δ |
|---|---:|---:|---:|
| K=3, freq=50 uniform | 8.7 MB | **3.3 MB** | **-62%** |
| K=3 mixed freq | 3.8 MB | **1.3 MB** | **-66%** |
| K=8, freq=50 uniform | 23.1 MB | **8.7 MB** | **-62%** |
| K=8 mixed freq | 11.5 MB | **4.2 MB** | **-64%** |

Snapshot frequency now counts **channel updates** instead of
**supersteps**. On graphs where supersteps outpace per-channel updates
(e.g., input/end steps that don't write to channels), branch stores ~3×
fewer snapshot blobs.

### Tradeoff worth flagging

InMemory K=3 with mixed frequencies (50/200/1000) at 500 turns: **+64%
read latency** (46.6 → 76.5 ms). The mixed scenario has a channel with
`freq=1000` that goes the entire 500-turn run with no snapshot. On main,
the old superstep-counted cadence happened to fire at step=500 anyway.
New cadence gives users explicit control over walk depth via
`snapshot_frequency`. The K=8 mixed case still wins overall (-27%); this
regression is specific to the K=3 mixed shape.

Default `snapshot_frequency=1000` is the upper bound on walk depth —
it's a tunable knob.

## Tests

- New sqlite smoke test (`test_get_delta_channel_history.py`) exercises
the inherited default `BaseCheckpointSaver` impl via `SqliteSaver` /
`AsyncSqliteSaver` end-to-end with a real `DeltaChannel`-backed graph.
Sqlite uses the default unchanged — this validates the default path
actually works on a real second saver, not just on the optimized
override.
- Module-level `pytest.importorskip("langgraph.channels.delta")` guards
the test for sqlite's standalone CI environment (matches the postgres
pattern).

## Test plan

- [x] `libs/checkpoint`: 150 passed, 16 skipped
- [x] `libs/langgraph` (channels + delta migration): 41/41 (post-merge)
- [x] `libs/langgraph` (full pregel suite): 1784 passing — 6 "failures"
verified via `env -i` clean shell are local LangSmith env vars + `git
describe revision_id` polluting LangChain metadata fixtures; CI is
unaffected
- [x] `libs/checkpoint-postgres`: 40/40 saver tests + 3/3 delta channel
reconstruction tests against local Postgres
- [x] `libs/checkpoint-sqlite`: 105/105 (incl. retry-passed flake
`test_ttl_refresh`, unrelated to this PR)
- [x] Lint clean across all four libs (`ruff format`, `ruff check`,
`mypy`)
- [x] Branch-vs-main benchmarks — see results above

---------

Co-authored-by: Quanzheng Long <long@langchain.dev>
Co-authored-by: Cursor <cursoragent@cursor.com>
2026-05-04 15:18:43 -04:00

673 lines
23 KiB
Python

import operator
from collections.abc import Sequence
from typing import Annotated
import pytest
from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import NotRequired, TypedDict
from langgraph._internal._typing import MISSING
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.last_value import LastValue
from langgraph.channels.topic import Topic
from langgraph.channels.untracked_value import UntrackedValue
from langgraph.errors import EmptyChannelError, InvalidUpdateError
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
from langgraph.graph.state import _get_channel
from langgraph.types import Overwrite
pytestmark = pytest.mark.anyio
# ---------------------------------------------------------------------------
# Core channel primitives
# ---------------------------------------------------------------------------
def test_last_value() -> None:
channel = LastValue(int).from_checkpoint(MISSING)
assert channel.ValueType is int
assert channel.UpdateType is int
with pytest.raises(EmptyChannelError):
channel.get()
with pytest.raises(InvalidUpdateError):
channel.update([5, 6])
channel.update([3])
assert channel.get() == 3
channel.update([4])
assert channel.get() == 4
checkpoint = channel.checkpoint()
channel = LastValue(int).from_checkpoint(checkpoint)
assert channel.get() == 4
def test_topic() -> None:
channel = Topic(str).from_checkpoint(MISSING)
assert channel.ValueType == Sequence[str]
assert channel.UpdateType == str | list[str]
assert channel.update(["a", "b"])
assert channel.get() == ["a", "b"]
assert channel.update([["c", "d"], "d"])
assert channel.get() == ["c", "d", "d"]
assert channel.update([])
with pytest.raises(EmptyChannelError):
channel.get()
assert not channel.update([]), "channel already empty"
assert channel.update(["e"])
assert channel.get() == ["e"]
checkpoint = channel.checkpoint()
channel = Topic(str).from_checkpoint(checkpoint)
assert channel.get() == ["e"]
channel_copy = Topic(str).from_checkpoint(checkpoint)
channel_copy.update(["f"])
assert channel_copy.get() == ["f"]
assert channel.get() == ["e"]
def test_topic_accumulate() -> None:
channel = Topic(str, accumulate=True).from_checkpoint(MISSING)
assert channel.ValueType == Sequence[str]
assert channel.UpdateType == str | list[str]
assert channel.update(["a", "b"])
assert channel.get() == ["a", "b"]
assert channel.update(["b", ["c", "d"], "d"])
assert channel.get() == ["a", "b", "b", "c", "d", "d"]
assert not channel.update([])
assert channel.get() == ["a", "b", "b", "c", "d", "d"]
checkpoint = channel.checkpoint()
channel = Topic(str, accumulate=True).from_checkpoint(checkpoint)
assert channel.get() == ["a", "b", "b", "c", "d", "d"]
assert channel.update(["e"])
assert channel.get() == ["a", "b", "b", "c", "d", "d", "e"]
def test_binop() -> None:
channel = BinaryOperatorAggregate(int, operator.add).from_checkpoint(MISSING)
assert channel.ValueType is int
assert channel.UpdateType is int
assert channel.get() == 0
channel.update([1, 2, 3])
assert channel.get() == 6
channel.update([4])
assert channel.get() == 10
checkpoint = channel.checkpoint()
channel = BinaryOperatorAggregate(int, operator.add).from_checkpoint(checkpoint)
assert channel.get() == 10
def test_untracked_value() -> None:
channel = UntrackedValue(dict).from_checkpoint(MISSING)
assert channel.ValueType is dict
assert channel.UpdateType is dict
with pytest.raises(EmptyChannelError):
channel.get()
test_data = {"session": "test", "temp": "dir"}
channel.update([test_data])
assert channel.get() == test_data
new_data = {"session": "updated", "temp": "newdir"}
channel.update([new_data])
assert channel.get() == new_data
checkpoint = channel.checkpoint()
assert checkpoint is MISSING
new_channel = UntrackedValue(dict).from_checkpoint(checkpoint)
with pytest.raises(EmptyChannelError):
new_channel.get()
# ---------------------------------------------------------------------------
# DeltaChannel — message reducer
# ---------------------------------------------------------------------------
def test_delta_channel_basic_two_steps() -> None:
ch = DeltaChannel(_messages_delta_reducer, list).from_checkpoint(MISSING)
ch.update([HumanMessage(content="hi", id="h1")])
d1 = ch.checkpoint()
assert d1 is MISSING
ch.update([AIMessage(content="hello", id="a1")])
d2 = ch.checkpoint()
assert d2 is MISSING
assert len(ch.get()) == 2
assert ch.get()[0].content == "hi"
assert ch.get()[1].content == "hello"
def test_delta_channel_from_checkpoint_writes_list() -> None:
"""replay_writes on a fresh channel replays through the operator."""
spec = DeltaChannel(_messages_delta_reducer, list)
ch = spec.from_checkpoint(MISSING)
ch.replay_writes(
[
("t0", "messages", HumanMessage(content="hi", id="h1")),
("t1", "messages", AIMessage(content="hello", id="a1")),
("t2", "messages", HumanMessage(content="bye", id="h2")),
]
)
msgs = ch.get()
assert len(msgs) == 3
assert msgs[0].content == "hi"
assert msgs[1].content == "hello"
assert msgs[2].content == "bye"
def test_delta_channel_from_checkpoint_backwards_compat() -> None:
spec = DeltaChannel(_messages_delta_reducer, list)
old_value = [HumanMessage(content="old", id="h1")]
ch = spec.from_checkpoint(old_value)
assert ch.get() == old_value
def test_delta_channel_overwrite() -> None:
ch = DeltaChannel(_messages_delta_reducer, list).from_checkpoint(MISSING)
ch.update([HumanMessage(content="old", id="h1")])
ch.update([Overwrite([HumanMessage(content="new", id="h2")])])
d = ch.checkpoint()
assert d is MISSING
assert len(ch.get()) == 1
assert ch.get()[0].content == "new"
def test_delta_channel_remove_message_and_replay() -> None:
"""RemoveMessage must round-trip correctly when writes are replayed."""
spec = DeltaChannel(_messages_delta_reducer, list)
ch = spec.from_checkpoint(MISSING)
ch.update([HumanMessage(content="hi", id="h1")])
ch.update([AIMessage(content="hello", id="a1")])
assert ch.get() == [
HumanMessage(content="hi", id="h1"),
AIMessage(content="hello", id="a1"),
]
ch.update([RemoveMessage(id="a1")])
assert ch.get() == [HumanMessage(content="hi", id="h1")]
ch2 = spec.from_checkpoint(MISSING)
ch2.replay_writes(
[
("t0", "messages", HumanMessage(content="hi", id="h1")),
("t1", "messages", AIMessage(content="hello", id="a1")),
("t2", "messages", RemoveMessage(id="a1")),
]
)
assert ch2.get() == [HumanMessage(content="hi", id="h1")]
def test_delta_channel_update_by_id_and_replay() -> None:
"""Updating a message by ID must round-trip correctly through writes replay."""
spec = DeltaChannel(_messages_delta_reducer, list)
ch = spec.from_checkpoint(MISSING)
ch.update([HumanMessage(content="original", id="h1")])
ch.update([HumanMessage(content="updated", id="h1")])
assert ch.get() == [HumanMessage(content="updated", id="h1")]
ch2 = spec.from_checkpoint(MISSING)
ch2.replay_writes(
[
("t0", "messages", HumanMessage(content="original", id="h1")),
("t1", "messages", HumanMessage(content="updated", id="h1")),
]
)
assert len(ch2.get()) == 1
assert ch2.get()[0].content == "updated"
def test_delta_channel_dict_coercion() -> None:
"""_messages_delta_reducer coerces dict writes to BaseMessage objects.
HTTP-driven input always arrives as JSON dicts. The reducer must coerce
them (same contract as add_messages) so graphs work without a separate
coercion step.
"""
ch = DeltaChannel(_messages_delta_reducer, list).from_checkpoint(MISSING)
# dict input — simulates what arrives from the HTTP API
ch.update([{"role": "human", "content": "hello", "id": "h1"}])
assert len(ch.get()) == 1
assert isinstance(ch.get()[0], HumanMessage)
assert ch.get()[0].content == "hello"
assert ch.get()[0].id == "h1"
# update by ID via dict
ch.update([{"role": "ai", "content": "world", "id": "h1"}])
assert len(ch.get()) == 1
assert ch.get()[0].content == "world"
# remove via RemoveMessage instance (same contract as add_messages)
ch.update([RemoveMessage(id="h1")])
assert ch.get() == []
def test_messages_delta_reducer_coerces_state() -> None:
"""State (left side) is coerced when raw — supports raw initial input
and deserialized blobs. The steady-state path (state already typed)
short-circuits and skips coercion.
"""
state = [{"role": "human", "content": "hello", "id": "h1"}]
writes = [[{"role": "ai", "content": "world", "id": "h1"}]]
result = _messages_delta_reducer(state, writes) # type: ignore[arg-type]
assert len(result) == 1
assert isinstance(result[0], AIMessage)
assert result[0].content == "world"
assert result[0].id == "h1"
def test_messages_delta_reducer_tuple_write_is_one_message() -> None:
"""A top-level tuple write is one message-like, not a sequence to flatten.
`("user", "hi")` is a valid `MessageLikeRepresentation`; flattening it
would produce two HumanMessages ("user", "hi") instead of one.
"""
result = _messages_delta_reducer([], [("user", "hi")]) # type: ignore[arg-type]
assert len(result) == 1
assert isinstance(result[0], HumanMessage)
assert result[0].content == "hi"
def test_delta_channel_checkpoint_returns_missing() -> None:
"""checkpoint() always returns MISSING regardless of state.
Pregel writes `_DeltaSnapshot(ch.get())` directly into `channel_values`
on snapshot steps; the channel itself never participates in snapshot
serialization, so its `checkpoint()` is always the absence sentinel.
"""
ch = DeltaChannel(_messages_delta_reducer, list).from_checkpoint(MISSING)
assert ch.checkpoint() is MISSING
ch.update([HumanMessage(content="hi", id="h1")])
assert ch.checkpoint() is MISSING
# ---------------------------------------------------------------------------
# DeltaChannel — snapshot frequency
# ---------------------------------------------------------------------------
def test_delta_channel_snapshot_version_based() -> None:
"""Snapshots fire when a channel accumulates `snapshot_frequency` updates.
Under the version-delta cadence, every time the channel's
`current_version - last_snapshot_version >= snapshot_frequency` a
`_DeltaSnapshot` blob is written. Bounds the ancestor walk to at most
`snapshot_frequency` steps on any read for that channel.
"""
class State(TypedDict):
messages: Annotated[
list, DeltaChannel(_messages_delta_reducer, snapshot_frequency=5)
]
other: str
def node_a(state: State) -> dict:
i = len(state["messages"]) // 2
return {"messages": [AIMessage(content=f"a{i}", id=f"a{i}")]}
def node_b(state: State) -> dict:
return {"other": "y"}
g = StateGraph(State)
g.add_node("a", node_a)
g.add_node("b", node_b)
g.add_edge(START, "a")
g.add_edge("a", "b")
saver = InMemorySaver()
graph = g.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "t1"}}
for i in range(6):
graph.invoke(
{"messages": [HumanMessage(content=f"h{i}", id=f"h{i}")], "other": ""},
config,
)
msg_blob_values = [
saver.serde.loads_typed((type_tag, blob))
for k, (type_tag, blob) in saver.blobs.items()
if k[2] == "messages" and type_tag == "msgpack" and blob
]
snapshots = [v for v in msg_blob_values if isinstance(v, _DeltaSnapshot)]
assert snapshots, "expected at least one _DeltaSnapshot blob for messages"
state = graph.get_state(config)
assert len(state.values["messages"]) == 12 # 6 human + 6 AI
# TODO(delta-channel-cadence): the previous "snapshot fires even when channel
# was not written" test asserted eager step-based snapshotting; under the new
# version-delta cadence (`should_snapshot` triggers on per-channel update
# count, not superstep count), no snapshot fires for an unwritten channel.
# Replace with a test that exercises the version-delta trigger plus the
# durability="exit" force-snapshot branch — see
# `docs/superpowers/specs/2026-05-04-delta-channel-batched-reads-design.md`
# section "Snapshot cadence".
# ---------------------------------------------------------------------------
# DeltaChannel — end-to-end (InMemorySaver)
# ---------------------------------------------------------------------------
def test_delta_channel_inmemory_saver_assembles_writes() -> None:
"""InMemorySaver assembles writes from checkpoint_writes inside get_tuple."""
class State(TypedDict):
messages: Annotated[list, DeltaChannel(_messages_delta_reducer, list)]
n = {"v": 0}
def respond(state: State) -> dict:
n["v"] += 1
return {"messages": [AIMessage(content=f"ok{n['v']}", id=f"ai{n['v']}")]}
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
saver = InMemorySaver()
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "t1"}}
graph.invoke({"messages": [HumanMessage(content="hi", id="h1")]}, config)
graph.invoke({"messages": [HumanMessage(content="bye", id="h2")]}, config)
saved = saver.get_tuple(config)
assert saved is not None
assert "messages" not in saved.checkpoint["channel_values"]
state = graph.get_state(config)
assert len(state.values["messages"]) == 4 # 2 human + 2 AI
# ---------------------------------------------------------------------------
# DeltaChannel — dict reducer
# ---------------------------------------------------------------------------
def _delta_channel_with_type(op, typ):
"""Build a DeltaChannel with an explicit type via the Annotated injection path."""
return _get_channel("_test", Annotated[typ, DeltaChannel(op)])
def test_delta_channel_dict_reducer_fresh_channel() -> None:
"""DeltaChannel with a dict reducer starts as empty dict on MISSING checkpoint."""
def merge_dicts(state: dict, writes: list) -> dict:
result = dict(state)
for w in writes:
result.update(w)
return result
ch = _delta_channel_with_type(merge_dicts, dict).from_checkpoint(MISSING)
assert ch.is_available()
assert ch.get() == {}
def test_delta_channel_dict_reducer_basic_updates() -> None:
"""DeltaChannel with a dict reducer accumulates key/value pairs across steps."""
def merge_dicts(state: dict, writes: list) -> dict:
result = dict(state)
for w in writes:
result.update(w)
return result
ch = _delta_channel_with_type(merge_dicts, dict).from_checkpoint(MISSING)
ch.update([{"a": 1}])
d1 = ch.checkpoint()
assert d1 is MISSING
ch.update([{"b": 2}])
d2 = ch.checkpoint()
assert d2 is MISSING
assert ch.get() == {"a": 1, "b": 2}
def test_delta_channel_dict_reducer_writes_reconstruction() -> None:
"""replay_writes on a fresh channel replays through a dict merge reducer."""
def merge_dicts(state: dict, writes: list) -> dict:
result = dict(state)
for w in writes:
result.update(w)
return result
spec = _delta_channel_with_type(merge_dicts, dict)
ch = spec.from_checkpoint(MISSING)
ch.replay_writes(
[
("t0", "files", {"a": 1}),
("t1", "files", {"b": 2}),
("t2", "files", {"c": 3}),
]
)
assert ch.get() == {"a": 1, "b": 2, "c": 3}
def test_delta_channel_dict_reducer_with_deletions() -> None:
"""Dict reducer that treats None values as deletions works end-to-end."""
def merge_files(state: dict, writes: list) -> dict:
result = dict(state)
for w in writes:
for k, v in w.items():
if v is None:
result.pop(k, None)
else:
result[k] = v
return result
ch = _delta_channel_with_type(merge_files, dict).from_checkpoint(MISSING)
ch.update([{"file1.py": "content1", "file2.py": "content2"}])
ch.update([{"file1.py": None, "file3.py": "content3"}])
assert ch.get() == {"file2.py": "content2", "file3.py": "content3"}
spec = _delta_channel_with_type(merge_files, dict)
ch2 = spec.from_checkpoint(MISSING)
ch2.replay_writes(
[
("t0", "files", {"file1.py": "content1", "file2.py": "content2"}),
("t1", "files", {"file1.py": None, "file3.py": "content3"}),
]
)
assert ch2.get() == {"file2.py": "content2", "file3.py": "content3"}
def test_delta_channel_dict_reducer_overwrite_in_update() -> None:
"""Overwrite(dict) in update() must preserve dict shape, not coerce to list."""
def merge_dicts(state: dict, writes: list) -> dict:
result = dict(state)
for w in writes:
result.update(w)
return result
ch = _delta_channel_with_type(merge_dicts, dict).from_checkpoint(MISSING)
ch.update([{"a": 1}])
ch.update([Overwrite({"b": 2, "c": 3})])
assert ch.get() == {"b": 2, "c": 3}
def test_delta_channel_dict_reducer_overwrite_in_writes_replay() -> None:
"""Overwrite(dict) embedded in replayed writes must reconstruct as dict."""
def merge_dicts(state: dict, writes: list) -> dict:
result = dict(state)
for w in writes:
result.update(w)
return result
spec = _delta_channel_with_type(merge_dicts, dict)
ch = spec.from_checkpoint(MISSING)
ch.replay_writes(
[
("t0", "files", {"a": 1}),
("t1", "files", Overwrite({"x": 10, "y": 20})),
("t2", "files", {"z": 30}),
]
)
assert ch.get() == {"x": 10, "y": 20, "z": 30}
def test_delta_channel_dict_reducer_with_notrequired_annotation() -> None:
"""DeltaChannel infers dict type through `Annotated[NotRequired[dict[...]], ch]`."""
def merge_dicts(state: dict, writes: list) -> dict:
result = dict(state)
for w in writes:
result.update(w)
return result
annotation = Annotated[NotRequired[dict[str, int]], DeltaChannel(merge_dicts)]
ch = _get_channel("files", annotation).from_checkpoint(MISSING)
assert ch.get() == {}
ch.update([{"a": 1}])
ch.update([{"b": 2}])
assert ch.get() == {"a": 1, "b": 2}
def test_delta_channel_dict_reducer_end_to_end_filesystem() -> None:
"""End-to-end: graph with dict-reducer (filesystem-style) channel wrapped in DeltaChannel."""
def merge_files(state: dict, writes: list) -> dict:
result = dict(state)
for w in writes:
for k, v in w.items():
if v is None:
result.pop(k, None)
else:
result[k] = v
return result
class State(TypedDict):
files: Annotated[dict[str, str], DeltaChannel(merge_files)]
turn = {"v": 0}
def write_file(state: State) -> dict:
turn["v"] += 1
n = turn["v"]
return {"files": {f"/doc_{n}.txt": f"content for turn {n}"}}
builder = StateGraph(State)
builder.add_node("write_file", write_file)
builder.add_edge(START, "write_file")
saver = InMemorySaver()
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "fs"}}
for _ in range(3):
graph.invoke({"files": {}}, config)
saved = saver.get_tuple(config)
assert saved is not None
assert "files" not in saved.checkpoint["channel_values"]
state = graph.get_state(config)
assert state.values["files"] == {
"/doc_1.txt": "content for turn 1",
"/doc_2.txt": "content for turn 2",
"/doc_3.txt": "content for turn 3",
}
def delete_file(state: State) -> dict:
return {"files": {"/doc_1.txt": None}}
builder2 = StateGraph(State)
builder2.add_node("write_file", write_file)
builder2.add_node("delete_file", delete_file)
builder2.add_edge(START, "write_file")
builder2.add_edge("write_file", "delete_file")
turn["v"] = 0
saver2 = InMemorySaver()
graph2 = builder2.compile(checkpointer=saver2)
config2 = {"configurable": {"thread_id": "fs2"}}
graph2.invoke({"files": {}}, config2)
state2 = graph2.get_state(config2)
assert state2.values["files"] == {}
def test_delta_channel_dict_reducer_backwards_compat() -> None:
"""A pre-DeltaChannel dict checkpoint must load as a dict, not be listified."""
def merge_dicts(state: dict, writes: list) -> dict:
result = dict(state)
for w in writes:
result.update(w)
return result
spec = _delta_channel_with_type(merge_dicts, dict)
old_value = {"a": 1, "b": 2}
ch = spec.from_checkpoint(old_value)
assert ch.get() == {"a": 1, "b": 2}
# ---------------------------------------------------------------------------
# DeltaChannel — seed / pre-delta migration
# ---------------------------------------------------------------------------
def test_delta_channel_from_checkpoint_honors_seed() -> None:
"""A non-sentinel value to from_checkpoint is used as the pre-delta seed.
Guards the pre-delta migration path: when the saver's ancestor walk hits
a pre-DeltaChannel blob it passes it as `seed` so replay reconstructs
the post-migration state correctly rather than replaying from empty.
"""
spec = DeltaChannel(_messages_delta_reducer, list)
seed = [HumanMessage(content="pre-delta", id="p1")]
ch = spec.from_checkpoint(seed)
ch.replay_writes(
[
("t0", "messages", AIMessage(content="delta-1", id="d1")),
("t1", "messages", HumanMessage(content="delta-2", id="d2")),
]
)
msgs = ch.get()
assert [m.content for m in msgs] == ["pre-delta", "delta-1", "delta-2"]
def test_delta_channel_from_checkpoint_seed_without_writes() -> None:
"""Reconstruction at a pre-delta ancestor with no newer deltas returns
just the seed — the saver's terminator fired immediately."""
spec = DeltaChannel(_messages_delta_reducer, list)
seed = [HumanMessage(content="only-snap", id="s1")]
ch = spec.from_checkpoint(seed)
ch.replay_writes([])
assert ch.get() == seed
def test_delta_channel_from_checkpoint_seed_none_is_distinct_from_sentinel() -> None:
"""`seed=None` must start replay from None, not from an empty channel.
The `MISSING` absence sentinel means 'no seed'; passing `None`
explicitly should feed None to the reducer as the left operand.
"""
def replace(state, writes):
return writes[-1] if writes else state
spec = DeltaChannel(replace, list)
ch = spec.from_checkpoint(None)
ch.replay_writes([("t0", "x", "after")])
assert ch.get() == "after"