diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index c336f75a6..86000ced7 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -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]], diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index e28e7daba..8afafb62f 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -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, diff --git a/libs/langgraph/tests/test_delta_channel_exit_mode.py b/libs/langgraph/tests/test_delta_channel_exit_mode.py index b96fb8bb6..710eff48e 100644 --- a/libs/langgraph/tests/test_delta_channel_exit_mode.py +++ b/libs/langgraph/tests/test_delta_channel_exit_mode.py @@ -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"]