diff --git a/libs/langgraph/langgraph/_internal/_constants.py b/libs/langgraph/langgraph/_internal/_constants.py index 00a2db4b4..97b228fea 100644 --- a/libs/langgraph/langgraph/_internal/_constants.py +++ b/libs/langgraph/langgraph/_internal/_constants.py @@ -42,7 +42,7 @@ CONFIG_KEY_CACHE = sys.intern("__pregel_cache") CONFIG_KEY_RESUMING = sys.intern("__pregel_resuming") # holds a boolean indicating if subgraphs should resume from a previous checkpoint CONFIG_KEY_REPLAYING = sys.intern("__pregel_replaying") -# holds a boolean indicating if subgraphs should replay (re-run tasks, drop cached RESUME writes) +# signals to subgraphs that the parent is replaying from an earlier 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") diff --git a/libs/langgraph/langgraph/pregel/_algo.py b/libs/langgraph/langgraph/pregel/_algo.py index 55d212391..01ef55321 100644 --- a/libs/langgraph/langgraph/pregel/_algo.py +++ b/libs/langgraph/langgraph/pregel/_algo.py @@ -40,7 +40,6 @@ from langgraph._internal._constants import ( CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_CHECKPOINTER, CONFIG_KEY_READ, - CONFIG_KEY_REPLAYING, CONFIG_KEY_RESUME_MAP, CONFIG_KEY_RUNTIME, CONFIG_KEY_SCRATCHPAD, @@ -581,14 +580,12 @@ def prepare_single_task( if checkpoint_null_version is None: return # If any of the channels read by this process were updated. - is_replaying = configurable.get(CONFIG_KEY_REPLAYING, False) if _triggers( channels, checkpoint["channel_versions"], checkpoint["versions_seen"].get(name), checkpoint_null_version, proc, - is_replaying=is_replaying, ): triggers = tuple(sorted(proc.triggers)) # create task id @@ -1039,10 +1036,8 @@ def _triggers( seen: ChannelVersions | None, null_version: V, proc: PregelNode, - *, - is_replaying: bool = False, ) -> bool: - if is_replaying or seen is None: + if seen is None: for chan in proc.triggers: if channels[chan].is_available(): return True diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 6f294aea4..dfa89503b 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -748,11 +748,14 @@ class PregelLoop: elif CONFIG_KEY_RESUMING not in configurable: raise EmptyInputError(f"Received no input for {input_keys}") # Propagate resuming and replaying flags to subgraphs. + # When replaying, don't tell subgraphs to resume — they should + # re-apply input so that triggers fire naturally. if not self.is_nested: self.config = patch_configurable( self.config, { - CONFIG_KEY_RESUMING: is_resuming, + CONFIG_KEY_RESUMING: is_resuming + and not self.is_replaying, CONFIG_KEY_REPLAYING: self.is_replaying, }, ) @@ -1106,22 +1109,23 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): }, ) - def _get_checkpoint_after_parent(self) -> CheckpointTuple | None: - """Find the right subgraph checkpoint to restore when the parent replays. + def _get_checkpoint_before_invocation(self) -> RunnableConfig | None: + """Find the config for the subgraph checkpoint from just before a + previous invocation. - Each time the parent invokes a subgraph, the subgraph creates a series - of checkpoints. Every checkpoint records which parent checkpoint was - active when it was created (in `metadata["parents"]`). The first - checkpoint in each invocation has `source="input"` and contains the - accumulated channel_values from prior invocations but hasn't run any - nodes yet. + When a parent graph replays from an earlier checkpoint, it re-invokes + the subgraph with the same input. Instead of loading the subgraph's + latest checkpoint (which may be from a later parent step), we find the + state the subgraph was in *before* the original invocation so that + `_first()` can re-apply the input naturally — no trigger hacks needed. - We query for `source="input"` + `parents={parent_ns: parent_id}` to - find the starting checkpoint from the invocation that ran under the - given parent checkpoint. Node re-triggering is handled by `_triggers` - which skips `versions_seen` when `CONFIG_KEY_REPLAYING` is set. + We find the `source="input"` checkpoint that was created under the + matching parent checkpoint, then return the config for its parent + (the pre-input state). The caller uses this config with `get_tuple()` + to load the actual checkpoint. - Returns None to start fresh if no such checkpoint exists.""" + Returns None to start fresh if no matching checkpoint exists or if + this is the subgraph's first invocation.""" checkpoint_map = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {}) parent_ns = NS_SEP.join(self.checkpoint_ns[:-1]) if self.checkpoint_ns else "" parent_checkpoint_id = checkpoint_map.get(parent_ns) @@ -1135,8 +1139,8 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): }, limit=1, ): - return saved - return None + return saved.parent_config # None for first invocation → start fresh + return None # no matching checkpoint (e.g. fork) — start fresh # context manager @@ -1148,11 +1152,14 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): if not self.checkpointer: saved = None elif is_subgraph_replay: - # Subgraph replay: the parent graph is replaying from an earlier - # checkpoint, so we need to restore the subgraph checkpoint that - # corresponds to that parent checkpoint — not the subgraph's - # latest. We find it by matching the parent checkpoint timeline. - saved = self._get_checkpoint_after_parent() + # Subgraph replay: load the pre-input checkpoint so _first() + # can re-apply input naturally — triggers fire without hacks. + pre_input_config = self._get_checkpoint_before_invocation() + saved = ( + self.checkpointer.get_tuple(pre_input_config) + if pre_input_config + else None + ) else: # Normal case: fetch the most recent checkpoint for this # graph/thread. If a specific checkpoint_id is in the config, @@ -1331,8 +1338,8 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): }, ) - async def _aget_checkpoint_after_parent(self) -> CheckpointTuple | None: - """Async version of `_get_checkpoint_after_parent`.""" + async def _aget_checkpoint_before_invocation(self) -> RunnableConfig | None: + """Async version of `_get_checkpoint_before_invocation`.""" checkpoint_map = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {}) parent_ns = NS_SEP.join(self.checkpoint_ns[:-1]) if self.checkpoint_ns else "" parent_checkpoint_id = checkpoint_map.get(parent_ns) @@ -1346,8 +1353,8 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): }, limit=1, ): - return saved - return None + return saved.parent_config # None for first invocation → start fresh + return None # no matching checkpoint (e.g. fork) — start fresh # context manager @@ -1359,11 +1366,14 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): if not self.checkpointer: saved = None elif is_subgraph_replay: - # Subgraph replay: the parent graph is replaying from an earlier - # checkpoint, so we need to restore the subgraph checkpoint that - # corresponds to that parent checkpoint — not the subgraph's - # latest. We find it by matching the parent checkpoint timeline. - saved = await self._aget_checkpoint_after_parent() + # Subgraph replay: load the pre-input checkpoint so _first() + # can re-apply input naturally — triggers fire without hacks. + pre_input_config = await self._aget_checkpoint_before_invocation() + saved = ( + await self.checkpointer.aget_tuple(pre_input_config) + if pre_input_config + else None + ) else: # Normal case: fetch the most recent checkpoint for this # graph/thread. If a specific checkpoint_id is in the config,