Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_checkpoint.py239 linesDownload Raw Back to pregel
1from __future__ import annotations2 3from collections.abc import Callable, Mapping4from datetime import datetime, timezone5from typing import Any, cast6 7from langchain_core.runnables import RunnableConfig8from langgraph.checkpoint.base import (9    BaseCheckpointSaver,10    Checkpoint,11)12from langgraph.checkpoint.base.id import uuid613from langgraph.checkpoint.serde.types import _DeltaSnapshot14 15from langgraph._internal._config import DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT16from langgraph._internal._typing import MISSING17from langgraph.channels.base import BaseChannel18from langgraph.channels.delta import DeltaChannel19from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec20 21LATEST_VERSION = 422 23GetNextVersion = Callable[[Any, None], Any]24 25 26def empty_checkpoint() -> Checkpoint:27    return Checkpoint(28        v=LATEST_VERSION,29        id=str(uuid6(clock_seq=-2)),30        ts=datetime.now(timezone.utc).isoformat(),31        channel_values={},32        channel_versions={},33        versions_seen={},34    )35 36 37def delta_channels_to_snapshot(38    channels: Mapping[str, BaseChannel],39    counters_since_delta_snapshot: Mapping[str, tuple[int, int]],40) -> set[str]:41    """Return the set of DeltaChannel names that should snapshot now.42 43    A channel snapshots when EITHER its accumulated update count reaches44    `snapshot_frequency` OR the total supersteps since its last snapshot45    reaches `DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT`. This is a pure46    predicate — no mutation.47    """48    result: set[str] = set()49    for name, ch in channels.items():50        if not isinstance(ch, DeltaChannel) or not ch.is_available():51            continue52        updates, supersteps = counters_since_delta_snapshot.get(name, (0, 0))53        if (54            updates >= ch.snapshot_frequency55            or supersteps >= DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT56        ):57            result.add(name)58    return result59 60 61def create_checkpoint(62    checkpoint: Checkpoint,63    channels: Mapping[str, BaseChannel] | None,64    step: int,65    *,66    id: str | None = None,67    updated_channels: set[str] | None = None,68    get_next_version: GetNextVersion | None = None,69    channels_to_snapshot: set[str] | None = None,70) -> Checkpoint:71    """Build a new Checkpoint from the previous one and live channel state.72 73    For each name in `channels_to_snapshot`, a `_DeltaSnapshot(value)` blob74    is written into `channel_values[k]`. Other delta channels are omitted75    from `channel_values` — the ancestor walk reconstructs their state76    from `checkpoint_writes`. Callers compute the set via77    `delta_channels_to_snapshot(channels, counters)`; defaults to empty78    (no snapshots) when not provided.79    """80    ts = datetime.now(timezone.utc).isoformat()81    channels_to_snapshot = channels_to_snapshot or set()82    if channels is None:83        values = checkpoint["channel_values"]84        channel_versions = checkpoint["channel_versions"]85    else:86        values = {}87        channel_versions = dict(checkpoint["channel_versions"])88        for k in channels:89            if k not in channel_versions:90                continue91            ch = channels[k]92            if k in channels_to_snapshot:93                # In exit mode, the snapshot decision is deferred to exit94                # time (intermediate steps have do_checkpoint=False). The95                # channel's count may have reached snapshot_frequency over96                # several supersteps, but the LAST superstep may not have97                # written to this channel. In that case apply_writes()98                # (in _algo.py) didn't bump this channel's version, so99                # saver.put() wouldn't include it in new_versions and100                # the snapshot blob would be silently dropped. The manual101                # bump below closes the gap. In sync/async durability this102                # branch is effectively dead code (the step that pushes103                # the count to freq always writes the channel).104                if get_next_version is not None and (105                    updated_channels is None or k not in updated_channels106                ):107                    channel_versions[k] = get_next_version(channel_versions[k], None)108                values[k] = _DeltaSnapshot(ch.get())109            else:110                v = ch.checkpoint()111                if v is not MISSING:112                    values[k] = v113    return Checkpoint(114        v=LATEST_VERSION,115        ts=ts,116        id=id or str(uuid6(clock_seq=step)),117        channel_values=values,118        channel_versions=channel_versions,119        versions_seen=checkpoint["versions_seen"],120        updated_channels=None if updated_channels is None else sorted(updated_channels),121    )122 123 124def _needs_replay(spec: BaseChannel, stored: object) -> bool:125    """True if `spec` is a `DeltaChannel` and no value is stored at this126    checkpoint, requiring an ancestor walk to reconstruct.127 128    `_DeltaSnapshot` blobs and plain values (migration) resolve directly via129    `from_checkpoint` — only absence (`MISSING`) triggers replay.130    """131    if not isinstance(spec, DeltaChannel):132        return False133    return stored is MISSING134 135 136def channels_from_checkpoint(137    specs: Mapping[str, BaseChannel | ManagedValueSpec],138    checkpoint: Checkpoint,139    *,140    saver: BaseCheckpointSaver | None = None,141    config: RunnableConfig | None = None,142) -> tuple[Mapping[str, BaseChannel], ManagedValueMapping]:143    """Hydrate channels from a checkpoint.144 145    For most channels, `spec.from_checkpoint(checkpoint["channel_values"][k])`146    is sufficient. `DeltaChannel` is the exception: when the channel is147    absent from `channel_values`, an ancestor walk via148    `saver.get_delta_channel_history` is required to find the nearest seed149    (`_DeltaSnapshot` blob or pre-migration plain value) and accumulate150    the writes between it and the target. All delta channels needing151    replay are batched into a single saver call.152    """153    channel_specs: dict[str, BaseChannel] = {}154    managed_specs: dict[str, ManagedValueSpec] = {}155    for k, v in specs.items():156        if isinstance(v, BaseChannel):157            channel_specs[k] = v158        else:159            managed_specs[k] = v160 161    delta_channels: list[str] = [162        k163        for k, spec in channel_specs.items()164        if _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING))165    ]166    histories: Mapping[str, Any] = {}167    if delta_channels and saver is not None and config is not None:168        histories = saver.get_delta_channel_history(169            config=config, channels=delta_channels170        )171 172    channels: dict[str, BaseChannel] = {}173    for k, spec in channel_specs.items():174        ch: BaseChannel175        if k in histories:176            delta_spec = cast(DeltaChannel, spec)177            history = histories[k]178            replay_ch = delta_spec.from_checkpoint(history.get("seed", MISSING))179            replay_ch.replay_writes(history["writes"])180            ch = replay_ch181        else:182            ch = spec.from_checkpoint(checkpoint["channel_values"].get(k, MISSING))183        channels[k] = ch184    return channels, managed_specs185 186 187async def achannels_from_checkpoint(188    specs: Mapping[str, BaseChannel | ManagedValueSpec],189    checkpoint: Checkpoint,190    *,191    saver: BaseCheckpointSaver | None = None,192    config: RunnableConfig | None = None,193) -> tuple[Mapping[str, BaseChannel], ManagedValueMapping]:194    """Async version of `channels_from_checkpoint`. See docstring there."""195    channel_specs: dict[str, BaseChannel] = {}196    managed_specs: dict[str, ManagedValueSpec] = {}197    for k, v in specs.items():198        if isinstance(v, BaseChannel):199            channel_specs[k] = v200        else:201            managed_specs[k] = v202 203    delta_channels: list[str] = [204        k205        for k, spec in channel_specs.items()206        if _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING))207    ]208    histories: Mapping[str, Any] = {}209    if delta_channels and saver is not None and config is not None:210        histories = await saver.aget_delta_channel_history(211            config=config, channels=delta_channels212        )213 214    channels: dict[str, BaseChannel] = {}215    for k, spec in channel_specs.items():216        ch: BaseChannel217        if k in histories:218            delta_spec = cast(DeltaChannel, spec)219            history = histories[k]220            replay_ch = delta_spec.from_checkpoint(history.get("seed", MISSING))221            replay_ch.replay_writes(history["writes"])222            ch = replay_ch223        else:224            ch = spec.from_checkpoint(checkpoint["channel_values"].get(k, MISSING))225        channels[k] = ch226    return channels, managed_specs227 228 229def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:230    return Checkpoint(231        v=checkpoint["v"],232        ts=checkpoint["ts"],233        id=checkpoint["id"],234        channel_values=checkpoint["channel_values"].copy(),235        channel_versions=checkpoint["channel_versions"].copy(),236        versions_seen={k: v.copy() for k, v in checkpoint["versions_seen"].items()},237        updated_channels=checkpoint.get("updated_channels", None),238    )239 
codekingpro/portable-devtools · Team Ai