Support float and str versions

This commit is contained in:
Nuno Campos
2024-06-13 18:14:30 -07:00
parent a8b14a4e04
commit d98e90a13e
2 changed files with 11 additions and 7 deletions
+8 -4
View File
@@ -11,6 +11,8 @@ from typing import (
NamedTuple,
Optional,
TypedDict,
TypeVar,
Union,
)
from langchain_core.runnables import ConfigurableFieldSpec, RunnableConfig
@@ -21,6 +23,8 @@ from langgraph.constants import Send
from langgraph.serde.base import SerializerProtocol
from langgraph.serde.jsonplus import JsonPlusSerializer
V = TypeVar("V", int, float, str)
# Marked as total=False to allow for future expansion.
class CheckpointMetadata(TypedDict, total=False):
@@ -63,13 +67,13 @@ class Checkpoint(TypedDict):
Mapping from channel name to channel snapshot value.
"""
channel_versions: defaultdict[str, int]
channel_versions: defaultdict[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, int]]
versions_seen: defaultdict[str, defaultdict[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.
@@ -203,5 +207,5 @@ class BaseCheckpointSaver(ABC):
) -> RunnableConfig:
raise NotImplementedError
def get_next_version(self, current: int, channel: BaseChannel) -> int:
return current + 1
def get_next_version(self, current: Optional[V], channel: BaseChannel) -> V:
return current + 1 if current is not None else 1
+3 -3
View File
@@ -1593,8 +1593,8 @@ def _local_write(
commit(writes)
def _increment(current: int, channel: BaseChannel) -> int:
return current + 1
def _increment(current: Optional[int], channel: BaseChannel) -> int:
return current + 1 if current is not None else 1
def _apply_writes(
@@ -1631,7 +1631,7 @@ def _apply_writes(
f"Invalid update for channel {chan} with values {vals}"
) from e
checkpoint["channel_versions"][chan] = get_next_version(
max_version, channels[chan]
max_version or None, channels[chan]
)
updated_channels.add(chan)
# Channels that weren't updated in this step are notified of a new step