From 89ff2d33de39a138485dee575adf1ad8d9384443 Mon Sep 17 00:00:00 2001 From: Elior Nataf Lackritz Date: Tue, 29 Sep 2026 11:45:45 -0400 Subject: [PATCH] fix(langgraph): replay a resumed exit-mode run's delta writes in live order Exit durability stores a run's delta writes on the checkpoint it started from, under step-prefixed task ids. When that checkpoint already held writes (a resume after a parallel interrupt), the accumulator stored the loaded writes a second time, so every later read replayed them twice. Skipping them alone is not enough: the loaded writes keep their real task ids, which sort after the step-prefixed ones. Store the checkpoint's own superstep as sync durability does (real task id and path), so it interleaves with the loaded writes by task path, and give later supersteps a task path that sorts after every real one, in step order. Savers that take no task path keep the previous encoding. --- libs/langgraph/langgraph/pregel/_loop.py | 65 ++++++++++++----- .../tests/test_delta_channel_exit_mode.py | 70 +++++++++++++++++++ 2 files changed, 117 insertions(+), 18 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 371037768..e28e7daba 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -220,7 +220,12 @@ class PregelLoop: # Each tuple is `(step, task_id, channel, value)` — `step` drives the # synthetic step-prefixed task_id used to preserve chronological order # under the saver's `ORDER BY task_id, idx` sorting. - _exit_delta_writes: list[tuple[int, str, str, Any]] | None = None + _exit_delta_writes: list[tuple[int, str, str, str, Any]] | None = None + + # The pending writes loaded with the checkpoint are already stored on it, + # and the first superstep that captured writes is the checkpoint's own. + _loaded_write_ids: set[int] + _exit_first_step: int | None = None # Delta channels that saw an Overwrite since the last checkpoint. These # channels must snapshot after live update applies overwrite semantics so @@ -707,9 +712,21 @@ 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): - self._exit_delta_writes.append((self.step, tid, ch, v)) + if self._exit_first_step is None: + self._exit_first_step = self.step + for w in self.checkpoint_pending_writes: + tid, ch, v = w + if not isinstance(self.specs.get(ch), DeltaChannel): + continue + if ( + self.checkpointer_put_writes_accepts_task_path + and id(w) in self._loaded_write_ids + ): + continue + task = self.tasks.get(tid) + path = task_path_str(task.path) if task else "" + self._exit_delta_writes.append((self.step, tid, path, 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 @@ -848,6 +865,7 @@ class PregelLoop: def _first( self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None ) -> set[str] | None: + self._loaded_write_ids = {id(w) for w in self.checkpoint_pending_writes} # Resuming from a previous checkpoint requires two things: # 1. A prior checkpoint exists (channel_versions is non-empty) # 2. The input signals continuation (not a fresh run with new input) @@ -1017,7 +1035,9 @@ class PregelLoop: if self._exit_delta_writes is not None: for c, v in input_writes: if isinstance(self.specs.get(c), DeltaChannel): - self._exit_delta_writes.append((self.step, NULL_TASK_ID, c, v)) + self._exit_delta_writes.append( + (self.step, NULL_TASK_ID, "", c, v) + ) # Persist delta-channel input writes so sub-freq inputs are # recoverable via ancestor walk (mirrors the Command input path). if self.durability != "exit": @@ -1243,9 +1263,7 @@ class PregelLoop: ) pending = [ - (step, tid, ch, v) - for (step, tid, ch, v) in self._exit_delta_writes - if ch not in channels_to_snapshot + w for w in self._exit_delta_writes if w[3] not in channels_to_snapshot ] if not pending: return @@ -1280,11 +1298,23 @@ class PregelLoop: # sees the stub as its parent. self.checkpoint_config = anchor_config - # Step-prefixed synthetic task_id preserves chronological superstep - # order under the saver's ORDER BY task_id, idx sorting. - grouped: dict[tuple[int, str], list[tuple[str, Any]]] = {} - for step, tid, ch, v in pending: - grouped.setdefault((step, tid), []).append((ch, v)) + # Replay orders a checkpoint's writes by task path, then task id. The + # checkpoint's own superstep is stored as sync durability stores it, so + # it interleaves with the writes a resume loaded from it; later + # supersteps sort after every real task path, in step order. A saver + # that takes no task path orders by task id alone, so it gets the + # step-prefixed id for every write. + grouped: dict[tuple[str, str], list[tuple[str, Any]]] = {} + for step, tid, path, ch, v in pending: + if not self.checkpointer_put_writes_accepts_task_path: + key = (exit_delta_task_id(step, tid), "") + elif tid == NULL_TASK_ID: + key = (exit_delta_task_id(step, tid), "") + elif step == self._exit_first_step: + key = (tid, path) + else: + key = (exit_delta_task_id(step, tid), f"~~{step:010d}{path}") + grouped.setdefault(key, []).append((ch, v)) anchor_write_config = patch_configurable( anchor_config, { @@ -1294,22 +1324,21 @@ class PregelLoop: CONFIG_KEY_CHECKPOINT_ID: anchor_config[CONF][CONFIG_KEY_CHECKPOINT_ID], }, ) - for (step, tid), entries in grouped.items(): - synth_tid = exit_delta_task_id(step, tid) + for (tid, path), entries in grouped.items(): if self.checkpointer_put_writes_accepts_task_path: fut = self.submit( self.checkpointer_put_writes, anchor_write_config, entries, - synth_tid, - "", + tid, + path, ) else: fut = self.submit( self.checkpointer_put_writes, anchor_write_config, entries, - synth_tid, + tid, ) if self._delta_write_futs is not None: self._delta_write_futs.append(fut) diff --git a/libs/langgraph/tests/test_delta_channel_exit_mode.py b/libs/langgraph/tests/test_delta_channel_exit_mode.py index 3d18ce708..b96fb8bb6 100644 --- a/libs/langgraph/tests/test_delta_channel_exit_mode.py +++ b/libs/langgraph/tests/test_delta_channel_exit_mode.py @@ -6,11 +6,13 @@ channel), lazy stub creation when no parent exists, and proper read-path reconstruction via ancestor walks. """ +import operator import uuid from typing import Annotated, Any import pytest from langchain_core.messages import AIMessage, HumanMessage +from langgraph.checkpoint.base import BaseCheckpointSaver from langgraph.checkpoint.memory import InMemorySaver from langgraph.checkpoint.serde.types import _DeltaSnapshot from typing_extensions import TypedDict @@ -19,6 +21,7 @@ from langgraph.channels.delta import DeltaChannel from langgraph.graph import START, StateGraph from langgraph.graph.message import _messages_delta_reducer from langgraph.pregel._checkpoint import exit_delta_task_id +from langgraph.types import Command, Durability, interrupt pytestmark = pytest.mark.anyio @@ -389,3 +392,70 @@ async def test_exit_snapshot_then_tail_deltas() -> None: assert "seed-msg" in contents assert "tail-msg" in contents assert contents.index("seed-msg") < contents.index("tail-msg") + + +def _append(current: list, writes: list) -> list: + out = list(current) + for write in writes: + out.extend(write) + return out + + +class _ResumeState(TypedDict): + log: Annotated[list, DeltaChannel(_append)] + plain: Annotated[list, operator.add] + + +def _both(marker: str) -> dict: + return {"log": [marker], "plain": [marker]} + + +def _ask(marker: str) -> Any: + def ask(state: _ResumeState) -> dict: + interrupt("approve?") + return _both(marker) + + return ask + + +@pytest.mark.parametrize("addressed", [False, True]) +def test_resume_after_a_parallel_interrupt_replays_in_live_order( + sync_checkpointer: BaseCheckpointSaver, durability: Durability, addressed: bool +) -> None: + builder = StateGraph(_ResumeState) + builder.add_node("done", lambda state: _both("done")) + builder.add_node("ask", _ask("ask")) + builder.add_node("after", lambda state: _both("after")) + builder.add_edge(START, "done") + builder.add_edge(START, "ask") + builder.add_edge("ask", "after") + graph = builder.compile(checkpointer=sync_checkpointer) + config = {"configurable": {"thread_id": "t"}} + graph.invoke(_both("in"), config, durability=durability) + head = graph.get_state(config).config + + graph.invoke( + Command(resume="yes"), head if addressed else config, durability=durability + ) + + state = graph.get_state(config) + assert state.values["log"] == state.values["plain"] + assert sorted(state.values["log"]) == ["after", "ask", "done", "in"] + + +def test_resume_interleaves_the_resumed_superstep_by_task_path( + sync_checkpointer: BaseCheckpointSaver, durability: Durability +) -> None: + builder = StateGraph(_ResumeState) + builder.add_node("z_done", lambda state: _both("z")) + builder.add_node("a_asks", _ask("a")) + builder.add_edge(START, "z_done") + builder.add_edge(START, "a_asks") + graph = builder.compile(checkpointer=sync_checkpointer) + config = {"configurable": {"thread_id": "t"}} + graph.invoke(_both("in"), config, durability=durability) + + graph.invoke(Command(resume="yes"), config, durability=durability) + + state = graph.get_state(config) + assert state.values["log"] == state.values["plain"] == ["in", "a", "z"]