chore(langgraph): re-implement exit mode for delta channel (#7730)

## Summary

Replaces `durability="exit"`'s blanket force-snapshot of every
`DeltaChannel` with proper write persistence that honors per-channel
`snapshot_frequency`, plus closes two latent bugs the force-snapshot was
masking.

Before: every exit-mode run wrote a full `_DeltaSnapshot` blob for every
delta channel, even when the channel had zero updates this run and was
nowhere near its `snapshot_frequency`. After: the same count-based
decision used by `durability="sync"`/`"async"` applies — channels at or
above `snapshot_frequency` snapshot; channels below it persist their
accumulated writes via a lazy "stub" anchor; untouched channels write
nothing.

## What changed

**Core redesign** (`pregel/_loop.py`, `pregel/_checkpoint.py`)

- Drop `force_delta_snapshot` from `create_checkpoint` and
`_should_snapshot_delta`.
- Add `decide_delta_snapshots(channels, counts)` pure helper used by
both `create_checkpoint` and the new exit-mode peek-ahead path.
- Add `_exit_delta_writes` accumulator: every delta-channel write
produced during a `durability="exit"` run (input writes from `_first` +
per-superstep writes captured before `pending_writes.clear()` in
`after_tick`) is collected into this list.
- Add `_put_exit_delta_writes` (sync + async): runs from
`_suppress_interrupt` BEFORE `_put_checkpoint(exiting=True)`. Filters
out channels that will snapshot, then persists remaining writes to
`checkpoint_writes` under an anchor parent. The anchor is the existing
saved parent on resumed runs, or a lazily-created empty stub on first
runs.
- Visibility ordering: stub put goes onto `_put_checkpoint_fut` (becomes
the next put's `prev`); exit-write futures go onto `_delta_write_futs`.
The existing `_checkpointer_put_after_previous` already drains both
before calling `saver.put`, so `final_checkpoint` is structurally
guaranteed to land last — readers never see a partial view.

**Latent bugs fixed (previously masked by force-snapshot)**

- **Sync drain race**: `SyncPregelLoop` now initializes
`_delta_write_futs = []` in `__enter__` and drains it in sync
`_checkpointer_put_after_previous` before `put`, mirroring the async
version. Without this, a multi-worker `BackgroundExecutor` could publish
a checkpoint before the writes that produced it.
- **Count double-bump in exit mode**: in `_put_checkpoint`,
`delta_updates_since_snapshot` was being incremented twice for the last
superstep — once by the intermediate `after_tick` call, once by
`_suppress_interrupt`. Force-snapshot used to reset all counts to 0 so
this never persisted; without it, snapshots would fire one superstep
early after every exit-mode run. Fixed by gating the count-bump behind
`not exiting`.

**Pre-existing input-durability gap**

- In the plain (non-Command) input path of `_first`, delta-channel input
writes are now persisted via `put_writes` (mirroring the Command path),
so sub-frequency inputs survive a `get_state` on resumed runs in
`sync`/`async` durability. Note: first-run `sync`/`async` still has the
same gap (writes orphan on the synthetic-empty parent id). That's
flagged as a follow-up — out of scope for this PR.

## Test plan

- Existing `tests/test_pregel.py` and `tests/test_pregel_async.py` pass
unchanged.
- Existing `tests/test_channels.py` (29 tests) and
`tests/test_delta_channel_migration.py` pass unchanged.
- New `tests/test_exit_delta_persistence.py` (11 tests) covers:
- **Write-path**: zero-write exit (no stub), all-snapshot first run (no
stub), sub-freq first run (single shared stub), sub-freq resumed run
(anchor on saved parent), sync-vs-exit count parity, mixed
snapshot/non-snapshot channels, snapshot fires at frequency.
- **Read-path**: K-run replay chain reads correctly across
stub→saved-parent transition; metadata `delta_updates_since_snapshot`
round-trips correctly; mixed sync/exit durability alternation produces
correct final state; snapshot+tail-deltas combination reads correctly.
- `make format && make lint && make test` in `libs/langgraph/`.

---------

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: Sydney Runkle <sydneymarierunkle@gmail.com>
This commit is contained in:
Quanzheng Long
2026-05-07 09:47:01 -07:00
committed by GitHub
co-authored by Cursor Sydney Runkle
parent 95b41d058f
commit 69f2d3a430
3 changed files with 617 additions and 79 deletions
+36 -58
View File
@@ -34,28 +34,23 @@ def empty_checkpoint() -> Checkpoint:
)
def _should_snapshot_delta(
name: str,
ch: DeltaChannel,
updates_since_snapshot: Mapping[str, int],
*,
force: bool,
) -> bool:
"""Decide whether `ch` should write a `_DeltaSnapshot` this step.
def delta_channels_to_snapshot(
channels: Mapping[str, BaseChannel],
counts: Mapping[str, int],
) -> set[str]:
"""Return the set of DeltaChannel names that should snapshot now.
Triggers:
* `force` — always snapshot (used by `durability="exit"`).
* Update-count: this channel has accumulated at least
`snapshot_frequency` updates since its last snapshot. The count
is supplied by the caller via `updates_since_snapshot[name]` and
is reset to `0` whenever a snapshot fires.
Version-format-independent: works for `int`, `float`, and `str`
versioning schemes alike.
A channel snapshots when its accumulated update count (since the last
snapshot) reaches or exceeds `snapshot_frequency`. This is a pure
predicate — no mutation.
"""
if force:
return True
return updates_since_snapshot.get(name, 0) >= ch.snapshot_frequency
return {
name
for name, ch in channels.items()
if isinstance(ch, DeltaChannel)
and ch.is_available()
and counts.get(name, 0) >= ch.snapshot_frequency
}
def create_checkpoint(
@@ -66,34 +61,19 @@ def create_checkpoint(
id: str | None = None,
updated_channels: set[str] | None = None,
get_next_version: GetNextVersion | None = None,
force_delta_snapshot: bool = False,
updates_since_snapshot: Mapping[str, int] | None = None,
new_updates_since_snapshot: dict[str, int] | None = None,
channels_to_snapshot: set[str] | None = None,
) -> Checkpoint:
"""Create a checkpoint for the given channels.
"""Build a new Checkpoint from the previous one and live channel state.
For each `DeltaChannel`, a `_DeltaSnapshot(value)` blob is written into
`channel_values[k]` when this channel has accumulated at least
`snapshot_frequency` updates since its last snapshot (counter supplied
via `updates_since_snapshot`). Otherwise the channel is omitted from
`channel_values`; its `channel_versions` entry still bumps so that the
saver tracks the channel and the ancestor walk can replay writes.
Snapshots are eager: even if the channel had no write this step, a
version bump is forced (via `get_next_version`) so `put()` includes
the channel in `new_versions` and stores the blob.
`force_delta_snapshot` ignores the cadence and always snapshots —
used by `durability="exit"` where intermediate writes are not stored
as ancestor `checkpoint_writes`.
If `new_updates_since_snapshot` is provided, the function resets the
counter to `0` for any channel that snapshotted this step. Counters
for channels that did not snapshot are left untouched (the caller is
responsible for incrementing them based on `updated_channels`).
For each name in `channels_to_snapshot`, a `_DeltaSnapshot(value)` blob
is written into `channel_values[k]`. Other delta channels are omitted
from `channel_values` — the ancestor walk reconstructs their state
from `checkpoint_writes`. Callers compute the set via
`delta_channels_to_snapshot(channels, counts)`; defaults to empty
(no snapshots) when not provided.
"""
ts = datetime.now(timezone.utc).isoformat()
counts = updates_since_snapshot or {}
channels_to_snapshot = channels_to_snapshot or set()
if channels is None:
values = checkpoint["channel_values"]
channel_versions = checkpoint["channel_versions"]
@@ -104,25 +84,23 @@ def create_checkpoint(
if k not in channel_versions:
continue
ch = channels[k]
if (
isinstance(ch, DeltaChannel)
and ch.is_available()
and _should_snapshot_delta(
k,
ch,
counts,
force=force_delta_snapshot,
)
):
# Eager snapshot: bump version if not already written this step
# so put() includes this channel in new_versions and stores blob.
if k in channels_to_snapshot:
# In exit mode, the snapshot decision is deferred to exit
# time (intermediate steps have do_checkpoint=False). The
# channel's count may have reached snapshot_frequency over
# several supersteps, but the LAST superstep may not have
# written to this channel. In that case apply_writes()
# (in _algo.py) didn't bump this channel's version, so
# saver.put() wouldn't include it in new_versions and
# the snapshot blob would be silently dropped. The manual
# bump below closes the gap. In sync/async durability this
# branch is effectively dead code (the step that pushes
# the count to freq always writes the channel).
if get_next_version is not None and (
updated_channels is None or k not in updated_channels
):
channel_versions[k] = get_next_version(channel_versions[k], None)
values[k] = _DeltaSnapshot(ch.get())
if new_updates_since_snapshot is not None:
new_updates_since_snapshot[k] = 0
else:
v = ch.checkpoint()
if v is not MISSING:
+216 -21
View File
@@ -100,6 +100,7 @@ from langgraph.pregel._checkpoint import (
channels_from_checkpoint,
copy_checkpoint,
create_checkpoint,
delta_channels_to_snapshot,
empty_checkpoint,
)
from langgraph.pregel._executor import (
@@ -194,8 +195,40 @@ class PregelLoop:
_migrate_checkpoint: Callable[[Checkpoint], None] | None
submit: Submit
channels: Mapping[str, BaseChannel]
# Only set on AsyncPregelLoop; sync loops keep this as None.
# Futures from `checkpointer.put_writes` calls that produced delta-channel
# writes. `_checkpointer_put_after_previous` drains this list (swap to a
# local `futs` then reset to `[]` and wait/gather) before putting the
# next checkpoint, so a checkpoint never becomes durable before the
# writes that produced it. Initialised to `[]` in both sync and async
# `__enter__`; stays `None` only when no checkpointer.
_delta_write_futs: list[Any] | None = None
# Exit-mode accumulator: every delta-channel write produced during this
# run (input writes from `_first` + per-superstep writes captured in
# `after_tick`). At exit, `_put_exit_delta_writes` filters out channels
# that will snapshot, then persists the rest under an anchor parent.
# `None` when not in exit mode (so the capture sites are no-ops).
# Each tuple is `(step, task_id, channel, value)` — `step` drives the
# synthetic step-prefixed task_id used to preserve chronological order
# under the saver's `ORDER BY task_id, idx` sorting.
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
# The checkpoint_config that points at the parent loaded at `__enter__`
# (or the synthetic-empty checkpoint, on first run). We capture it
# eagerly because every `_put_checkpoint` advances `self.checkpoint_config`
# to the newly-saved checkpoint's id — by exit time the original parent
# config would otherwise be lost. `_put_exit_delta_writes` uses this:
# on resumed runs as the anchor for exit delta writes; on first runs
# to derive the lazy stub's config (its `checkpoint_id` is the
# synthetic-empty id we want the stub persisted under).
_initial_checkpoint_config: RunnableConfig
# True iff the saver actually returned a tuple at `__enter__`. False
# on the first-ever run for a thread (no parent persisted yet).
# `_put_exit_delta_writes` uses this to decide between anchoring on
# the existing parent (True) or creating a lazy stub (False).
_has_persisted_parent: bool = False
managed: ManagedValueMapping
checkpoint: Checkpoint
checkpoint_id_saved: str
@@ -637,6 +670,11 @@ class PregelLoop:
self._emit(
"values", map_output_values, self.output_keys, writes, self.channels
)
# capture delta-channel writes for exit-mode accumulator before clearing
if self._exit_delta_writes is not None:
for tid, ch, v in self.checkpoint_pending_writes:
if isinstance(self.specs.get(ch), DeltaChannel):
self._exit_delta_writes.append((self.step, tid, ch, v))
# clear pending writes
self.checkpoint_pending_writes.clear()
# only replay (re-execute) done tasks on the first tick
@@ -854,6 +892,27 @@ class PregelLoop:
self.checkpointer_get_next_version,
self.trigger_to_nodes,
)
# Input writes go through `apply_writes` directly (above) — they
# never enter `checkpoint_pending_writes`, so the after_tick
# capture site does not see them. In exit mode, capture them
# here so `_exit_delta_writes` includes the input's delta writes
# alongside per-superstep writes; otherwise the input would be
# lost on read (it's not in final_checkpoint.channel_values for
# sub-freq channels, and walks ignore target.pending_writes).
if self._exit_delta_writes is not None:
for c, v in input_writes:
if isinstance(self.specs.get(c), DeltaChannel):
self._exit_delta_writes.append((self.step, NULL_TASK_ID, c, v))
# Persist delta-channel input writes so sub-freq inputs are
# recoverable via ancestor walk (mirrors the Command input path).
if self.durability != "exit":
delta_input = [
(c, v)
for c, v in input_writes
if isinstance(self.specs.get(c), DeltaChannel)
]
if delta_input:
self.put_writes(NULL_TASK_ID, delta_input)
# save input checkpoint
self.updated_channels = updated_channels
self._put_checkpoint({"source": "input"})
@@ -905,36 +964,60 @@ class PregelLoop:
return updated_channels
def _put_checkpoint(self, metadata: CheckpointMetadata) -> None:
# assign step and parents
# `is` (object identity) — not `==`. Three of four call sites pass a
# fresh dict ({"source":"input"|"loop"|"fork"}); only
# `_suppress_interrupt`(will rename to _on_loop_exit soon)
# at exit reuses the existing `self.checkpoint_metadata` instance. So
# `metadata is self.checkpoint_metadata` is True only on the exit call,
# which is what we use to gate exit-only behaviour (skip count-bump,
# don't replace metadata). Could be replaced by an explicit
# `exiting: bool = False` parameter; left as-is to match the existing
# idiom in this file.
# TODO: replace with an explicit `exiting: bool = False` parameter.
exiting = metadata is self.checkpoint_metadata
if exiting and self.checkpoint["id"] == self.checkpoint_id_saved:
# checkpoint already saved
return
# Carry per-delta-channel update bookkeeping forward across
# supersteps. Capture from the OLD metadata before potentially
# replacing it with a fresh dict that wouldn't contain it. Then
# increment for any delta channel updated this step (so the count
# reflects "supersteps that wrote to this channel since last
# snapshot"). create_checkpoint will reset entries to 0 for any
# channel that fires a snapshot this step.
prev_counts = dict(
self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
)
new_counts = dict(prev_counts)
if self.updated_channels:
for ch_name in self.updated_channels:
ch_obj = self.channels.get(ch_name)
if isinstance(ch_obj, DeltaChannel):
new_counts[ch_name] = new_counts.get(ch_name, 0) + 1
# Per-delta-channel update bookkeeping.
#
# `_put_checkpoint` is called once per superstep with a fresh
# metadata dict (source="input"|"loop"|"fork") — those are the
# intermediate calls that bump the count by +1 for each delta
# channel touched that step. In exit mode,
# `_suppress_interrupt`(will rename to _on_loop_exit soon)
# additionally calls `_put_checkpoint(self.checkpoint_metadata)` AT
# EXIT to commit the final checkpoint — this runs *after* the last
# intermediate call already counted the last superstep. So the
# exit call must NOT bump again or it would double-count the last
# superstep. (Sync/async durability does not call `_put_checkpoint`
# at exit, so the issue only surfaces in exit mode. force_delta_snapshot
# used to mask this latent bug by resetting every count to 0.)
if not exiting:
prev_counts = dict(
self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
)
new_counts = dict(prev_counts)
if self.updated_channels:
for ch_name in self.updated_channels:
if isinstance(self.channels.get(ch_name), DeltaChannel):
new_counts[ch_name] = new_counts.get(ch_name, 0) + 1
metadata["step"] = self.step
metadata["parents"] = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {})
self.checkpoint_metadata = metadata
else:
new_counts = dict(
self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
)
# do checkpoint?
do_checkpoint = self._checkpointer_put_after_previous is not None and (
exiting or self.durability != "exit"
)
# create new checkpoint
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, new_counts)
if do_checkpoint
else set()
)
self.checkpoint = create_checkpoint(
self.checkpoint,
self.channels if do_checkpoint else None,
@@ -944,10 +1027,10 @@ class PregelLoop:
get_next_version=self.checkpointer_get_next_version
if do_checkpoint
else None,
force_delta_snapshot=exiting and self.durability == "exit",
updates_since_snapshot=new_counts,
new_updates_since_snapshot=new_counts,
channels_to_snapshot=channels_to_snapshot,
)
for k in channels_to_snapshot:
new_counts[k] = 0
if new_counts:
self.checkpoint_metadata["delta_updates_since_snapshot"] = new_counts
elif "delta_updates_since_snapshot" in self.checkpoint_metadata:
@@ -1010,6 +1093,97 @@ class PregelLoop:
# increment step
self.step += 1
def _put_exit_delta_writes(self) -> None:
"""Stage stub + accumulated delta writes so final_checkpoint's put
waits on them (visibility invariant: both must be durable before
final_checkpoint becomes visible to readers).
Stub is created lazily — only when no persisted parent exists AND at
least one delta channel has writes that won't be snapshotted.
"""
if (
not self._exit_delta_writes
or self.checkpointer is None
or self._checkpointer_put_after_previous is None
or self.checkpointer_put_writes is None
):
return
counts = self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
channels_to_snapshot = delta_channels_to_snapshot(self.channels, counts)
pending = [
(step, tid, ch, v)
for (step, tid, ch, v) in self._exit_delta_writes
if ch not in channels_to_snapshot
]
if not pending:
return
if self._has_persisted_parent:
# _initial_checkpoint_config's checkpoint_id is the saved parent's
# id (saver returned a real tuple at __enter__).
anchor_config = self._initial_checkpoint_config
else:
stub_cp = empty_checkpoint()
stub_cp["id"] = self.checkpoint_id_saved
stub_cp["ts"] = datetime.now(timezone.utc).isoformat()
# Stub has no parent (checkpoint_id=None in config).
stub_put_config = patch_configurable(
self._initial_checkpoint_config,
{CONFIG_KEY_CHECKPOINT_ID: None},
)
# Anchor config for put_writes: checkpoint_id = stub's id.
anchor_config = patch_configurable(
self._initial_checkpoint_config,
{CONFIG_KEY_CHECKPOINT_ID: stub_cp["id"]},
)
self._put_checkpoint_fut = self.submit(
self._checkpointer_put_after_previous,
getattr(self, "_put_checkpoint_fut", None),
stub_put_config,
stub_cp,
{"step": -2},
{},
)
# Set checkpoint_config so final_checkpoint's _put_checkpoint
# sees the stub as its parent.
self.checkpoint_config = anchor_config
# Step-prefixed synthetic task_id preserves chronological superstep
# order under the saver's ORDER BY task_id, idx sorting.
grouped: dict[tuple[int, str], list[tuple[str, Any]]] = {}
for step, tid, ch, v in pending:
grouped.setdefault((step, tid), []).append((ch, v))
anchor_write_config = patch_configurable(
anchor_config,
{
CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get(
CONFIG_KEY_CHECKPOINT_NS, ""
),
CONFIG_KEY_CHECKPOINT_ID: anchor_config[CONF][CONFIG_KEY_CHECKPOINT_ID],
},
)
for (step, tid), entries in grouped.items():
synth_tid = f"{step:08d}-{tid}"
if self.checkpointer_put_writes_accepts_task_path:
fut = self.submit(
self.checkpointer_put_writes,
anchor_write_config,
entries,
synth_tid,
"",
)
else:
fut = self.submit(
self.checkpointer_put_writes,
anchor_write_config,
entries,
synth_tid,
)
if self._delta_write_futs is not None:
self._delta_write_futs.append(fut)
def _suppress_interrupt(
self,
exc_type: type[BaseException] | None,
@@ -1025,6 +1199,7 @@ class PregelLoop:
# or a nested graph with checkpointer=True
or all(NS_END not in part for part in self.checkpoint_ns)
):
self._put_exit_delta_writes()
self._put_checkpoint(self.checkpoint_metadata)
self._put_pending_writes()
# suppress interrupt
@@ -1230,6 +1405,9 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
if self._delta_write_futs:
futs, self._delta_write_futs = self._delta_write_futs, []
concurrent.futures.wait(futs)
try:
if prev is not None:
prev.result()
@@ -1347,6 +1525,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
# graph/thread. Returns None on first invocation.
saved = self.checkpointer.get_tuple(self.checkpoint_config)
# Capture before the synthetic-empty fallback below overwrites `saved`.
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
# to lazy-create a stub instead of anchoring delta writes on a parent.
self._has_persisted_parent = saved is not None
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1362,6 +1544,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
**saved.config.get(CONF, {}),
},
}
self._initial_checkpoint_config = self.checkpoint_config
self.prev_checkpoint_config = saved.parent_config
self.checkpoint_id_saved = saved.checkpoint["id"]
self.checkpoint = saved.checkpoint
@@ -1371,6 +1554,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if saved.pending_writes is not None
else []
)
self._delta_write_futs = []
self._exit_delta_writes = (
[] if self.durability == "exit" and self.checkpointer is not None else None
)
self.submit = self.stack.enter_context(BackgroundExecutor(self.config))
self.channels, self.managed = channels_from_checkpoint(
self.specs,
@@ -1596,6 +1783,10 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
# graph/thread. Returns None on first invocation.
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
# Capture before the synthetic-empty fallback below overwrites `saved`.
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
# to lazy-create a stub instead of anchoring delta writes on a parent.
self._has_persisted_parent = saved is not None
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1611,6 +1802,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
**saved.config.get(CONF, {}),
},
}
self._initial_checkpoint_config = self.checkpoint_config
self.prev_checkpoint_config = saved.parent_config
self.checkpoint_id_saved = saved.checkpoint["id"]
self.checkpoint = saved.checkpoint
@@ -1621,6 +1813,9 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
else []
)
self._delta_write_futs = []
self._exit_delta_writes = (
[] if self.durability == "exit" and self.checkpointer is not None else None
)
self.submit = await self.stack.enter_async_context(
AsyncBackgroundExecutor(self.config)
)
@@ -0,0 +1,365 @@
"""Tests for exit-mode delta channel persistence redesign.
Validates that `durability="exit"` correctly persists delta-channel writes
using count-based snapshot decisions (rather than force-snapshotting every
channel), lazy stub creation when no parent exists, and proper read-path
reconstruction via ancestor walks.
"""
from typing import Annotated, Any
import pytest
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
pytestmark = pytest.mark.anyio
def _build_graph(
checkpointer: InMemorySaver,
*,
freq: int = 1000,
) -> Any:
channel = DeltaChannel(_messages_delta_reducer, snapshot_frequency=freq)
# Functional TypedDict form: class form can't reference `channel` (a
# local variable) inside Annotated due to forward-ref evaluation rules.
State = TypedDict("State", {"messages": Annotated[list, channel]}) # type: ignore[call-overload] # noqa: UP013
def respond(state: dict) -> dict:
i = len(state["messages"])
return {"messages": [AIMessage(content=f"reply-{i}", id=f"ai{i}")]}
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
return builder.compile(checkpointer=checkpointer)
# ---------------------------------------------------------------------------
# 8a. Write-path / structural tests
# ---------------------------------------------------------------------------
async def test_exit_first_run_no_delta_writes() -> None:
"""Graph with delta channel invoked with input that doesn't touch it.
Only one checkpoint row, no stub."""
State = TypedDict( # noqa: UP013
"State",
{
"messages": Annotated[list, DeltaChannel(_messages_delta_reducer)],
"value": str,
},
) # type: ignore[call-overload]
def noop(state: dict) -> dict:
return {"value": "done"}
saver = InMemorySaver()
builder = StateGraph(State)
builder.add_node("noop", noop)
builder.add_edge(START, "noop")
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "no-delta-writes"}}
graph.invoke({"value": "start"}, config, durability="exit")
checkpoints = list(saver.list(config))
assert len(checkpoints) == 1
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
assert len(stubs) == 0
async def test_exit_first_run_all_snapshot() -> None:
"""snapshot_frequency=1 forces every channel to snapshot.
No stub needed; final_checkpoint has _DeltaSnapshot."""
saver = InMemorySaver()
graph = _build_graph(saver, freq=1)
config = {"configurable": {"thread_id": "all-snapshot"}}
result = graph.invoke(
{"messages": [HumanMessage(content="hi", id="h1")]},
config,
durability="exit",
)
assert len(result["messages"]) == 2
checkpoints = list(saver.list(config))
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
assert len(stubs) == 0
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
state = graph.get_state(config)
assert [m.content for m in state.values["messages"]] == ["hi", "reply-1"]
async def test_exit_first_run_sub_freq_with_writes() -> None:
"""First run with default snapshot_frequency (1000), writes below threshold.
A stub is created; writes are anchored under it; get_state reconstructs."""
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "sub-freq-first"}}
result = graph.invoke(
{"messages": [HumanMessage(content="hello", id="h1")]},
config,
durability="exit",
)
assert [m.content for m in result["messages"]] == ["hello", "reply-1"]
checkpoints = list(saver.list(config))
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
assert len(stubs) == 1, f"Expected 1 stub, got {len(stubs)}"
head = saver.get_tuple(config)
assert head is not None
assert "messages" not in head.checkpoint["channel_values"]
assert "messages" in head.checkpoint["channel_versions"]
state = graph.get_state(config)
assert [m.content for m in state.values["messages"]] == ["hello", "reply-1"]
async def test_exit_resumed_run_sub_freq() -> None:
"""Two consecutive exit runs. Second run anchors on the first's
final_checkpoint (no new stub). Ordering preserved."""
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "resumed-sub-freq"}}
graph.invoke(
{"messages": [HumanMessage(content="msg1", id="h1")]},
config,
durability="exit",
)
graph.invoke(
{"messages": [HumanMessage(content="msg2", id="h2")]},
config,
durability="exit",
)
checkpoints = list(saver.list(config))
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
assert len(stubs) == 1
state = graph.get_state(config)
contents = [m.content for m in state.values["messages"]]
assert len(contents) == 4
assert contents[0] == "msg1"
assert contents[2] == "msg2"
assert contents[0:4:2] == ["msg1", "msg2"]
async def test_exit_count_parity_sync_vs_exit() -> None:
"""Sync and exit durability produce the same delta_updates_since_snapshot
after an equivalent run."""
for durability in ("sync", "exit"):
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": f"parity-{durability}"}}
graph.invoke(
{"messages": [HumanMessage(content="hi", id="h1")]},
config,
durability=durability,
)
head = saver.get_tuple(config)
assert head is not None
counts = head.metadata.get("delta_updates_since_snapshot", {})
assert counts.get("messages") == 2, (
f"durability={durability}: expected count=2, got {counts}"
)
async def test_exit_snapshot_fires_at_frequency() -> None:
"""With snapshot_frequency=3, after 3 exit runs (each incrementing count
by 2: input + superstep), the 2nd run hits count=4>=3, triggering snapshot.
After that run, count resets to 0 and channel_values has _DeltaSnapshot."""
saver = InMemorySaver()
graph = _build_graph(saver, freq=3)
config = {"configurable": {"thread_id": "snapshot-at-freq"}}
graph.invoke(
{"messages": [HumanMessage(content="m1", id="h1")]},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
count1 = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0)
assert count1 == 2
graph.invoke(
{"messages": [HumanMessage(content="m2", id="h2")]},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
count2 = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0)
assert count2 == 0, f"Expected reset to 0 after snapshot, got {count2}"
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
async def test_exit_mixed_snapshot_and_non_snapshot() -> None:
"""One delta channel at freq=1 (always snapshot) and one at freq=1000
(never snapshot within this test). Verify correct behavior for both."""
fast_ch = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1)
slow_ch = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1000)
State = TypedDict( # noqa: UP013
"State",
{"fast": Annotated[list, fast_ch], "slow": Annotated[list, slow_ch]},
) # type: ignore[call-overload]
def respond(state: dict) -> dict:
return {
"fast": [AIMessage(content="fast-reply", id="f1")],
"slow": [AIMessage(content="slow-reply", id="s1")],
}
saver = InMemorySaver()
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "mixed-freq"}}
graph.invoke(
{
"fast": [HumanMessage(content="fast-in", id="fi")],
"slow": [HumanMessage(content="slow-in", id="si")],
},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("fast"), _DeltaSnapshot)
assert "slow" not in head.checkpoint["channel_values"]
state = graph.get_state(config)
assert [m.content for m in state.values["fast"]] == ["fast-in", "fast-reply"]
assert [m.content for m in state.values["slow"]] == ["slow-in", "slow-reply"]
# ---------------------------------------------------------------------------
# 8b. Read-path tests
# ---------------------------------------------------------------------------
async def test_exit_multi_run_replay_chain() -> None:
"""K=4 consecutive exit runs, each adding a message. After each run,
get_state returns all messages in chronological order."""
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "replay-chain"}}
for i in range(4):
graph.invoke(
{"messages": [HumanMessage(content=f"user-{i}", id=f"h{i}")]},
config,
durability="exit",
)
state = graph.get_state(config)
contents = [m.content for m in state.values["messages"]]
user_msgs = [c for c in contents if c.startswith("user-")]
assert user_msgs == [f"user-{j}" for j in range(i + 1)], (
f"After run {i}: user messages out of order: {user_msgs}"
)
assert len(contents) == (i + 1) * 2
async def test_exit_metadata_round_trip() -> None:
"""K=5 consecutive exit runs with snapshot_frequency=5. Verify metadata
delta_updates_since_snapshot increments correctly across runs."""
freq = 5
saver = InMemorySaver()
graph = _build_graph(saver, freq=freq)
config = {"configurable": {"thread_id": "metadata-rt"}}
for i in range(1, 6):
graph.invoke(
{"messages": [HumanMessage(content=f"m{i}", id=f"h{i}")]},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
count = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0)
cumulative = i * 2
if cumulative >= freq:
assert count == 0 or count == cumulative % freq or count < freq, (
f"After run {i}: count={count} should have reset or be partial"
)
else:
assert count == cumulative, (
f"After run {i}: expected {cumulative}, got {count}"
)
async def test_exit_mixed_durability_round_trip() -> None:
"""Alternate sync and exit durability; verify counts stay monotonic
and state accumulates correctly."""
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "mixed-durability"}}
for i, dur in enumerate(["sync", "exit", "sync", "exit"]):
graph.invoke(
{"messages": [HumanMessage(content=f"msg-{i}", id=f"h{i}")]},
config,
durability=dur,
)
state = graph.get_state(config)
contents = [m.content for m in state.values["messages"]]
user_msgs = [c for c in contents if c.startswith("msg-")]
assert user_msgs == [f"msg-{j}" for j in range(i + 1)], (
f"After run {i} (durability={dur}): {user_msgs}"
)
assert len(contents) == (i + 1) * 2
async def test_exit_snapshot_then_tail_deltas() -> None:
"""Run 1 forces snapshot (freq=1). Run 2 at freq=1000 adds more writes
that don't snapshot. Reading after run 2 must combine the snapshot seed
with the tail deltas."""
saver = InMemorySaver()
graph1 = _build_graph(saver, freq=1)
config = {"configurable": {"thread_id": "snapshot-then-tail"}}
graph1.invoke(
{"messages": [HumanMessage(content="seed-msg", id="h1")]},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
graph2 = _build_graph(saver, freq=1000)
graph2.invoke(
{"messages": [HumanMessage(content="tail-msg", id="h2")]},
config,
durability="exit",
)
state = graph2.get_state(config)
contents = [m.content for m in state.values["messages"]]
assert "seed-msg" in contents
assert "tail-msg" in contents
assert contents.index("seed-msg") < contents.index("tail-msg")