Implement checkpoint migration

- Migrate start:{node} channels to branch:to:{node}
- Migrate {node} channels to branch:to:{node}
This commit is contained in:
Nuno Campos
2025-04-02 12:44:59 -07:00
parent 6fe319ed1b
commit 4dda404da3
7 changed files with 1500 additions and 38 deletions
@@ -455,6 +455,7 @@ def get_checkpoint_metadata(
) -> CheckpointMetadata:
"""Get checkpoint metadata in a backwards-compatible manner."""
metadata = metadata.copy()
print("metadata", metadata)
for obj in (config.get("metadata"), config.get("configurable")):
if not obj:
continue
+89 -1
View File
@@ -2,6 +2,7 @@ import inspect
import logging
import typing
import warnings
from collections import defaultdict
from functools import partial
from inspect import isclass, isfunction, ismethod, signature
from types import FunctionType
@@ -35,7 +36,16 @@ from langgraph.channels.dynamic_barrier_value import DynamicBarrierValue, WaitFo
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.channels.last_value import LastValue
from langgraph.channels.named_barrier_value import NamedBarrierValue
from langgraph.constants import EMPTY_SEQ, MISSING, NS_END, NS_SEP, SELF, TAG_HIDDEN
from langgraph.checkpoint.base import Checkpoint
from langgraph.constants import (
EMPTY_SEQ,
INTERRUPT,
MISSING,
NS_END,
NS_SEP,
SELF,
TAG_HIDDEN,
)
from langgraph.errors import (
ErrorCode,
InvalidUpdateError,
@@ -922,6 +932,84 @@ class CompiledStateGraph(CompiledGraph):
)
)
def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None:
"""Migrate a checkpoint to new channel layout."""
values = checkpoint["channel_values"]
versions = checkpoint["channel_versions"]
seen = checkpoint["versions_seen"]
# empty checkpoints do not need migration
if not versions:
return
# current version
if checkpoint["v"] >= 3:
return
# migrate from v2 to v3
if any(k.startswith("start:") for k in versions):
# Migrate from start:node to branch:to:node
for k in list(versions):
if k.startswith("start:"):
# confirm node is present
node = k.split(":")[1]
if node not in self.nodes:
continue
# get next version
new_k = f"branch:to:{node}"
new_v = (
max(versions[new_k], versions.pop(k))
if new_k in versions
else versions.pop(k)
)
# update seen
for ss in (seen.get(node, {}), seen.get(INTERRUPT, {})):
if k in ss:
s = ss.pop(k)
if new_k in ss:
ss[new_k] = max(s, ss[new_k])
else:
ss[new_k] = s
# update value
if new_k not in values and k in values:
values[new_k] = values.pop(k)
# update version
versions[new_k] = new_v
if not set(self.nodes).isdisjoint(versions):
# Migrate from "node" to "branch:to:node"
source_to_target = defaultdict(list)
for start, end in self.builder.edges:
if start != START and end != END:
source_to_target[start].append(end)
for k in list(versions):
if k == START:
continue
if k in self.nodes:
v = versions.pop(k)
c = values.pop(k, MISSING)
for end in source_to_target[k]:
# get next version
new_k = f"branch:to:{end}"
new_v = max(versions[new_k], v) if new_k in versions else v
# update seen
for ss in (seen.get(end, {}), seen.get(INTERRUPT, {})):
if k in ss:
s = ss.pop(k)
if new_k in ss:
ss[new_k] = max(s, ss[new_k])
else:
ss[new_k] = s
# update value
if new_k not in values and c is not MISSING:
values[new_k] = c
# update version
versions[new_k] = new_v
# pop interrupt seen
if INTERRUPT in seen:
seen[INTERRUPT].pop(k, MISSING)
def _get_state_reader(
builder: StateGraph, schema: Type[Any]
+16 -6
View File
@@ -48,6 +48,7 @@ from langgraph.channels.base import (
)
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
Checkpoint,
CheckpointTuple,
copy_checkpoint,
)
@@ -766,6 +767,10 @@ class Pregel(PregelProtocol):
for name, node in self.get_subgraphs(namespace=namespace, recurse=recurse):
yield name, node
def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None:
"""Migrate a saved checkpoint to new channel layout."""
pass
def _prepare_state_snapshot(
self,
config: RunnableConfig,
@@ -784,6 +789,9 @@ class Pregel(PregelProtocol):
tasks=(),
)
# migrate checkpoint if needed
self._migrate_checkpoint(saved.checkpoint)
with ChannelsManager(
self.channels,
saved.checkpoint,
@@ -897,6 +905,9 @@ class Pregel(PregelProtocol):
tasks=(),
)
# migrate checkpoint if needed
self._migrate_checkpoint(saved.checkpoint)
async with AsyncChannelsManager(
self.channels,
saved.checkpoint,
@@ -1222,6 +1233,7 @@ class Pregel(PregelProtocol):
# get last checkpoint
config = ensure_config(self.config, input_config)
saved = checkpointer.get_tuple(config)
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint()
)
@@ -1632,6 +1644,7 @@ class Pregel(PregelProtocol):
# get last checkpoint
config = ensure_config(self.config, input_config)
saved = await checkpointer.aget_tuple(config)
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint()
)
@@ -2277,6 +2290,7 @@ class Pregel(PregelProtocol):
manager=run_manager,
debug=debug,
trigger_to_nodes=self.trigger_to_nodes,
migrate_checkpoint=self._migrate_checkpoint,
) as loop:
# create runner
runner = PregelRunner(
@@ -2570,12 +2584,8 @@ class Pregel(PregelProtocol):
interrupt_after=interrupt_after_,
manager=run_manager,
debug=debug,
# `self.nodes` can be modified after creation of `Pregel`. For example,
# that's how StateGraph compilation currently works.
# For now, we recompute the trigger_to_nodes mapping every time the
# loop is created. We could potentially memoize this if it becomes a
# performance issue.
trigger_to_nodes=_trigger_to_nodes(self.nodes),
trigger_to_nodes=self.trigger_to_nodes,
migrate_checkpoint=self._migrate_checkpoint,
) as loop:
# create runner
runner = PregelRunner(
+1 -31
View File
@@ -5,9 +5,8 @@ from langgraph.channels.base import BaseChannel
from langgraph.checkpoint.base import Checkpoint
from langgraph.checkpoint.base.id import uuid6
from langgraph.constants import MISSING
from langgraph.pregel.read import PregelNode
LATEST_VERSION = 2
LATEST_VERSION = 3
def empty_checkpoint() -> Checkpoint:
@@ -50,32 +49,3 @@ def create_checkpoint(
versions_seen=checkpoint["versions_seen"],
pending_sends=checkpoint.get("pending_sends", []),
)
def migrate_checkpoint(
checkpoint: Checkpoint,
channels: Mapping[str, BaseChannel],
nodes: Mapping[str, PregelNode],
) -> None:
"""Migrate a checkpoint to new channel layout."""
values = checkpoint["channel_values"]
versions = checkpoint["channel_versions"]
seen = checkpoint["versions_seen"]
if any(k.startswith("start:") for k in versions):
# Migrate from start:node to branch:to:node
for k in values:
if k.startswith("start:"):
node = k.split(":")[1]
new_k = f"branch:to:{node}"
if node not in nodes:
continue
v = versions.pop(k)
s = seen.get(node, {}).pop(k, None)
if s is None or s < v:
values[new_k] = values.pop(k)
# TODO handle s == v
# TODO handle new_k already in values
# TODO Migrate from "node" to "branch:to:node"
+14
View File
@@ -30,6 +30,7 @@ from typing_extensions import ParamSpec, Self
from langgraph.channels.base import BaseChannel
from langgraph.checkpoint.base import (
EXCLUDED_METADATA_KEYS,
WRITES_IDX_MAP,
BaseCheckpointSaver,
ChannelVersions,
@@ -174,6 +175,7 @@ class PregelLoop(LoopProtocol):
Any,
]
]
_migrate_checkpoint: Optional[Callable[[Checkpoint], None]]
submit: Submit
channels: Mapping[str, BaseChannel]
managed: ManagedValueMapping
@@ -211,6 +213,7 @@ class PregelLoop(LoopProtocol):
manager: Union[None, AsyncParentRunManager, ParentRunManager] = None,
input_model: Optional[Type[BaseModel]] = None,
debug: bool = False,
migrate_checkpoint: Optional[Callable[[Checkpoint], None]] = None,
trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None,
checkpoint_every_step: bool = True,
) -> None:
@@ -236,6 +239,7 @@ class PregelLoop(LoopProtocol):
CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
or CONFIG_KEY_DEDUPE_TASKS in config[CONF]
)
self._migrate_checkpoint = migrate_checkpoint
self.trigger_to_nodes = trigger_to_nodes
self.checkpoint_every_step = checkpoint_every_step
self.debug = debug
@@ -723,6 +727,8 @@ class PregelLoop(LoopProtocol):
# bail if no checkpointer
if self._checkpointer_put_after_previous is not None:
for k, v in self.config["metadata"].items():
if k in EXCLUDED_METADATA_KEYS:
continue
metadata.setdefault(k, v) # type: ignore
# create new checkpoint
@@ -899,6 +905,7 @@ class SyncPregelLoop(PregelLoop, ContextManager):
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
input_model: Optional[Type[BaseModel]] = None,
debug: bool = False,
migrate_checkpoint: Optional[Callable[[Checkpoint], None]] = None,
trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None,
) -> None:
super().__init__(
@@ -916,6 +923,7 @@ class SyncPregelLoop(PregelLoop, ContextManager):
interrupt_before=interrupt_before,
manager=manager,
debug=debug,
migrate_checkpoint=migrate_checkpoint,
trigger_to_nodes=trigger_to_nodes,
)
self.stack = ExitStack()
@@ -984,6 +992,8 @@ class SyncPregelLoop(PregelLoop, ContextManager):
saved = CheckpointTuple(
self.config, empty_checkpoint(), {"step": -2}, None, []
)
elif self._migrate_checkpoint is not None:
self._migrate_checkpoint(saved.checkpoint)
self.checkpoint_config = {
**self.config,
**saved.config,
@@ -1042,6 +1052,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
input_model: Optional[Type[BaseModel]] = None,
debug: bool = False,
migrate_checkpoint: Optional[Callable[[Checkpoint], None]] = None,
trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None,
) -> None:
super().__init__(
@@ -1059,6 +1070,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
interrupt_before=interrupt_before,
manager=manager,
debug=debug,
migrate_checkpoint=migrate_checkpoint,
trigger_to_nodes=trigger_to_nodes,
)
self.stack = AsyncExitStack()
@@ -1127,6 +1139,8 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
saved = CheckpointTuple(
self.config, empty_checkpoint(), {"step": -2}, None, []
)
elif self._migrate_checkpoint is not None:
self._migrate_checkpoint(saved.checkpoint)
self.checkpoint_config = {
**self.config,
**saved.config,
+5
View File
@@ -4,6 +4,11 @@ from typing import Any, Sequence, Union
from typing_extensions import Self
class AnyObject:
def __eq__(self, value):
return True
class FloatBetween(float):
def __new__(cls, min_value: float, max_value: float) -> Self:
return super().__new__(cls, min_value)
File diff suppressed because it is too large Load Diff