mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-23 16:12:25 +02:00
Implement checkpoint migration
- Migrate start:{node} channels to branch:to:{node}
- Migrate {node} channels to branch:to:{node}
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user