revert(langgraph): leave exit-mode resume writes as they are on main

Skipping the loaded writes in the exit accumulator removed the duplicate
replay but reordered it: the loaded write keeps its task id while the
run's later writes get step-prefixed ids that sort first. That ordering
problem belongs to exit mode, not to forking, and exists on main without
a fork, so it gets its own fix. The resume test marks exit durability as
an expected failure until then.
This commit is contained in:
Elior Nataf Lackritz
2026-09-29 11:39:36 -04:00
parent 719a4d71bc
commit e59ecbc23d
2 changed files with 16 additions and 12 deletions
+2 -12
View File
@@ -223,10 +223,6 @@ class PregelLoop:
# under the saver's `ORDER BY task_id, idx` sorting.
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
# ids of the pending writes loaded with the checkpoint. They are already
# stored on it, so the exit accumulator must not store them again.
_loaded_write_ids: set[int]
# Delta channels that must snapshot at the next checkpoint, whatever their
# cadence counters say:
# * an Overwrite arrived since the last checkpoint, so sparse replay has to
@@ -715,14 +711,9 @@ class PregelLoop:
)
# capture delta-channel writes for exit-mode accumulator before clearing
if self._exit_delta_writes is not None:
for w in self.checkpoint_pending_writes:
tid, ch, v = w
if (
isinstance(self.specs.get(ch), DeltaChannel)
and id(w) not in self._loaded_write_ids
):
for tid, ch, v in self.checkpoint_pending_writes:
if isinstance(self.specs.get(ch), DeltaChannel):
self._exit_delta_writes.append((self.step, tid, ch, v))
self._loaded_write_ids = set()
# clear pending writes
self.checkpoint_pending_writes.clear()
# only replay (re-execute) done tasks on the first tick
@@ -869,7 +860,6 @@ class PregelLoop:
# - None input: resume after interrupt (invoke(None, config))
# - Command input: any Command operates on existing state
# - Same run_id: re-entry into an ongoing run (e.g. stream reconnect)
self._loaded_write_ids = {id(w) for w in self.checkpoint_pending_writes}
configurable = self.config.get(CONF, {})
input_is_command = isinstance(self.input, Command)
is_resuming = bool(self.checkpoint["channel_versions"]) and bool(
@@ -445,6 +445,20 @@ def _build_deferred_after_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
return builder.compile(checkpointer=checkpointer, interrupt_after=["a"])
@pytest.mark.parametrize(
"durability",
[
"sync",
"async",
pytest.param(
"exit",
marks=pytest.mark.xfail(
reason="exit durability stores a resumed run's loaded writes twice",
strict=True,
),
),
],
)
def test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None: