From a9a10dedaf9ccc9749ab5bc40cb9cb41fc9695e3 Mon Sep 17 00:00:00 2001 From: Elior Nataf Lackritz Date: Tue, 29 Sep 2026 12:44:14 -0400 Subject: [PATCH] fix(langgraph): order later exit supersteps after real task ids too Later supersteps now get task ids that start with ffffffff, so they sort after every real task id as well as after every real task path. Savers that replay by task id, including released checkpoint packages that do not order by path, would otherwise see a multi-step exit run's later writes before its first superstep, fresh runs included. Loaded writes are skipped for every saver, since they are stored on the checkpoint either way, and kept alive so their ids stay unique for the tick. --- .../langgraph/langgraph/pregel/_checkpoint.py | 10 +++++ libs/langgraph/langgraph/pregel/_loop.py | 38 +++++++---------- .../tests/test_delta_channel_exit_mode.py | 42 +++++++++++++++++++ 3 files changed, 68 insertions(+), 22 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index c336f75a6..86000ced7 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -47,6 +47,16 @@ def exit_delta_task_id(step: int, task_id: str) -> str: return f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}" +def exit_delta_late_task_id(step: int, task_id: str) -> str: + """Synthetic task id for exit-mode writes of a superstep after the anchor's own. + + Sorts after every real task id, in step order, so replay keeps them after + the anchor's own superstep whether a saver orders by task path or task id. + """ + parts = str(uuid.UUID(task_id)).split("-") + return f"ffffffff-{step >> 16:04x}-{step & 0xFFFF:04x}-{parts[3]}-{parts[4]}" + + def delta_channels_to_snapshot( channels: Mapping[str, BaseChannel], counters_since_delta_snapshot: Mapping[str, tuple[int, int]], diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index e28e7daba..8afafb62f 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -103,6 +103,7 @@ from langgraph.pregel._checkpoint import ( create_checkpoint, delta_channels_to_snapshot, empty_checkpoint, + exit_delta_late_task_id, exit_delta_task_id, ) from langgraph.pregel._executor import ( @@ -217,14 +218,14 @@ class PregelLoop: # `after_tick`). At exit, `_put_exit_delta_writes` filters out channels # that will snapshot, then persists the rest under an anchor parent. # `None` when not in exit mode (so the capture sites are no-ops). - # 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. + # Each tuple is `(step, task_id, task_path, channel, value)`; see + # `_put_exit_delta_writes` for how they are ordered. _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] + # The pending writes loaded with the checkpoint, already stored on it, kept + # alive so their ids stay unique; and the checkpoint's own superstep, the + # first one this run ticks. + _loaded_write_ids: dict[int, tuple[str, str, Any]] _exit_first_step: int | None = None # Delta channels that saw an Overwrite since the last checkpoint. These @@ -718,15 +719,12 @@ class PregelLoop: 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 - ): + if 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() + self._loaded_write_ids = {} # clear pending writes self.checkpoint_pending_writes.clear() # only replay (re-execute) done tasks on the first tick @@ -865,7 +863,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} + self._loaded_write_ids = {id(w): 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) @@ -1298,22 +1296,18 @@ class PregelLoop: # sees the stub as its parent. self.checkpoint_config = anchor_config - # 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. + # 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 and task id, in step + # order, so this holds whether a saver orders by path or by id. 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: + if 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}") + key = (exit_delta_late_task_id(step, tid), f"~~{step:010d}{path}") grouped.setdefault(key, []).append((ch, v)) anchor_write_config = patch_configurable( anchor_config, diff --git a/libs/langgraph/tests/test_delta_channel_exit_mode.py b/libs/langgraph/tests/test_delta_channel_exit_mode.py index b96fb8bb6..710eff48e 100644 --- a/libs/langgraph/tests/test_delta_channel_exit_mode.py +++ b/libs/langgraph/tests/test_delta_channel_exit_mode.py @@ -459,3 +459,45 @@ def test_resume_interleaves_the_resumed_superstep_by_task_path( state = graph.get_state(config) assert state.values["log"] == state.values["plain"] == ["in", "a", "z"] + + +class _TaskIdOrderSaver(InMemorySaver): + """Replays each checkpoint's writes by task id, as savers without task path + ordering do.""" + + def get_tuple(self, config: Any) -> Any: + tup = super().get_tuple(config) + if tup and tup.pending_writes: + tup = tup._replace(pending_writes=sorted(tup.pending_writes)) + return tup + + get_delta_channel_history = BaseCheckpointSaver.get_delta_channel_history + + +def test_exit_run_replays_supersteps_in_order_on_a_task_id_ordered_saver() -> None: + builder = StateGraph(_ResumeState) + 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") + graph = builder.compile(checkpointer=_TaskIdOrderSaver()) + config = {"configurable": {"thread_id": "t"}} + + graph.invoke(_both("in"), config, durability="exit") + + assert graph.get_state(config).values["log"] == ["in", "a", "b"] + + +def test_exit_resume_replays_supersteps_in_order_on_a_task_id_ordered_saver() -> None: + builder = StateGraph(_ResumeState) + builder.add_node("ask", _ask("ask")) + builder.add_node("after", lambda state: _both("after")) + builder.add_edge(START, "ask") + builder.add_edge("ask", "after") + graph = builder.compile(checkpointer=_TaskIdOrderSaver()) + config = {"configurable": {"thread_id": "t"}} + graph.invoke(_both("in"), config, durability="exit") + + graph.invoke(Command(resume="yes"), config, durability="exit") + + assert graph.get_state(config).values["log"] == ["in", "ask", "after"]