fix: replay behavior for parent + subgraphs! (#7038)

## Summary

Fix time travel (replay and fork) for graphs with interrupts and
subgraphs.

## Problem

Two issues with replaying/forking from earlier checkpoints:

1. **Stale interrupt values during replay** — Replays incorrectly reused
cached `RESUME` values from prior `interrupt()` calls, so interrupts
silently returned stale answers instead of re-firing.

2. **Wrong subgraph state during time travel** — Subgraphs always loaded
their **latest** checkpoint instead of the one corresponding to the
parent's historical state. This caused subgraphs to skip execution or
produce incorrect results during replay/fork.

## Changes

Code changes span `libs/langgraph/langgraph/pregel/_loop.py`,
`libs/langgraph/langgraph/_internal/_constants.py`, and a new
`libs/langgraph/langgraph/_internal/_replay.py` module:

- **Strip stale `RESUME` writes on replay** — During replays, cached
`RESUME` writes are filtered out so `interrupt()` re-fires instead of
returning old values. Genuine resumes (`Command(resume=...)`) preserve
these writes.

- **Rename `skip_done_tasks` → `is_replaying`** — Clearer naming for the
flag that tracks whether the current run is replaying from a specific
checkpoint.

- **New `ReplayState` class (`_replay.py`)** — Encapsulates subgraph
checkpoint loading during time-travel. Tracks a parent checkpoint ID
upper bound and which subgraph namespaces have already loaded their
pre-replay checkpoint. On the first visit to a subgraph namespace, it
loads the latest checkpoint created *before* the replay point (via
`checkpointer.list(..., before=...)` with `limit=1`). On subsequent
visits (e.g. the same subgraph in a later loop iteration), it falls back
to normal latest-checkpoint loading. The task-id suffix is stripped from
namespaces so the same logical subgraph is recognized across loop
iterations.

- **New `CONFIG_KEY_REPLAY_STATE` config key** — The parent graph
creates a `ReplayState` instance and passes it to subgraphs via config.
For forks (`source=update`), the replay state uses the fork's parent
checkpoint ID since the fork was created after the subgraph's original
checkpoints. The single `ReplayState` instance is shared by reference
across all derived configs within one parent execution.

- **Subgraph checkpoint loading in `__enter__`/`__aenter__`** — When a
subgraph detects a `ReplayState` in its config, it delegates checkpoint
loading to `ReplayState.get_checkpoint`/`aget_checkpoint` instead of
using the default `get_tuple`. It also clears `CONFIG_KEY_RESUMING` so
`_first` re-applies input and recreates ephemeral routing channels.

## Tests

New test files `test_time_travel.py` (~2500 lines) and
`test_time_travel_async.py` (~2200 lines) covering:
- Replay and fork with interrupts (single and multiple)
- Replay and fork for graphs with and without subgraphs
- Correct subgraph checkpoint restoration during parent time travel
- `get_state` with subgraph state during replay
This commit is contained in:
Sydney Runkle
2026-03-09 21:21:26 -04:00
committed by GitHub
parent 2638ff715a
commit 9c2deacb28
5 changed files with 5019 additions and 23 deletions
@@ -41,6 +41,9 @@ CONFIG_KEY_CACHE = sys.intern("__pregel_cache")
# holds a `BaseCache` made available to subgraphs
CONFIG_KEY_RESUMING = sys.intern("__pregel_resuming")
# holds a boolean indicating if subgraphs should resume from a previous checkpoint
CONFIG_KEY_REPLAY_STATE = sys.intern("__pregel_replay_state")
# holds a ReplayState tracking the parent checkpoint_id upper bound and which
# subgraph namespaces have already loaded their pre-replay checkpoint
CONFIG_KEY_TASK_ID = sys.intern("__pregel_task_id")
# holds the task ID for the current task
CONFIG_KEY_THREAD_ID = sys.intern("thread_id")
@@ -98,6 +101,7 @@ RESERVED = {
CONFIG_KEY_STREAM,
CONFIG_KEY_CHECKPOINT_MAP,
CONFIG_KEY_RESUMING,
CONFIG_KEY_REPLAY_STATE,
CONFIG_KEY_TASK_ID,
CONFIG_KEY_CHECKPOINT_MAP,
CONFIG_KEY_CHECKPOINT_ID,
@@ -0,0 +1,90 @@
"""Replay state for subgraph checkpoint loading during time-travel."""
from __future__ import annotations
from typing import TYPE_CHECKING
from langgraph._internal._constants import NS_END
if TYPE_CHECKING:
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointTuple
class ReplayState:
"""Tracks which subgraphs have already loaded their pre-replay checkpoint.
During a parent replay, each subgraph's first invocation should restore the
checkpoint from before the replay point. Subsequent invocations of the same
subgraph (e.g. in a loop) should use normal checkpoint loading so they pick
up freshly created checkpoints.
The single `ReplayState` instance is shared by reference across all derived
configs within one parent execution.
"""
__slots__ = ("checkpoint_id", "_visited_ns")
def __init__(self, checkpoint_id: str) -> None:
self.checkpoint_id = checkpoint_id
# DO NOT CHANGE THIS VARIABLE it may need to be rehydrated
# in other runtimes
self._visited_ns: set[str] = set()
def _is_first_visit(self, checkpoint_ns: str) -> bool:
"""Return True the first time a subgraph namespace is seen.
The task-id suffix is stripped so that the same logical subgraph
(e.g. ``"sub_node"``) is recognized across loop iterations even
though each iteration has a different task id.
"""
# "sub_node:task_id" -> "sub_node"
stable_ns = (
checkpoint_ns.rsplit(NS_END, 1)[0]
if NS_END in checkpoint_ns
else checkpoint_ns
)
if stable_ns in self._visited_ns:
return False
self._visited_ns.add(stable_ns)
return True
def get_checkpoint(
self,
checkpoint_ns: str,
checkpointer: BaseCheckpointSaver,
checkpoint_config: RunnableConfig,
) -> CheckpointTuple | None:
"""Load the right checkpoint for a subgraph during replay.
On the first call for a given subgraph namespace, returns the latest
checkpoint created *before* the replay point. On subsequent calls
(e.g. the same subgraph in a later loop iteration), falls back to
normal latest-checkpoint loading.
"""
if self._is_first_visit(checkpoint_ns):
for saved in checkpointer.list(
checkpoint_config,
before={"configurable": {"checkpoint_id": self.checkpoint_id}},
limit=1,
):
return saved
return None
return checkpointer.get_tuple(checkpoint_config)
async def aget_checkpoint(
self,
checkpoint_ns: str,
checkpointer: BaseCheckpointSaver,
checkpoint_config: RunnableConfig,
) -> CheckpointTuple | None:
"""Async version of `get_checkpoint`."""
if self._is_first_visit(checkpoint_ns):
async for saved in checkpointer.alist(
checkpoint_config,
before={"configurable": {"checkpoint_id": self.checkpoint_id}},
limit=1,
):
return saved
return None
return await checkpointer.aget_tuple(checkpoint_config)
+99 -23
View File
@@ -42,6 +42,7 @@ from langgraph._internal._constants import (
CONFIG_KEY_CHECKPOINT_ID,
CONFIG_KEY_CHECKPOINT_MAP,
CONFIG_KEY_CHECKPOINT_NS,
CONFIG_KEY_REPLAY_STATE,
CONFIG_KEY_RESUME_MAP,
CONFIG_KEY_RESUMING,
CONFIG_KEY_SCRATCHPAD,
@@ -58,6 +59,7 @@ from langgraph._internal._constants import (
RESUME,
TASKS,
)
from langgraph._internal._replay import ReplayState
from langgraph._internal._scratchpad import PregelScratchpad
from langgraph._internal._typing import EMPTY_SEQ, MISSING
from langgraph.channels.base import BaseChannel
@@ -152,7 +154,7 @@ class PregelLoop:
input_keys: str | Sequence[str]
output_keys: str | Sequence[str]
stream_keys: str | Sequence[str]
skip_done_tasks: bool
is_replaying: bool
is_nested: bool
manager: None | AsyncParentRunManager | ParentRunManager
interrupt_after: All | Sequence[str]
@@ -244,7 +246,7 @@ class PregelLoop:
self.interrupt_before = interrupt_before
self.manager = manager
self.is_nested = CONFIG_KEY_TASK_ID in self.config.get(CONF, {})
self.skip_done_tasks = CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
self.is_replaying = CONFIG_KEY_CHECKPOINT_ID in config[CONF]
self._migrate_checkpoint = migrate_checkpoint
self.trigger_to_nodes = trigger_to_nodes
self.retry_policy = retry_policy
@@ -451,7 +453,7 @@ class PregelLoop:
# save the new task
self.tasks[pushed.id] = pushed
# match any pending writes to the new task
if self.skip_done_tasks:
if not self.is_replaying:
self._match_writes({pushed.id: pushed})
# return the new task, to be started if not run before
return pushed
@@ -515,7 +517,7 @@ class PregelLoop:
return False
# if there are pending writes from a previous loop, apply them
if self.skip_done_tasks and self.checkpoint_pending_writes:
if not self.is_replaying and self.checkpoint_pending_writes:
self._match_writes(self.tasks)
# before execution, check if we should interrupt
@@ -557,8 +559,8 @@ class PregelLoop:
)
# clear pending writes
self.checkpoint_pending_writes.clear()
# "not skip_done_tasks" only applies to first tick after resuming
self.skip_done_tasks = True
# only replay (re-execute) done tasks on the first tick
self.is_replaying = False
# save checkpoint
self._put_checkpoint({"source": "loop"})
# after execution, check if we should interrupt
@@ -618,15 +620,21 @@ class PregelLoop:
def _first(
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
) -> set[str] | None:
# resuming from previous checkpoint requires
# - finding a previous checkpoint
# - receiving None input (outer graph) or RESUMING flag (subgraph)
# 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)
# For subgraphs, the parent explicitly sets CONFIG_KEY_RESUMING.
# For the outer graph, we infer from the input:
# - 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)
configurable = self.config.get(CONF, {})
input_is_command = isinstance(self.input, Command)
is_resuming = bool(self.checkpoint["channel_versions"]) and bool(
configurable.get(
CONFIG_KEY_RESUMING,
self.input is None
or isinstance(self.input, Command)
or input_is_command
or (
not self.is_nested
and self.config.get("metadata", {}).get("run_id")
@@ -635,9 +643,25 @@ class PregelLoop:
)
)
# When replaying from a specific checkpoint, drop cached RESUME
# writes so that interrupt() calls re-fire instead of returning
# stale values. But if we're actively resuming, keep them —
# multi-interrupt scenarios need previously resolved values preserved.
# We check two conditions because resume signals arrive differently:
# - Command(resume=...): the outer graph receives resume via input
# - CONFIG_KEY_RESUMING: child subgraphs receive it via config from
# the parent (their input is a Send arg, not a Command)
if self.is_replaying and not (
(input_is_command and cast(Command, self.input).resume is not None)
or configurable.get(CONFIG_KEY_RESUMING, False)
):
self.checkpoint_pending_writes = [
w for w in self.checkpoint_pending_writes if w[1] != RESUME
]
# map command to writes
if isinstance(self.input, Command):
if (resume := self.input.resume) is not None:
if input_is_command:
if (resume := cast(Command, self.input).resume) is not None:
if not self.checkpointer:
raise RuntimeError(
"Cannot use Command(resume=...) without checkpointer"
@@ -657,7 +681,7 @@ class PregelLoop:
writes: defaultdict[str, list[tuple[str, Any]]] = defaultdict(list)
# group writes by task ID
for tid, c, v in map_command(cmd=self.input):
for tid, c, v in map_command(cmd=cast(Command, self.input)):
if not (c == RESUME and resume_is_map):
writes[tid].append((c, v))
if not writes and not resume_is_map:
@@ -723,10 +747,30 @@ class PregelLoop:
self._put_checkpoint({"source": "input"})
elif CONFIG_KEY_RESUMING not in configurable:
raise EmptyInputError(f"Received no input for {input_keys}")
# update config
# Propagate resuming and replaying flags to subgraphs.
if not self.is_nested:
# Pass the resolved before-bound checkpoint ID so subgraphs can
# find their corresponding checkpoint without re-fetching the
# parent. For forks (source=update), use the fork's parent
# checkpoint ID since the fork was created after the subgraph's
# checkpoints from the original execution.
replay_state: ReplayState | None = None
if self.is_replaying:
replay_checkpoint_id = self.checkpoint["id"]
if (
self.checkpoint_metadata.get("source") == "update"
and self.prev_checkpoint_config
):
replay_checkpoint_id = self.prev_checkpoint_config[CONF].get(
CONFIG_KEY_CHECKPOINT_ID, replay_checkpoint_id
)
replay_state = ReplayState(replay_checkpoint_id)
self.config = patch_configurable(
self.config, {CONFIG_KEY_RESUMING: is_resuming}
self.config,
{
CONFIG_KEY_RESUMING: is_resuming,
CONFIG_KEY_REPLAY_STATE: replay_state,
},
)
# set flag
self.status = "pending"
@@ -1081,10 +1125,27 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
# context manager
def __enter__(self) -> Self:
if self.checkpointer:
saved = self.checkpointer.get_tuple(self.checkpoint_config)
else:
if not self.checkpointer:
saved = None
elif self.is_nested and (
replay_state := self.config[CONF].get(CONFIG_KEY_REPLAY_STATE)
):
saved = replay_state.get_checkpoint(
self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, ""),
self.checkpointer,
self.checkpoint_config,
)
# Clear RESUMING so _first re-applies input instead of resuming.
# This recreates ephemeral routing channels so nodes trigger
# naturally via version comparison.
self.config[CONF].pop(CONFIG_KEY_RESUMING, None)
else:
# Normal case: fetch the most recent checkpoint for this
# graph/thread. If a specific checkpoint_id is in the config,
# fetch that exact checkpoint; otherwise fetch the latest one.
# Returns None on first invocation (no checkpoints exist yet).
saved = self.checkpointer.get_tuple(self.checkpoint_config)
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1109,7 +1170,6 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if saved.pending_writes is not None
else []
)
self.submit = self.stack.enter_context(BackgroundExecutor(self.config))
self.channels, self.managed = channels_from_checkpoint(
self.specs, self.checkpoint
@@ -1260,10 +1320,27 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
# context manager
async def __aenter__(self) -> Self:
if self.checkpointer:
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
else:
if not self.checkpointer:
saved = None
elif self.is_nested and (
replay_state := self.config[CONF].get(CONFIG_KEY_REPLAY_STATE)
):
saved = await replay_state.aget_checkpoint(
self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, ""),
self.checkpointer,
self.checkpoint_config,
)
# Clear RESUMING so _first re-applies input instead of resuming.
# This recreates ephemeral routing channels so nodes trigger
# naturally via version comparison.
self.config[CONF].pop(CONFIG_KEY_RESUMING, None)
else:
# Normal case: fetch the most recent checkpoint for this
# graph/thread. If a specific checkpoint_id is in the config,
# fetch that exact checkpoint; otherwise fetch the latest one.
# Returns None on first invocation (no checkpoints exist yet).
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1288,7 +1365,6 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
if saved.pending_writes is not None
else []
)
self.submit = await self.stack.enter_async_context(
AsyncBackgroundExecutor(self.config)
)
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff