diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 7f3141af6..142a420f2 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -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 = ( diff --git a/libs/langgraph/tests/test_delta_channel_fork.py b/libs/langgraph/tests/test_delta_channel_fork.py index 6390e35fd..12d01f8c0 100644 --- a/libs/langgraph/tests/test_delta_channel_fork.py +++ b/libs/langgraph/tests/test_delta_channel_fork.py @@ -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: