diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 8afafb62f..c1c423ec0 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -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)) diff --git a/libs/langgraph/tests/test_delta_channel_exit_mode.py b/libs/langgraph/tests/test_delta_channel_exit_mode.py index 710eff48e..bcd22d09c 100644 --- a/libs/langgraph/tests/test_delta_channel_exit_mode.py +++ b/libs/langgraph/tests/test_delta_channel_exit_mode.py @@ -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"]