fix(langgraph): seal what a resume drops, and keep storage-only bumps out of scheduling

- A resume that reapplies the head's pending writes now seals the delta
  channels of the loaded writes no task of the run claimed, decided in
  `after_tick` once the tasks are known. `Command(resume=..., goto=[Send(...)])`
  replaces the fan-out, so the finished task's writes stayed on the head and
  the reload replayed them.
- One predicate, `_reapplies_pending_writes`, answers whether the loaded
  writes go back to their tasks, for every site that reapplies them.
- The snapshot cadence skips delta channels without a version. Past
  DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT, every checkpoint minted a version and
  an empty snapshot for each never-written delta channel.
- Versions minted only to store a snapshot are recorded under
  `SNAPSHOT_BUMPS` in `versions_seen`, and `update_state` skips them when it
  infers `as_node`. Advancing `versions_seen` over a bump made a raw `Pregel`
  subscriber look like the last writer after an exit-durability run.
- `_mark_bumps_seen` no longer synthesizes an `INTERRUPT` entry.
This commit is contained in:
Elior Nataf Lackritz
2026-10-01 15:14:22 -04:00
parent 858e55f232
commit e4d77cd222
7 changed files with 239 additions and 29 deletions
@@ -23,6 +23,8 @@ RETURN = sys.intern("__return__")
# for writes of a task where we simply record the return value
PREVIOUS = sys.intern("__previous__")
# the implicit branch that handles each node's Control values
SNAPSHOT_BUMPS = sys.intern("__snapshot_bumps__")
# `versions_seen` key for channel versions minted only to store a snapshot
# --- Reserved cache namespaces ---
@@ -116,6 +118,7 @@ RESERVED = {
ERROR,
ERROR_SOURCE_NODE,
NO_WRITES,
SNAPSHOT_BUMPS,
# reserved config.configurable keys
CONFIG_KEY_SEND,
CONFIG_KEY_READ,
+41 -14
View File
@@ -16,7 +16,7 @@ from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from langgraph._internal._config import DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT
from langgraph._internal._constants import INTERRUPT, PUSH
from langgraph._internal._constants import PUSH, SNAPSHOT_BUMPS
from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel
from langgraph.channels.delta import DeltaChannel
@@ -52,17 +52,23 @@ def exit_delta_task_id(step: int, task_id: str) -> str:
def delta_channels_to_snapshot(
channels: Mapping[str, BaseChannel],
counters_since_delta_snapshot: Mapping[str, tuple[int, int]],
channel_versions: ChannelVersions,
) -> set[str]:
"""Return the set of DeltaChannel names that should snapshot now.
A channel snapshots when EITHER its accumulated update count reaches
`snapshot_frequency` OR the total supersteps since its last snapshot
reaches `DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT`. This is a pure
predicate — no mutation.
reaches `DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT`. A channel without a version
was never written on this branch, so it has nothing to snapshot. This is a
pure predicate — no mutation.
"""
result: set[str] = set()
for name, ch in channels.items():
if not isinstance(ch, DeltaChannel) or not ch.is_available():
if (
not isinstance(ch, DeltaChannel)
or not ch.is_available()
or name not in channel_versions
):
continue
updates, supersteps = counters_since_delta_snapshot.get(name, (0, 0))
if (
@@ -142,6 +148,7 @@ def create_checkpoint_plan_for_update_state_api(
saved_metadata: Mapping[str, Any] | None,
is_fresh_thread: bool,
fork_channels: set[str],
channel_versions: ChannelVersions,
) -> tuple[set[str], dict[str, Any]]:
"""Return ``(channels_to_snapshot, metadata)`` for an update_state head."""
metadata: dict[str, Any] = {
@@ -158,7 +165,8 @@ def create_checkpoint_plan_for_update_state_api(
prev_metadata=saved_metadata,
)
channels_to_snapshot = (
delta_channels_to_snapshot(channels, new_counters) | fork_channels
delta_channels_to_snapshot(channels, new_counters, channel_versions)
| fork_channels
)
for k in channels_to_snapshot:
new_counters[k] = (0, 0)
@@ -205,9 +213,7 @@ def create_checkpoint(
if k in channels_to_snapshot and get_next_version is not None:
channel_versions[k] = get_next_version(None, None)
bumped[k] = (None, channel_versions[k])
values[k] = _DeltaSnapshot(
ch.get() if ch.is_available() else ch.typ()
)
values[k] = _DeltaSnapshot(ch.get())
continue
if k in channels_to_snapshot:
# `put` only stores a blob for a channel whose version moved,
@@ -243,18 +249,39 @@ def _mark_bumps_seen(
"""Advance whoever had seen a bumped channel's old version to the new one.
A bump that only stores a snapshot is not a write. Left unseen, it would
re-fire `interrupt_before` and rerun the channel's subscribers.
re-fire `interrupt_before` and rerun the channel's subscribers. The bumped
versions are also kept under `SNAPSHOT_BUMPS`, so inferring which node
wrote last can skip them.
"""
if not bumped:
return versions_seen
out: dict[str, ChannelVersions] = {}
for node, seen in {INTERRUPT: {}, **versions_seen}.items():
advanced = {k: new for k, (old, new) in bumped.items() if seen.get(k) == old}
if advanced or node in versions_seen:
out[node] = {**seen, **advanced}
out = {
node: {
**seen,
**{k: new for k, (old, new) in bumped.items() if seen.get(k) == old},
}
for node, seen in versions_seen.items()
}
out[SNAPSHOT_BUMPS] = {
**versions_seen.get(SNAPSHOT_BUMPS, {}),
**{k: new for k, (_, new) in bumped.items()},
}
return out
def versions_seen_without_bumps(
versions_seen: dict[str, ChannelVersions],
) -> dict[str, ChannelVersions]:
"""`versions_seen` without the versions minted only to store a snapshot."""
if not (bumps := versions_seen.get(SNAPSHOT_BUMPS)):
return versions_seen
return {
node: {k: v for k, v in seen.items() if bumps.get(k) != v}
for node, seen in versions_seen.items()
if node != SNAPSHOT_BUMPS
}
def _needs_replay(spec: BaseChannel, stored: object) -> bool:
"""True if `spec` is a `DeltaChannel` and no value is stored at this
checkpoint, requiring an ancestor walk to reconstruct.
+38 -9
View File
@@ -230,6 +230,9 @@ class PregelLoop:
# * the checkpoint this run starts from has pending writes to them; see
# `delta_channels_with_pending_writes`.
_delta_channels_forced_snapshot: set[str]
# Set by `_first` for a resume whose loaded writes `after_tick` still has
# to check against the tasks that ran.
_seal_unclaimed_writes: bool = False
# The checkpoint_config that points at the parent loaded at `__enter__`
# (or the synthetic-empty checkpoint, on first run). We capture it
@@ -585,7 +588,7 @@ class PregelLoop:
# save the new task
self.tasks[pushed.id] = pushed
# match any pending writes to the new task
if not self.is_replaying:
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({pushed.id: pushed})
# return the new task, to be started if not run before
return pushed
@@ -663,7 +666,7 @@ class PregelLoop:
return False
# if there are pending writes from a previous loop, apply them
if not self.is_replaying and self.checkpoint_pending_writes:
if self._reapplies_pending_writes and self.checkpoint_pending_writes:
self._reapply_writes_to_succeeded_nodes(self.tasks)
self._resume_error_handlers_if_applicable()
@@ -692,6 +695,20 @@ class PregelLoop:
for ch, v in writes
if isinstance(self.specs.get(ch), DeltaChannel) and _get_overwrite(v)[0]
)
if self._seal_unclaimed_writes:
# A loaded write no task of this run claimed belongs to a task the
# resume dropped, such as a `Send` that `Command(goto=...)` replaced.
self._delta_channels_forced_snapshot.update(
delta_channels_with_pending_writes(
self.specs,
[
w
for w in self.checkpoint_pending_writes
if w[0] != NULL_TASK_ID and w[0] not in self.tasks
],
)
)
self._seal_unclaimed_writes = False
# all tasks have finished
self.updated_channels = apply_writes(
self.checkpoint,
@@ -737,6 +754,12 @@ class PregelLoop:
# private
@property
def _reapplies_pending_writes(self) -> bool:
"""Whether the writes loaded with the checkpoint go back to the tasks
that made them, instead of those tasks rerunning."""
return not self.is_replaying
def _reapply_writes_to_succeeded_nodes(
self, tasks: Mapping[str, PregelExecutableTask]
) -> None:
@@ -902,11 +925,13 @@ class PregelLoop:
self.checkpoint_pending_writes = [
w for w in self.checkpoint_pending_writes if w[1] != RESUME
]
# A resume that is not replaying reuses the head's pending writes
# instead of rerunning their tasks, so none of them can leak.
# A resume that reapplies the head's pending writes only learns which
# of them its tasks claim once they are scheduled, so `after_tick`
# seals the rest.
self._seal_unclaimed_writes = is_resuming and self._reapplies_pending_writes
self._delta_channels_forced_snapshot = (
set()
if is_resuming and not self.is_replaying
if self._seal_unclaimed_writes
else delta_channels_with_pending_writes(
self.specs, self.checkpoint_pending_writes
)
@@ -1148,7 +1173,9 @@ class PregelLoop:
)
# create new checkpoint
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, new_counters)
delta_channels_to_snapshot(
self.channels, new_counters, self.checkpoint["channel_versions"]
)
| self._delta_channels_forced_snapshot
if do_checkpoint
else set()
@@ -1251,7 +1278,9 @@ class PregelLoop:
self.checkpoint_metadata.get("counters_since_delta_snapshot") or {}
)
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, counters)
delta_channels_to_snapshot(
self.channels, counters, self.checkpoint["channel_versions"]
)
| self._delta_channels_forced_snapshot
)
@@ -1613,7 +1642,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if handler_task is None:
return None
self.tasks[handler_task.id] = handler_task
if not self.is_replaying:
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
for task in self.match_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
@@ -1867,7 +1896,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
if handler_task is None:
return None
self.tasks[handler_task.id] = handler_task
if not self.is_replaying:
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
for task in await self.amatch_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
+12 -3
View File
@@ -136,6 +136,7 @@ from langgraph.pregel._checkpoint import (
delta_channels_with_pending_writes,
empty_checkpoint,
get_updated_channels_from_tasks,
versions_seen_without_bumps,
)
from langgraph.pregel._draw import draw_graph
from langgraph.pregel._io import map_input, read_channels
@@ -1936,7 +1937,9 @@ class Pregel(
as_node = tuple(self.nodes)[0]
elif as_node is None and not any(
v
for vv in checkpoint["versions_seen"].values()
for vv in versions_seen_without_bumps(
checkpoint["versions_seen"]
).values()
for v in vv.values()
):
if (
@@ -1947,7 +1950,9 @@ class Pregel(
elif as_node is None:
last_seen_by_node = sorted(
(v, n)
for n, seen in checkpoint["versions_seen"].items()
for n, seen in versions_seen_without_bumps(
checkpoint["versions_seen"]
).items()
if n in self.nodes
for v in seen.values()
)
@@ -2047,6 +2052,7 @@ class Pregel(
saved_metadata=saved.metadata if saved else None,
is_fresh_thread=saved is None,
fork_channels=fork_pending,
channel_versions=checkpoint["channel_versions"],
)
)
checkpoint = create_checkpoint(
@@ -2434,7 +2440,9 @@ class Pregel(
elif as_node is None:
last_seen_by_node = sorted(
(v, n)
for n, seen in checkpoint["versions_seen"].items()
for n, seen in versions_seen_without_bumps(
checkpoint["versions_seen"]
).items()
if n in self.nodes
for v in seen.values()
)
@@ -2534,6 +2542,7 @@ class Pregel(
saved_metadata=saved.metadata if saved else None,
is_fresh_thread=saved is None,
fork_channels=fork_pending,
channel_versions=checkpoint["channel_versions"],
)
)
checkpoint = create_checkpoint(
@@ -17,7 +17,14 @@ from typing_extensions import TypedDict
from langgraph._internal._constants import INPUT
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import END, START, StateGraph
from langgraph.types import Command, Durability, StateSnapshot, StateUpdate, interrupt
from langgraph.types import (
Command,
Durability,
Send,
StateSnapshot,
StateUpdate,
interrupt,
)
pytestmark = pytest.mark.anyio
@@ -470,6 +477,60 @@ def test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot(
assert not _snapshotted_checkpoints(sync_checkpointer, config)
def _build_send_fan_out(checkpointer: BaseCheckpointSaver) -> Any:
def q(state: _State) -> dict:
interrupt("continue?")
return _both("q")
builder = StateGraph(_State)
builder.add_node("p", lambda state: _both("p"))
builder.add_node("q", q)
builder.add_conditional_edges(
START, lambda state: [Send("p", state), Send("q", state)], ["p", "q"]
)
return builder.compile(checkpointer=checkpointer)
def test_resume_that_replaces_the_pending_sends_drops_the_finished_task(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
graph = _build_send_fan_out(sync_checkpointer)
config = _thread("t")
graph.invoke(_both("in"), config, durability=durability)
live = graph.invoke(
Command(resume="yes", goto=[Send("q", _both("unused"))]),
config,
durability=durability,
)
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"] == live["log"], (
f"'p' ran in the fan-out the resume replaced, but the reload reads "
f"{state.values['log']} against the live {live['log']}"
)
async def test_aresume_that_replaces_the_pending_sends_drops_the_finished_task(
async_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
graph = _build_send_fan_out(async_checkpointer)
config = _thread("t")
await graph.ainvoke(_both("in"), config, durability=durability)
live = await graph.ainvoke(
Command(resume="yes", goto=[Send("q", _both("unused"))]),
config,
durability=durability,
)
state = await graph.aget_state(config)
assert state.values["log"] == state.values["plain"] == live["log"], (
f"'p' ran in the fan-out the resume replaced, but the reload reads "
f"{state.values['log']} against the live {live['log']}"
)
def test_resume_addressed_at_an_interrupted_head_reruns_its_tasks_once(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
@@ -94,6 +94,28 @@ async def test_forced_snapshot_single_run() -> None:
assert "seed-a" in state.values["a"]
async def test_supersteps_bound_skips_a_channel_never_written() -> None:
with patch(
"langgraph.pregel._checkpoint.DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT",
3,
):
saver = InMemorySaver()
graph = _build_two_channel_graph(saver, n_loops=4)
config = {"configurable": {"thread_id": "never-written"}}
graph.invoke({"a": ["seed-a"]}, config)
minted = [
t.config["configurable"]["checkpoint_id"]
for t in saver.list(config)
if "b" in t.checkpoint["channel_versions"]
]
assert not minted, (
f"b was never written, but {len(minted)} checkpoints minted it a version"
)
assert graph.get_state(config).values["b"] == []
async def test_forced_snapshot_accumulates_across_runs() -> None:
"""Supersteps counter for an unwritten channel persists across separate
invoke() calls. After enough runs, the channel is force-snapshotted."""
@@ -139,13 +161,17 @@ async def test_predicate_fires_on_supersteps_overflow() -> None:
channels = {"x": ch_instance}
counters: dict[str, tuple[int, int]] = {"x": (0, 5000)}
result = delta_channels_to_snapshot(channels, counters)
result = delta_channels_to_snapshot(channels, counters, {"x": 1})
assert "x" in result
counters_below: dict[str, tuple[int, int]] = {"x": (0, 4999)}
result2 = delta_channels_to_snapshot(channels, counters_below)
result2 = delta_channels_to_snapshot(channels, counters_below, {"x": 1})
assert "x" not in result2
assert not delta_channels_to_snapshot(channels, counters, {}), (
"a channel with no version was never written, so it has nothing to snapshot"
)
async def test_counter_reset_after_supersteps_snapshot() -> None:
"""After the supersteps bound triggers a snapshot, the counters for
+55
View File
@@ -9410,6 +9410,61 @@ def test_fork_does_not_apply_pending_writes(
assert result == {"value": 121}
def _extend(state: list, writes: list[list]) -> list:
return [*state, *(v for write in writes for v in write)]
class _IntVersionSaver(InMemorySaver):
"""Integer versions tie exactly where `InMemorySaver`'s break at random."""
get_next_version = BaseCheckpointSaver.get_next_version
def _build_chain_after_a_delta_channel() -> Pregel:
return Pregel(
nodes={
"a": NodeBuilder().subscribe_only("inp").do(lambda _: ["a"]).write_to("d"),
"b": NodeBuilder().subscribe_only("d").do(lambda _: "b").write_to("x"),
"c": NodeBuilder().subscribe_only("x").do(lambda _: "c").write_to("out"),
},
channels={
"inp": LastValue(str),
"d": DeltaChannel(_extend, snapshot_frequency=1),
"x": LastValue(str),
"out": LastValue(str),
},
input_channels=["inp"],
output_channels=["out"],
checkpointer=_IntVersionSaver(),
)
def test_update_state_after_an_exit_snapshot_infers_the_last_writer() -> None:
graph = _build_chain_after_a_delta_channel()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"inp": "go"}, config, durability="exit")
graph.update_state(config, "u")
values = graph.get_state(config).values
assert values["out"] == "u", (
f"the update should apply as c, the last node to write, but state is {values}"
)
async def test_aupdate_state_after_an_exit_snapshot_infers_the_last_writer() -> None:
graph = _build_chain_after_a_delta_channel()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"inp": "go"}, config, durability="exit")
await graph.aupdate_state(config, "u")
values = (await graph.aget_state(config)).values
assert values["out"] == "u", (
f"the update should apply as c, the last node to write, but state is {values}"
)
async def test_delta_channel_end_to_end_inmemory() -> None:
"""Full graph run: DeltaChannel accumulates correctly across multiple turns."""