mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
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:
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user