codekingpro/portable-devtools
114k
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 