fix(langgraph): snapshot only what a fork can leak, and hide the bump

A fork used to snapshot every DeltaChannel whenever the caller passed a
checkpoint_id. That fired on every turn a client addresses the head
(storing a full copy of the channel per turn), missed new input sent to an
interrupted head without an id, and its storage-only version bump read as
a real write: interrupt_before fired again on resume, and a replay paused
at a node it had already passed.

Snapshot the delta channels the base checkpoint has pending writes for,
since only those can leak into a branch that does not consume them, and
advance versions_seen past any bump that only stores a snapshot, including
the interrupt tracker. update_state no longer records its narrower
updated_channels when it snapshots, so a deferred node listed in next still
runs on resume (#9089).
This commit is contained in:
Elior Nataf Lackritz
2026-09-28 19:54:23 -04:00
parent 9d5b0f1991
commit ddaf708cd0
5 changed files with 273 additions and 101 deletions
+61 -40
View File
@@ -8,13 +8,15 @@ from typing import Any, cast
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
ChannelVersions,
Checkpoint,
PendingWrite,
)
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 PUSH
from langgraph._internal._constants import INTERRUPT, PUSH
from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel
from langgraph.channels.delta import DeltaChannel
@@ -80,14 +82,29 @@ def get_updated_channels_from_tasks(
def get_delta_channels_from_all_channels(
channels: Mapping[str, BaseChannel],
*,
include_unavailable: bool = False,
) -> set[str]:
"""DeltaChannels to snapshot on the first update_state of a fresh thread or fork."""
"""DeltaChannels to snapshot on the first update_state of a fresh thread."""
return {
k
for k, ch in channels.items()
if isinstance(ch, DeltaChannel) and (include_unavailable or ch.is_available())
if isinstance(ch, DeltaChannel) and ch.is_available()
}
def delta_channels_with_pending_writes(
specs: Mapping[str, Any],
pending_writes: Iterable[PendingWrite] | None,
) -> set[str]:
"""DeltaChannels a branch starting from this checkpoint must snapshot.
A checkpoint's pending writes belong to the child that consumed them, and
nothing records which child that was. A new branch snapshots every delta
channel they touch, so its ancestor walk never replays them.
"""
return {
ch
for _, ch, _ in pending_writes or ()
if isinstance(specs.get(ch), DeltaChannel)
}
@@ -124,29 +141,25 @@ def create_checkpoint_plan_for_update_state_api(
parents: dict[str, Any],
saved_metadata: Mapping[str, Any] | None,
is_fresh_thread: bool,
is_fork: bool,
fork_channels: set[str],
) -> tuple[set[str], dict[str, Any]]:
"""Return ``(channels_to_snapshot, metadata)`` for an update_state head.
A fork snapshots everything, like a fresh thread: its base also holds the
writes of the branch it abandons, so the ancestor walk must stop here.
"""
"""Return ``(channels_to_snapshot, metadata)`` for an update_state head."""
metadata: dict[str, Any] = {
"source": "update",
"step": step,
"parents": parents,
}
if is_fresh_thread or is_fork:
return get_delta_channels_from_all_channels(
channels, include_unavailable=is_fork
), metadata
if is_fresh_thread:
return get_delta_channels_from_all_channels(channels), metadata
new_counters = create_metadata_for_update_state_api(
channels,
updated_channels,
prev_metadata=saved_metadata,
)
channels_to_snapshot = delta_channels_to_snapshot(channels, new_counters)
channels_to_snapshot = (
delta_channels_to_snapshot(channels, new_counters) | fork_channels
)
for k in channels_to_snapshot:
new_counters[k] = (0, 0)
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
@@ -160,7 +173,7 @@ def create_fork_checkpoint(
channels: Mapping[str, BaseChannel],
step: int,
*,
is_fork: bool,
fork_channels: set[str],
get_next_version: GetNextVersion,
) -> Checkpoint:
"""``create_checkpoint`` for the update_state paths that skip the plan.
@@ -170,16 +183,14 @@ def create_fork_checkpoint(
paths never write the delta channel, so its version must be bumped here
or ``put`` drops the blob; derive ``new_versions`` from the result.
"""
if not is_fork:
if not fork_channels:
return create_checkpoint(checkpoint, channels, step)
return create_checkpoint(
checkpoint,
channels,
step,
get_next_version=get_next_version,
channels_to_snapshot=get_delta_channels_from_all_channels(
channels, include_unavailable=True
),
channels_to_snapshot=fork_channels,
)
@@ -204,6 +215,7 @@ def create_checkpoint(
"""
ts = datetime.now(timezone.utc).isoformat()
channels_to_snapshot = channels_to_snapshot or set()
bumped: dict[str, tuple[Any, Any]] = {}
if channels is None:
values = checkpoint["channel_values"]
channel_versions = checkpoint["channel_versions"]
@@ -218,32 +230,22 @@ def create_checkpoint(
# for versioned channels.
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()
)
continue
if k in channels_to_snapshot:
# Callers force a full snapshot blob here: exit mode when a
# delta channel reaches its snapshot cadence, update_state on
# a fresh thread (no ancestor to replay writes from), and a
# fork. The manual version-bump below only applies to the
# exit-mode case.
#
# 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).
# `put` only stores a blob for a channel whose version moved,
# so snapshotting a channel this step did not write needs a
# bump: exit mode reaching the cadence on a superstep that
# skipped the channel, and a fork's first checkpoint.
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)
old = channel_versions[k]
channel_versions[k] = get_next_version(old, None)
bumped[k] = (old, channel_versions[k])
values[k] = _DeltaSnapshot(ch.get())
else:
v = ch.checkpoint()
@@ -255,11 +257,30 @@ def create_checkpoint(
id=id or str(uuid6(clock_seq=step)),
channel_values=values,
channel_versions=channel_versions,
versions_seen=checkpoint["versions_seen"],
versions_seen=_mark_bumps_seen(checkpoint["versions_seen"], bumped),
updated_channels=None if updated_channels is None else sorted(updated_channels),
)
def _mark_bumps_seen(
versions_seen: dict[str, ChannelVersions],
bumped: Mapping[str, tuple[Any, Any]],
) -> dict[str, ChannelVersions]:
"""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.
"""
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}
return out
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.
+9 -12
View File
@@ -102,6 +102,7 @@ from langgraph.pregel._checkpoint import (
copy_checkpoint,
create_checkpoint,
delta_channels_to_snapshot,
delta_channels_with_pending_writes,
empty_checkpoint,
exit_delta_task_id,
)
@@ -226,11 +227,8 @@ class PregelLoop:
# cadence counters say:
# * an Overwrite arrived since the last checkpoint, so sparse replay has to
# start from the post-overwrite value;
# * this run forked off an explicitly addressed checkpoint. That base also
# holds the writes of the branch the fork abandons, and nothing records
# which child consumed which, so the ancestor walk must stop inside the
# fork. Any addressed checkpoint counts, because telling a real fork
# apart would mean trusting the base's `pending_writes` to be complete.
# * the checkpoint this run starts from has pending writes to them; see
# `delta_channels_with_pending_writes`.
_delta_channels_forced_snapshot: set[str]
# The checkpoint_config that points at the parent loaded at `__enter__`
@@ -375,13 +373,6 @@ class PregelLoop:
if self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)
else ()
)
# Value, not key presence like `is_replaying`: subgraph task configs
# always carry an explicit `None` checkpoint_id.
self._delta_channels_forced_snapshot = (
{k for k, spec in specs.items() if isinstance(spec, DeltaChannel)}
if self.checkpoint_config[CONF].get(CONFIG_KEY_CHECKPOINT_ID)
else set()
)
self.prev_checkpoint_config = None
runtime = self.config[CONF].get(CONFIG_KEY_RUNTIME)
self.control = runtime.control if isinstance(runtime, Runtime) else None
@@ -1695,6 +1686,9 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if saved.pending_writes is not None
else []
)
self._delta_channels_forced_snapshot = delta_channels_with_pending_writes(
self.specs, saved.pending_writes
)
self._delta_write_futs = []
self._error_handler_write_futs = []
self._exit_delta_writes = (
@@ -1952,6 +1946,9 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
if saved.pending_writes is not None
else []
)
self._delta_channels_forced_snapshot = delta_channels_with_pending_writes(
self.specs, saved.pending_writes
)
self._delta_write_futs = []
self._error_handler_write_futs = []
self._exit_delta_writes = (
+29 -47
View File
@@ -108,7 +108,6 @@ from langgraph.callbacks import (
get_sync_graph_callback_manager_for_config,
)
from langgraph.channels.base import BaseChannel
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.topic import Topic
from langgraph.config import get_config
from langgraph.constants import END
@@ -135,6 +134,7 @@ from langgraph.pregel._checkpoint import (
create_checkpoint,
create_checkpoint_plan_for_update_state_api,
create_fork_checkpoint,
delta_channels_with_pending_writes,
empty_checkpoint,
get_updated_channels_from_tasks,
)
@@ -1639,25 +1639,21 @@ class Pregel(
else:
raise ValueError(f"Subgraph {recast} not found")
# Read once from the caller's config: every later superstep receives
# the config of the checkpoint just written, which always names one.
# Cleared by the first checkpoint that carries the snapshots, which
# `__copy__` does not write.
fork_pending: set[str] = (
{k for k, v in self.channels.items() if isinstance(v, DeltaChannel)}
if config[CONF].get(CONFIG_KEY_CHECKPOINT_ID)
else set()
)
# Taken from the first superstep's base, and cleared by the first
# checkpoint that carries the snapshots, which `__copy__` does not write.
fork_pending: set[str] | None = None
def perform_superstep(
input_config: RunnableConfig,
updates: Sequence[StateUpdate],
*,
is_fork: bool,
input_config: RunnableConfig, updates: Sequence[StateUpdate]
) -> RunnableConfig:
nonlocal fork_pending
# get last checkpoint
config = ensure_config(self.config, input_config)
saved = checkpointer.get_tuple(config)
if fork_pending is None:
fork_pending = delta_channels_with_pending_writes(
self.channels, saved.pending_writes if saved else None
)
if saved is not None:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
@@ -1745,7 +1741,7 @@ class Pregel(
checkpoint,
channels,
step,
is_fork=is_fork,
fork_channels=fork_pending,
get_next_version=checkpointer.get_next_version,
)
fork_pending.difference_update(next_checkpoint["channel_values"])
@@ -1792,7 +1788,7 @@ class Pregel(
checkpoint,
channels,
next_step,
is_fork=is_fork,
fork_channels=fork_pending,
get_next_version=checkpointer.get_next_version,
)
fork_pending.difference_update(next_checkpoint["channel_values"])
@@ -1904,7 +1900,6 @@ class Pregel(
return perform_superstep(
patch_checkpoint_map(next_config, saved.metadata),
[item for lst in user_group_by.values() for item in lst],
is_fork=is_fork,
)
return patch_checkpoint_map(next_config, saved.metadata)
@@ -2052,21 +2047,19 @@ class Pregel(
parents=saved.metadata.get("parents", {}) if saved else {},
saved_metadata=saved.metadata if saved else None,
is_fresh_thread=saved is None,
is_fork=is_fork,
fork_channels=fork_pending,
)
)
checkpoint = create_checkpoint(
checkpoint,
channels,
step + 1,
updated_channels=updated_channels if channels_to_snapshot else None,
get_next_version=checkpointer.get_next_version
if channels_to_snapshot
else None,
channels_to_snapshot=channels_to_snapshot,
)
if is_fork:
fork_pending.difference_update(checkpoint["channel_values"])
fork_pending.difference_update(checkpoint["channel_values"])
next_config = checkpointer.put(
checkpoint_config,
checkpoint,
@@ -2085,9 +2078,7 @@ class Pregel(
config, {CONFIG_KEY_THREAD_ID: str(config[CONF][CONFIG_KEY_THREAD_ID])}
)
for superstep in supersteps:
current_config = perform_superstep(
current_config, superstep, is_fork=bool(fork_pending)
)
current_config = perform_superstep(current_config, superstep)
return current_config
async def abulk_update_state(
@@ -2140,25 +2131,21 @@ class Pregel(
else:
raise ValueError(f"Subgraph {recast} not found")
# Read once from the caller's config: every later superstep receives
# the config of the checkpoint just written, which always names one.
# Cleared by the first checkpoint that carries the snapshots, which
# `__copy__` does not write.
fork_pending: set[str] = (
{k for k, v in self.channels.items() if isinstance(v, DeltaChannel)}
if config[CONF].get(CONFIG_KEY_CHECKPOINT_ID)
else set()
)
# Taken from the first superstep's base, and cleared by the first
# checkpoint that carries the snapshots, which `__copy__` does not write.
fork_pending: set[str] | None = None
async def aperform_superstep(
input_config: RunnableConfig,
updates: Sequence[StateUpdate],
*,
is_fork: bool,
input_config: RunnableConfig, updates: Sequence[StateUpdate]
) -> RunnableConfig:
nonlocal fork_pending
# get last checkpoint
config = ensure_config(self.config, input_config)
saved = await checkpointer.aget_tuple(config)
if fork_pending is None:
fork_pending = delta_channels_with_pending_writes(
self.channels, saved.pending_writes if saved else None
)
if saved is not None:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
@@ -2244,7 +2231,7 @@ class Pregel(
checkpoint,
channels,
step,
is_fork=is_fork,
fork_channels=fork_pending,
get_next_version=checkpointer.get_next_version,
)
fork_pending.difference_update(next_checkpoint["channel_values"])
@@ -2291,7 +2278,7 @@ class Pregel(
checkpoint,
channels,
next_step,
is_fork=is_fork,
fork_channels=fork_pending,
get_next_version=checkpointer.get_next_version,
)
fork_pending.difference_update(next_checkpoint["channel_values"])
@@ -2402,7 +2389,6 @@ class Pregel(
return await aperform_superstep(
patch_checkpoint_map(next_config, saved.metadata),
[item for lst in user_group_by.values() for item in lst],
is_fork=is_fork,
)
return patch_checkpoint_map(
@@ -2548,21 +2534,19 @@ class Pregel(
parents=saved.metadata.get("parents", {}) if saved else {},
saved_metadata=saved.metadata if saved else None,
is_fresh_thread=saved is None,
is_fork=is_fork,
fork_channels=fork_pending,
)
)
checkpoint = create_checkpoint(
checkpoint,
channels,
step + 1,
updated_channels=updated_channels if channels_to_snapshot else None,
get_next_version=checkpointer.get_next_version
if channels_to_snapshot
else None,
channels_to_snapshot=channels_to_snapshot,
)
if is_fork:
fork_pending.difference_update(checkpoint["channel_values"])
fork_pending.difference_update(checkpoint["channel_values"])
next_config = await checkpointer.aput(
checkpoint_config,
checkpoint,
@@ -2580,9 +2564,7 @@ class Pregel(
config, {CONFIG_KEY_THREAD_ID: str(config[CONF][CONFIG_KEY_THREAD_ID])}
)
for superstep in supersteps:
current_config = await aperform_superstep(
current_config, superstep, is_fork=bool(fork_pending)
)
current_config = await aperform_superstep(current_config, superstep)
return current_config
def update_state(
+148 -2
View File
@@ -16,8 +16,8 @@ from typing_extensions import TypedDict
from langgraph._internal._constants import INPUT
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import END, StateGraph
from langgraph.types import Durability, StateSnapshot, StateUpdate
from langgraph.graph import END, START, StateGraph
from langgraph.types import Durability, StateSnapshot, StateUpdate, interrupt
pytestmark = pytest.mark.anyio
@@ -344,3 +344,149 @@ def test_unaddressed_bulk_update_keeps_snapshot_cadence(
)
assert not _snapshotted_checkpoints(sync_checkpointer, config)
def _both(marker: str) -> dict:
return {"log": [marker], "plain": [marker]}
def _build_paused_before_b(checkpointer: BaseCheckpointSaver) -> Any:
builder = StateGraph(_State)
builder.add_node("a", lambda state: _both("a"))
builder.add_node("b", lambda state: _both("b"))
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("b", END)
return builder.compile(checkpointer=checkpointer, interrupt_before=["b"])
def _build_parallel_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
def ask(state: _State) -> dict:
interrupt("approve?")
return _both("q")
builder = StateGraph(_State)
builder.add_node("p", lambda state: _both("p"))
builder.add_node("q", ask)
builder.add_edge(START, "p")
builder.add_edge(START, "q")
return builder.compile(checkpointer=checkpointer)
def test_resume_at_interrupt_before_with_the_head_checkpoint_id_runs_the_node(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build_paused_before_b(sync_checkpointer)
graph.invoke(_both("in"), config, durability=durability)
graph.invoke(None, graph.get_state(config).config, durability=durability)
state = graph.get_state(config)
assert state.next == (), f"resume paused again before {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
async def test_aresume_at_interrupt_before_with_the_head_checkpoint_id_runs_the_node(
async_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build_paused_before_b(async_checkpointer)
await graph.ainvoke(_both("in"), config, durability=durability)
await graph.ainvoke(
None, (await graph.aget_state(config)).config, durability=durability
)
state = await graph.aget_state(config)
assert state.next == (), f"resume paused again before {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
def test_replay_from_a_paused_checkpoint_runs_the_node_once(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
config = _thread("t")
graph = _build_paused_before_b(sync_checkpointer)
graph.invoke(_both("in"), config)
paused = graph.get_state(config).config
graph.invoke(None, config)
graph.invoke(None, paused)
state = graph.get_state(config)
assert state.next == (), f"replay paused again before {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
@pytest.mark.parametrize("addressed", [False, True])
def test_new_input_on_an_interrupted_head_does_not_replay_its_pending_writes(
sync_checkpointer: BaseCheckpointSaver, durability: Durability, addressed: bool
) -> None:
config = _thread("t")
graph = _build_parallel_interrupt(sync_checkpointer)
graph.invoke(_both("in-1"), config, durability=durability)
head = graph.get_state(config).config
graph.invoke(_both("in-2"), head if addressed else config, durability=durability)
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"] == ["in-1", "in-2", "p"]
def _build_deferred_after_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
builder = StateGraph(_State)
builder.add_node("a", lambda state: _both("a"))
builder.add_node("b", lambda state: _both("b"), defer=True)
builder.add_node("c", lambda state: {})
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("a", "c")
return builder.compile(checkpointer=checkpointer, interrupt_after=["a"])
def test_update_state_with_the_head_checkpoint_id_keeps_a_deferred_node(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_deferred_after_interrupt(sync_checkpointer)
config = _thread("t")
graph.invoke(_both("in"), config)
graph.update_state(graph.get_state(config).config, _both("u"), as_node="c")
graph.invoke(None, config)
state = graph.get_state(config)
assert state.next == (), f"deferred node never ran, still pending: {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "u", "b"]
async def test_aupdate_state_with_the_head_checkpoint_id_keeps_a_deferred_node(
async_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_deferred_after_interrupt(async_checkpointer)
config = _thread("t")
await graph.ainvoke(_both("in"), config)
await graph.aupdate_state(
(await graph.aget_state(config)).config, _both("u"), as_node="c"
)
await graph.ainvoke(None, config)
state = await graph.aget_state(config)
assert state.next == (), f"deferred node never ran, still pending: {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "u", "b"]
def test_turns_addressed_at_the_head_store_no_snapshot(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "turn")
graph.invoke(_input("in-1"), config)
for turn in range(2, 5):
graph.invoke(_input(f"in-{turn}"), graph.get_state(config).config)
assert not _snapshotted_checkpoints(sync_checkpointer, config)
assert (
graph.get_state(config).values["log"] == graph.get_state(config).values["plain"]
)
@@ -338,3 +338,29 @@ def test_state_history_chain_after_fresh_update_state_delta_channel() -> None:
assert update_snapshot.metadata["step"] == 0
assert update_snapshot.parent_config is None
assert [m.content for m in update_snapshot.values["messages"]] == ["hello"]
def test_update_state_that_snapshots_keeps_a_deferred_node_pending() -> None:
channel = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1)
class State(TypedDict):
messages: Annotated[list, channel]
builder = StateGraph(State)
builder.add_node("a", lambda state: {"messages": [HumanMessage("a", id="a")]})
builder.add_node(
"b", lambda state: {"messages": [HumanMessage("b", id="b")]}, defer=True
)
builder.add_node("c", lambda state: {})
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("a", "c")
graph = builder.compile(checkpointer=InMemorySaver(), interrupt_after=["a"])
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage("s", id="s")]}, config)
graph.update_state(config, {"messages": [HumanMessage("u", id="u")]}, as_node="c")
final = graph.invoke(None, config)
assert [m.content for m in final["messages"]] == ["s", "a", "u", "b"]
assert graph.get_state(config).next == ()