alt fix idea

This commit is contained in:
Sydney Runkle
2026-03-05 08:54:10 -08:00
parent 207dccf5b3
commit 6a92b7ff3c
2 changed files with 95 additions and 39 deletions
@@ -41,6 +41,8 @@ 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_REPLAYING = sys.intern("__pregel_replaying")
# holds a boolean indicating if subgraphs should replay (re-run tasks, drop cached RESUME writes)
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 +100,7 @@ RESERVED = {
CONFIG_KEY_STREAM,
CONFIG_KEY_CHECKPOINT_MAP,
CONFIG_KEY_RESUMING,
CONFIG_KEY_REPLAYING,
CONFIG_KEY_TASK_ID,
CONFIG_KEY_CHECKPOINT_MAP,
CONFIG_KEY_CHECKPOINT_ID,
+92 -39
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_REPLAYING,
CONFIG_KEY_RESUME_MAP,
CONFIG_KEY_RESUMING,
CONFIG_KEY_SCRATCHPAD,
@@ -244,7 +245,9 @@ 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.skip_done_tasks = CONFIG_KEY_CHECKPOINT_ID not in config[
CONF
] and not config[CONF].get(CONFIG_KEY_REPLAYING, False)
self._migrate_checkpoint = migrate_checkpoint
self.trigger_to_nodes = trigger_to_nodes
self.retry_policy = retry_policy
@@ -559,16 +562,28 @@ class PregelLoop:
self.checkpoint_pending_writes.clear()
# "not skip_done_tasks" only applies to first tick after resuming
self.skip_done_tasks = True
# collect child checkpoint mappings from subgraph tasks
children: dict[str, str] = {}
parent_ns = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
for task in self.tasks.values():
task_map = task.config.get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_MAP, {})
for ns, ckpt_id in task_map.items():
if ns != parent_ns:
children[ns] = ckpt_id
# save checkpoint
self._put_checkpoint({"source": "loop"})
metadata: CheckpointMetadata = {"source": "loop"}
if children:
metadata["children"] = children
self._put_checkpoint(metadata)
# after execution, check if we should interrupt
if self.interrupt_after and should_interrupt(
self.checkpoint, self.interrupt_after, self.tasks.values()
):
self.status = "interrupt_after"
raise GraphInterrupt()
# unset resuming flag
# unset resuming/replaying flags
self.config[CONF].pop(CONFIG_KEY_RESUMING, None)
self.config[CONF].pop(CONFIG_KEY_REPLAYING, None)
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
raise NotImplementedError
@@ -729,18 +744,27 @@ class PregelLoop:
self._put_checkpoint({"source": "input"})
elif CONFIG_KEY_RESUMING not in configurable:
raise EmptyInputError(f"Received no input for {input_keys}")
# Propagate resuming flag to subgraphs (only the outer graph does this).
# Propagate resuming and replaying flags to subgraphs.
if not self.is_nested:
has_resume_value = (
isinstance(self.input, Command) and self.input.resume is not None
)
# When forking (skip_done_tasks=False, i.e. specific checkpoint_id),
# subgraphs should NOT resume — they start fresh.
# When genuinely resuming from latest, subgraphs should also resume.
is_fork = not self.skip_done_tasks
subgraph_should_resume = has_resume_value or (is_resuming and not is_fork)
is_replaying = not self.skip_done_tasks
patch: dict[str, Any] = {
CONFIG_KEY_RESUMING: is_resuming,
CONFIG_KEY_REPLAYING: is_replaying,
}
# Load child checkpoint mappings from metadata so that
# checkpointer=True subgraphs can load the correct checkpoint
# during replay (instead of always loading the latest).
if is_replaying:
children = self.checkpoint_metadata.get("children", {})
if children:
current_map = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {})
patch[CONFIG_KEY_CHECKPOINT_MAP] = {
**current_map,
**children,
}
self.config = patch_configurable(
self.config, {CONFIG_KEY_RESUMING: subgraph_should_resume}
self.config,
patch,
)
# set flag
self.status = "pending"
@@ -886,6 +910,19 @@ class PregelLoop:
elif exc_type is None:
# save final output
self.output = read_channels(self.channels, self.output_keys)
# Write back checkpoint_id to CONFIG_KEY_CHECKPOINT_MAP so the parent
# can record which subgraph checkpoint corresponds to its own checkpoint.
# Only for checkpointer=True subgraphs (recast ns has no NS_END).
if (
self.is_nested
and self.checkpoint_ns
and all(NS_END not in part for part in self.checkpoint_ns)
and CONFIG_KEY_CHECKPOINT_MAP in self.config[CONF]
):
recast_ns = NS_SEP.join(self.checkpoint_ns)
self.config[CONF][CONFIG_KEY_CHECKPOINT_MAP][recast_ns] = self.checkpoint[
"id"
]
def _emit(
self,
@@ -1099,6 +1136,14 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
saved = self.checkpointer.get_tuple(self.checkpoint_config)
else:
saved = None
# When replaying a subgraph that wasn't in the checkpoint map
# (parent checkpoint predates this subgraph), start fresh.
if (
saved is not None
and self.config[CONF].get(CONFIG_KEY_REPLAYING)
and not self.checkpoint_config.get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_ID)
):
saved = None
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1123,19 +1168,19 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if saved.pending_writes is not None
else []
)
# When replaying from a specific checkpoint (fork), drop cached
# RESUME writes so that interrupt() calls re-fire instead of
# returning stale values. But if the input directly carries a
# resume value, keep them — multi-interrupt scenarios need
# previously resolved RESUME values preserved.
is_replaying = not self.skip_done_tasks
has_resume_value = (
self.config.get(CONF, {}).get(CONFIG_KEY_RESUMING) is True
) or (isinstance(self.input, Command) and self.input.resume is not None)
if is_replaying and not has_resume_value:
self.checkpoint_pending_writes = [
w for w in self.checkpoint_pending_writes if w[1] != RESUME
]
# When replaying from a specific checkpoint, drop cached RESUME
# writes so that interrupt() calls re-fire instead of returning
# stale values. But if a resume value is being provided (e.g.
# Command(resume=...) on a specific checkpoint), keep them —
# multi-interrupt scenarios need previously resolved values preserved.
if not self.skip_done_tasks:
has_resume_value = (
isinstance(self.input, Command) and self.input.resume is not None
) or self.config.get(CONF, {}).get(CONFIG_KEY_RESUMING, False)
if not has_resume_value:
self.checkpoint_pending_writes = [
w for w in self.checkpoint_pending_writes if w[1] != RESUME
]
self.submit = self.stack.enter_context(BackgroundExecutor(self.config))
self.channels, self.managed = channels_from_checkpoint(
@@ -1291,6 +1336,14 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
else:
saved = None
# When replaying a subgraph that wasn't in the checkpoint map
# (parent checkpoint predates this subgraph), start fresh.
if (
saved is not None
and self.config[CONF].get(CONFIG_KEY_REPLAYING)
and not self.checkpoint_config.get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_ID)
):
saved = None
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1315,19 +1368,19 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
if saved.pending_writes is not None
else []
)
# When replaying from a specific checkpoint (fork), drop cached
# RESUME writes so that interrupt() calls re-fire instead of
# returning stale values. But if the input directly carries a
# resume value, keep them — multi-interrupt scenarios need
# previously resolved RESUME values preserved.
is_replaying = not self.skip_done_tasks
has_resume_value = (
self.config.get(CONF, {}).get(CONFIG_KEY_RESUMING) is True
) or (isinstance(self.input, Command) and self.input.resume is not None)
if is_replaying and not has_resume_value:
self.checkpoint_pending_writes = [
w for w in self.checkpoint_pending_writes if w[1] != RESUME
]
# When replaying from a specific checkpoint, drop cached RESUME
# writes so that interrupt() calls re-fire instead of returning
# stale values. But if a resume value is being provided (e.g.
# Command(resume=...) on a specific checkpoint), keep them —
# multi-interrupt scenarios need previously resolved values preserved.
if not self.skip_done_tasks:
has_resume_value = (
isinstance(self.input, Command) and self.input.resume is not None
) or self.config.get(CONF, {}).get(CONFIG_KEY_RESUMING, False)
if not has_resume_value:
self.checkpoint_pending_writes = [
w for w in self.checkpoint_pending_writes if w[1] != RESUME
]
self.submit = await self.stack.enter_async_context(
AsyncBackgroundExecutor(self.config)