Support non-int versions, add tests for an example str version

This commit is contained in:
Nuno Campos
2024-06-13 18:38:02 -07:00
parent d98e90a13e
commit 033c29f55a
3 changed files with 43 additions and 24 deletions
+7 -11
View File
@@ -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(),
)
+11
View File
@@ -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],
+25 -13
View File
@@ -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"]
}
)