fix(langgraph): skip the seal on a resume that reuses the head's writes

A resume that is not replaying reuses the head's pending writes instead of
rerunning their tasks, so none of them can leak; deciding in `_first`,
where that is known, stops a plain `Command(resume=...)` from storing a
snapshot. A replaying resume reruns the tasks, so it still seals.

Exit mode also re-recorded the writes a resume loaded from the head: they
were already stored there, and every later read replayed them twice. The
exit accumulator now skips them.
This commit is contained in:
Elior Nataf Lackritz
2026-09-29 09:42:55 -04:00
parent ddaf708cd0
commit 719a4d71bc
2 changed files with 54 additions and 10 deletions
+21 -8
View File
@@ -223,6 +223,10 @@ class PregelLoop:
# under the saver's `ORDER BY task_id, idx` sorting.
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
# ids of the pending writes loaded with the checkpoint. They are already
# stored on it, so the exit accumulator must not store them again.
_loaded_write_ids: set[int]
# Delta channels that must snapshot at the next checkpoint, whatever their
# cadence counters say:
# * an Overwrite arrived since the last checkpoint, so sparse replay has to
@@ -711,9 +715,14 @@ class PregelLoop:
)
# 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):
for w in self.checkpoint_pending_writes:
tid, ch, v = w
if (
isinstance(self.specs.get(ch), DeltaChannel)
and id(w) not in self._loaded_write_ids
):
self._exit_delta_writes.append((self.step, tid, ch, v))
self._loaded_write_ids = set()
# clear pending writes
self.checkpoint_pending_writes.clear()
# only replay (re-execute) done tasks on the first tick
@@ -860,6 +869,7 @@ class PregelLoop:
# - None input: resume after interrupt (invoke(None, config))
# - Command input: any Command operates on existing state
# - Same run_id: re-entry into an ongoing run (e.g. stream reconnect)
self._loaded_write_ids = {id(w) for w in self.checkpoint_pending_writes}
configurable = self.config.get(CONF, {})
input_is_command = isinstance(self.input, Command)
is_resuming = bool(self.checkpoint["channel_versions"]) and bool(
@@ -902,6 +912,15 @@ 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.
self._delta_channels_forced_snapshot = (
set()
if is_resuming and not self.is_replaying
else delta_channels_with_pending_writes(
self.specs, self.checkpoint_pending_writes
)
)
# map command to writes
if input_is_command:
@@ -1686,9 +1705,6 @@ 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 = (
@@ -1946,9 +1962,6 @@ 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 = (
@@ -17,7 +17,7 @@ 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 Durability, StateSnapshot, StateUpdate, interrupt
from langgraph.types import Command, Durability, StateSnapshot, StateUpdate, interrupt
pytestmark = pytest.mark.anyio
@@ -363,7 +363,7 @@ def _build_paused_before_b(checkpointer: BaseCheckpointSaver) -> Any:
def _build_parallel_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
def ask(state: _State) -> dict:
interrupt("approve?")
return _both("q")
return {"other": ["q"]}
builder = StateGraph(_State)
builder.add_node("p", lambda state: _both("p"))
@@ -445,6 +445,37 @@ def _build_deferred_after_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
return builder.compile(checkpointer=checkpointer, interrupt_after=["a"])
def test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build_parallel_interrupt(sync_checkpointer)
graph.invoke(_both("in-1"), config, durability=durability)
graph.invoke(Command(resume="yes"), config, durability=durability)
state = graph.get_state(config)
assert state.next == ()
assert state.values["log"] == state.values["plain"] == ["in-1", "p"]
assert not _snapshotted_checkpoints(sync_checkpointer, config)
def test_resume_addressed_at_an_interrupted_head_reruns_its_tasks_once(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build_parallel_interrupt(sync_checkpointer)
graph.invoke(_both("in-1"), config, durability=durability)
graph.invoke(
Command(resume="yes"), graph.get_state(config).config, durability=durability
)
state = graph.get_state(config)
assert state.next == ()
assert state.values["log"] == state.values["plain"] == ["in-1", "p"]
def test_update_state_with_the_head_checkpoint_id_keeps_a_deferred_node(
sync_checkpointer: BaseCheckpointSaver,
) -> None: