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.
This commit is contained in:
Elior Nataf Lackritz
2026-09-29 12:44:14 -04:00
parent 89ff2d33de
commit a9a10dedaf
3 changed files with 68 additions and 22 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]],
+16 -22
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,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,
@@ -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"]