Files
langgraph/langgraph/checkpoint/base.py
T
Nuno Campos f304908102 Add metadata to checkpoints
- not yet used in this PR
2024-05-06 11:52:49 -07:00

174 lines
4.7 KiB
Python

from abc import ABC
from collections import defaultdict
from datetime import datetime, timezone
from typing import (
Any,
AsyncIterator,
Iterator,
NamedTuple,
Optional,
TypedDict,
)
from langchain_core.runnables import ConfigurableFieldSpec, RunnableConfig
from langgraph.serde.base import SerializerProtocol
from langgraph.serde.jsonplus import JsonPlusSerializer
from langgraph.utils import StrEnum
class Checkpoint(TypedDict):
"""State snapshot at a given point in time."""
v: int
"""The version of the checkpoint format. Currently 1."""
ts: str
"""The timestamp of the checkpoint in ISO 8601 format."""
channel_values: dict[str, Any]
"""The values of the channels at the time of the checkpoint.
Mapping from channel name to channel snapshot value.
"""
channel_versions: defaultdict[str, int]
"""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]]
"""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.
Used to determine which nodes to execute next.
"""
def _seen_dict():
return defaultdict(int)
def empty_checkpoint() -> Checkpoint:
return Checkpoint(
v=1,
ts=datetime.now(timezone.utc).isoformat(),
channel_values={},
channel_versions=defaultdict(int),
versions_seen=defaultdict(_seen_dict),
)
def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
return Checkpoint(
v=checkpoint["v"],
ts=checkpoint["ts"],
channel_values=checkpoint["channel_values"].copy(),
channel_versions=defaultdict(int, checkpoint["channel_versions"]),
versions_seen=defaultdict(
_seen_dict,
{k: defaultdict(int, v) for k, v in checkpoint["versions_seen"].items()},
),
)
class CheckpointAt(StrEnum):
"""When to take a checkpoint."""
END_OF_STEP = "end_of_step"
"""Take a checkpoint at the end of each step."""
END_OF_RUN = "end_of_run"
"""Take a checkpoint at the end of the run."""
class CheckpointTuple(NamedTuple):
config: RunnableConfig
checkpoint: Checkpoint
metadata: Optional[dict[str, Any]]
parent_config: Optional[RunnableConfig] = None
CheckpointThreadId = ConfigurableFieldSpec(
id="thread_id",
annotation=str,
name="Thread ID",
description=None,
default="",
is_shared=True,
)
CheckpointThreadTs = ConfigurableFieldSpec(
id="thread_ts",
annotation=Optional[str],
name="Thread Timestamp",
description="Pass to fetch a past checkpoint. If None, fetches the latest checkpoint.",
default=None,
is_shared=True,
)
class BaseCheckpointSaver(ABC):
at: CheckpointAt = CheckpointAt.END_OF_STEP
serde: SerializerProtocol = JsonPlusSerializer()
def __init__(
self,
*,
serde: Optional[SerializerProtocol] = None,
at: Optional[CheckpointAt] = None,
) -> None:
self.serde = serde or self.serde
self.at = at or self.at
@property
def config_specs(self) -> list[ConfigurableFieldSpec]:
return [CheckpointThreadId, CheckpointThreadTs]
def get(self, config: RunnableConfig) -> Optional[Checkpoint]:
if value := self.get_tuple(config):
return value.checkpoint
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
raise NotImplementedError
def list(
self,
config: RunnableConfig,
*,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
) -> Iterator[CheckpointTuple]:
raise NotImplementedError
def put(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: dict[str, Any],
) -> RunnableConfig:
raise NotImplementedError
async def aget(self, config: RunnableConfig) -> Optional[Checkpoint]:
if value := await self.aget_tuple(config):
return value.checkpoint
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
raise NotImplementedError
async def alist(
self,
config: RunnableConfig,
*,
before: Optional[RunnableConfig] = None,
limit: Optional[int] = None,
) -> AsyncIterator[CheckpointTuple]:
raise NotImplementedError
async def aput(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: dict[str, Any],
) -> RunnableConfig:
raise NotImplementedError