fix(langgraph): keep synthetic task ids for the resumed exit superstep

Storing the resumed superstep's delta writes under their real task ids let
a failed final checkpoint put leave the resumed task looking done: the next
resume skipped it and lost its other writes. They keep their real task
paths, so the order holds, but get the step-prefixed id. A resume addressed
by checkpoint_id reruns tasks whose writes are already stored; the real id
used to dedupe those, so their writes are now skipped instead.
This commit is contained in:
Elior Nataf Lackritz
2026-09-30 15:40:18 -04:00
parent a9a10dedaf
commit b9cafb9f3a
2 changed files with 55 additions and 7 deletions
+20 -7
View File
@@ -223,9 +223,11 @@ class PregelLoop:
_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.
# alive so their ids stay unique; the tasks whose delta writes are among
# them, which a resume addressed by `checkpoint_id` reruns; and the
# checkpoint's own superstep, the first one this run ticks.
_loaded_write_ids: dict[int, tuple[str, str, Any]]
_stored_delta_task_ids: set[str]
_exit_first_step: int | None = None
# Delta channels that saw an Overwrite since the last checkpoint. These
@@ -719,12 +721,16 @@ class PregelLoop:
tid, ch, v = w
if not isinstance(self.specs.get(ch), DeltaChannel):
continue
if id(w) in self._loaded_write_ids:
if (
id(w) in self._loaded_write_ids
or tid in self._stored_delta_task_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 = {}
self._stored_delta_task_ids = set()
# clear pending writes
self.checkpoint_pending_writes.clear()
# only replay (re-execute) done tasks on the first tick
@@ -864,6 +870,11 @@ class PregelLoop:
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}
self._stored_delta_task_ids = {
tid
for tid, ch, _ in self.checkpoint_pending_writes
if tid != NULL_TASK_ID and isinstance(self.specs.get(ch), DeltaChannel)
}
# 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)
@@ -1296,16 +1307,18 @@ class PregelLoop:
# sees the stub as its parent.
self.checkpoint_config = anchor_config
# 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
# The checkpoint's own superstep keeps its real task paths, so it
# interleaves with the writes a resume loaded from it. Its task ids stay
# synthetic: under the real id, a run whose final checkpoint fails to
# save would leave the resumed task looking done to the next resume.
# 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)
key = (exit_delta_task_id(step, tid), path)
else:
key = (exit_delta_late_task_id(step, tid), f"~~{step:010d}{path}")
grouped.setdefault(key, []).append((ch, v))
@@ -501,3 +501,38 @@ def test_exit_resume_replays_supersteps_in_order_on_a_task_id_ordered_saver() ->
graph.invoke(Command(resume="yes"), config, durability="exit")
assert graph.get_state(config).values["log"] == ["in", "ask", "after"]
class _FailingPutSaver(InMemorySaver):
fail = False
def put(
self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any
) -> Any:
if self.fail:
raise RuntimeError("final checkpoint lost")
return super().put(config, checkpoint, metadata, new_versions)
def test_exit_resume_retried_after_its_final_checkpoint_fails_reruns_the_resumed_task() -> (
None
):
saver = _FailingPutSaver()
builder = StateGraph(_ResumeState)
builder.add_node("done", lambda state: _both("done"))
builder.add_node("ask", _ask("ask"))
builder.add_edge(START, "done")
builder.add_edge(START, "ask")
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "t"}}
graph.invoke(_both("in"), config, durability="exit")
saver.fail = True
with pytest.raises(RuntimeError, match="final checkpoint lost"):
graph.invoke(Command(resume="yes"), config, durability="exit")
saver.fail = False
graph.invoke(Command(resume="yes"), config, durability="exit")
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"]
assert sorted(state.values["plain"]) == ["ask", "done", "in"]