diff --git a/langgraph/checkpoint/base.py b/langgraph/checkpoint/base.py index 98f464731..74a77c6e0 100644 --- a/langgraph/checkpoint/base.py +++ b/langgraph/checkpoint/base.py @@ -67,13 +67,13 @@ class Checkpoint(TypedDict): Mapping from channel name to channel snapshot value. """ - channel_versions: defaultdict[str, Union[str, int, float]] + channel_versions: dict[str, Union[str, int, float]] """The versions of the channels at the time of the checkpoint. The keys are channel names and the values are the logical time step at which the channel was last updated. """ - versions_seen: defaultdict[str, defaultdict[str, Union[str, int, float]]] + versions_seen: defaultdict[str, dict[str, Union[str, int, float]]] """Map from node ID to map from channel name to version seen. This keeps track of the versions of the channels that each node has seen. @@ -85,18 +85,14 @@ class Checkpoint(TypedDict): Cleared by the next checkpoint.""" -def _seen_dict(): - return defaultdict(int) - - def empty_checkpoint() -> Checkpoint: return Checkpoint( v=1, id=str(uuid6(clock_seq=-2)), ts=datetime.now(timezone.utc).isoformat(), channel_values={}, - channel_versions=defaultdict(int), - versions_seen=defaultdict(_seen_dict), + channel_versions={}, + versions_seen=defaultdict(dict), pending_sends=[], ) @@ -107,10 +103,10 @@ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint: ts=checkpoint["ts"], id=checkpoint["id"], channel_values=checkpoint["channel_values"].copy(), - channel_versions=defaultdict(int, checkpoint["channel_versions"]), + channel_versions=checkpoint["channel_versions"].copy(), versions_seen=defaultdict( - _seen_dict, - {k: defaultdict(int, v) for k, v in checkpoint["versions_seen"].items()}, + dict, + {k: v.copy() for k, v in checkpoint["versions_seen"].items()}, ), pending_sends=checkpoint.get("pending_sends", []).copy(), ) diff --git a/langgraph/checkpoint/sqlite.py b/langgraph/checkpoint/sqlite.py index 3836a041f..d7ce4e9c8 100644 --- a/langgraph/checkpoint/sqlite.py +++ b/langgraph/checkpoint/sqlite.py @@ -3,12 +3,14 @@ import pickle import sqlite3 import threading from contextlib import AbstractContextManager, contextmanager +from hashlib import md5 from types import TracebackType from typing import Any, AsyncIterator, Dict, Iterator, Optional, Sequence, Tuple from langchain_core.runnables import RunnableConfig from typing_extensions import Self +from langgraph.channels.base import BaseChannel from langgraph.checkpoint.base import ( BaseCheckpointSaver, Checkpoint, @@ -431,6 +433,15 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): """ raise NotImplementedError(_AIO_ERROR_MSG) + def get_next_version(self, current: Optional[str], channel: BaseChannel) -> str: + if current is None: + current_v = 1 + else: + current_v = int(current.split(".")[0]) + next_v = current_v + 1 + next_h = md5(self.serde.dumps(channel.checkpoint())).hexdigest() + return f"{next_v:032}.{next_h}" + def _metadata_predicate( metadata_filter: Dict[str, Any], diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index f8e86206f..29c6ebc36 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -878,8 +878,9 @@ class Pregel( # past previous interrupt, if any checkpoint = copy_checkpoint(checkpoint) for k in self.stream_channels_list: - version = checkpoint["channel_versions"][k] - checkpoint["versions_seen"][INTERRUPT][k] = version + if k in checkpoint["channel_versions"]: + version = checkpoint["channel_versions"][k] + checkpoint["versions_seen"][INTERRUPT][k] = version # Similarly to Bulk Synchronous Parallel / Pregel model # computation proceeds in steps, while there are channel updates @@ -1218,8 +1219,9 @@ class Pregel( # past previous interrupt, if any checkpoint = copy_checkpoint(checkpoint) for k in self.stream_channels_list: - version = checkpoint["channel_versions"][k] - checkpoint["versions_seen"][INTERRUPT][k] = version + if k in checkpoint["channel_versions"]: + version = checkpoint["channel_versions"][k] + checkpoint["versions_seen"][INTERRUPT][k] = version # Similarly to Bulk Synchronous Parallel / Pregel model # computation proceeds in steps, while there are channel updates @@ -1537,12 +1539,15 @@ def _should_interrupt( snapshot_channels: Sequence[str], tasks: list[PregelExecutableTask], ) -> bool: + version_type = type(next(iter(checkpoint["channel_versions"].values()), None)) + null_version = version_type() # defaultdicts are mutated on access :( so we need to copy - seen = checkpoint["versions_seen"].copy()[INTERRUPT].copy() + seen = checkpoint["versions_seen"].copy()[INTERRUPT] return ( # interrupt if any of snapshopt_channels has been updated since last interrupt any( - checkpoint["channel_versions"][chan] > seen[chan] + checkpoint["channel_versions"].get(chan, null_version) + > seen.get(chan, null_version) for chan in snapshot_channels ) # and any triggered node is in interrupt_nodes list @@ -1568,7 +1573,7 @@ def _local_read( if fresh: checkpoint = create_checkpoint(checkpoint, channels, -1) with ChannelsManager(channels, checkpoint) as channels: - _apply_writes(copy_checkpoint(checkpoint), channels, writes, _increment) + _apply_writes(copy_checkpoint(checkpoint), channels, writes, None) return read_channels(channels, select) else: return read_channels(channels, select) @@ -1601,7 +1606,7 @@ def _apply_writes( checkpoint: Checkpoint, channels: Mapping[str, BaseChannel], pending_writes: Sequence[tuple[str, Any]], - get_next_version: Callable[[int, BaseChannel], int], + get_next_version: Optional[Callable[[int, BaseChannel], int]], ) -> None: if checkpoint["pending_sends"]: checkpoint["pending_sends"].clear() @@ -1618,7 +1623,7 @@ def _apply_writes( if checkpoint["channel_versions"]: max_version = max(checkpoint["channel_versions"].values()) else: - max_version = 0 + max_version = None updated_channels: set[str] = set() # Apply writes to channels @@ -1630,9 +1635,10 @@ def _apply_writes( raise InvalidUpdateError( f"Invalid update for channel {chan} with values {vals}" ) from e - checkpoint["channel_versions"][chan] = get_next_version( - max_version or None, channels[chan] - ) + if get_next_version is not None: + checkpoint["channel_versions"][chan] = get_next_version( + max_version, channels[chan] + ) updated_channels.add(chan) # Channels that weren't updated in this step are notified of a new step for chan in channels: @@ -1730,6 +1736,10 @@ def _prepare_next_tasks( channels_to_consume = set() # Check if any processes should be run in next step # If so, prepare the values to be passed to them + version_type = type(next(iter(checkpoint["channel_versions"].values()), None)) + null_version = version_type() + if null_version is None: + return checkpoint, tasks for name, proc in processes.items(): seen = checkpoint["versions_seen"][name] # If any of the channels read by this process were updated @@ -1739,7 +1749,8 @@ def _prepare_next_tasks( if not isinstance( read_channel(channels, chan, return_exception=True), EmptyChannelError ) - and checkpoint["channel_versions"][chan] > seen[chan] + and checkpoint["channel_versions"].get(chan, null_version) + > seen.get(chan, null_version) ]: channels_to_consume.update(triggers) try: @@ -1753,6 +1764,7 @@ def _prepare_next_tasks( { chan: checkpoint["channel_versions"][chan] for chan in proc.triggers + if chan in checkpoint["channel_versions"] } )