Compare commits

...
Author SHA1 Message Date
Elior Nataf Lackritz a9a10dedaf 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.
2026-09-29 12:44:14 -04:00
Elior Nataf Lackritz 89ff2d33de 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.
2026-09-29 12:41:16 -04:00
3 changed files with 166 additions and 21 deletions
@@ -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]],
+44 -21
View File
@@ -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,10 +218,15 @@ 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.
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
# 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, 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
# channels must snapshot after live update applies overwrite semantics so
@@ -707,9 +713,18 @@ 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 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 = {}
# clear pending writes
self.checkpoint_pending_writes.clear()
# only replay (re-execute) done tasks on the first tick
@@ -848,6 +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): 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 +1033,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 +1261,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 +1296,19 @@ 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))
# 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 tid == NULL_TASK_ID:
key = (exit_delta_task_id(step, tid), "")
elif step == self._exit_first_step:
key = (tid, path)
else:
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,
{
@@ -1294,22 +1318,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)
@@ -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,112 @@ 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"]
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"]