mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 05:25:05 +02:00
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:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user