From 9fd6374302ef4a465d028eee5d9896d41a87cbf9 Mon Sep 17 00:00:00 2001 From: Sydney Runkle Date: Tue, 21 Apr 2026 09:37:49 -0400 Subject: [PATCH] feat(pregel): add _assemble_delta_channels helpers for universal DeltaChannel support --- .../langgraph/langgraph/pregel/_checkpoint.py | 167 +++++++++++++++++- 1 file changed, 166 insertions(+), 1 deletion(-) diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index 1b327203d..cdb69ad81 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -1,9 +1,18 @@ from __future__ import annotations +import logging from collections.abc import Mapping from datetime import datetime, timezone +from typing import Any -from langgraph.checkpoint.base import Checkpoint +from langchain_core.runnables import RunnableConfig + +from langgraph.checkpoint.base import ( + BaseCheckpointSaver, + Checkpoint, + DeltaChainValue, + DeltaValue, +) from langgraph.checkpoint.base.id import uuid6 from langgraph._internal._typing import MISSING @@ -12,6 +21,162 @@ from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec LATEST_VERSION = 4 +logger = logging.getLogger(__name__) + +_MISSING_SENTINEL = object() + + +def _assemble_delta_channels( + checkpoint: "Checkpoint", + config: RunnableConfig, + checkpointer: BaseCheckpointSaver, +) -> dict[str, Any]: + """Resolve any DeltaValue entries in checkpoint channel_values to DeltaChainValue. + + Returns a dict of only the channels that needed assembly (others are untouched). + Tries get_channel_blob fast-path first; falls back to get_tuple traversal. + """ + thread_id = str(config["configurable"]["thread_id"]) + checkpoint_ns = config["configurable"].get("checkpoint_ns", "") + assembled: dict[str, Any] = {} + + for channel, value in checkpoint["channel_values"].items(): + if not isinstance(value, DeltaValue): + continue + + chain_deltas: list[list[Any]] = [] + base: list[Any] | None = None + cursor: DeltaValue = value + visited: set[str] = set() + + while True: + chain_deltas.append(cursor.delta) + prev_id = cursor.prev_checkpoint_id + if prev_id is None: + break # chain root + if prev_id in visited: + logger.warning( + "DeltaChannel chain cycle at checkpoint %r for channel %r; breaking", + prev_id, + channel, + ) + break + visited.add(prev_id) + + # Fast path: saver has a dedicated blob store. + blob = checkpointer.get_channel_blob(thread_id, checkpoint_ns, prev_id, channel) + if blob is not NotImplemented: + if isinstance(blob, DeltaValue): + cursor = blob + continue + else: + base = blob # plain list = snapshot root + break + + # Fallback: load the full checkpoint and extract channel value. + parent_config: RunnableConfig = { + "configurable": { + "thread_id": thread_id, + "checkpoint_ns": checkpoint_ns, + "checkpoint_id": prev_id, + } + } + parent_tuple = checkpointer.get_tuple(parent_config) + if parent_tuple is None: + logger.warning( + "DeltaChannel chain broken: checkpoint %r not found for channel %r", + prev_id, + channel, + ) + break + prev_val = parent_tuple.checkpoint["channel_values"].get(channel, _MISSING_SENTINEL) + if prev_val is _MISSING_SENTINEL: + break + elif isinstance(prev_val, DeltaValue): + cursor = prev_val + else: + base = prev_val + break + + chain_deltas.reverse() + assembled[channel] = DeltaChainValue(base=base, deltas=chain_deltas) + + return assembled + + +async def _aassemble_delta_channels( + checkpoint: "Checkpoint", + config: RunnableConfig, + checkpointer: BaseCheckpointSaver, +) -> dict[str, Any]: + """Async version of _assemble_delta_channels.""" + thread_id = str(config["configurable"]["thread_id"]) + checkpoint_ns = config["configurable"].get("checkpoint_ns", "") + assembled: dict[str, Any] = {} + + for channel, value in checkpoint["channel_values"].items(): + if not isinstance(value, DeltaValue): + continue + + chain_deltas: list[list[Any]] = [] + base: list[Any] | None = None + cursor: DeltaValue = value + visited: set[str] = set() + + while True: + chain_deltas.append(cursor.delta) + prev_id = cursor.prev_checkpoint_id + if prev_id is None: + break + if prev_id in visited: + logger.warning( + "DeltaChannel chain cycle at checkpoint %r for channel %r; breaking", + prev_id, + channel, + ) + break + visited.add(prev_id) + + blob = await checkpointer.aget_channel_blob( + thread_id, checkpoint_ns, prev_id, channel + ) + if blob is not NotImplemented: + if isinstance(blob, DeltaValue): + cursor = blob + continue + else: + base = blob + break + + parent_config: RunnableConfig = { + "configurable": { + "thread_id": thread_id, + "checkpoint_ns": checkpoint_ns, + "checkpoint_id": prev_id, + } + } + parent_tuple = await checkpointer.aget_tuple(parent_config) + if parent_tuple is None: + logger.warning( + "DeltaChannel chain broken: checkpoint %r not found for channel %r", + prev_id, + channel, + ) + break + prev_val = parent_tuple.checkpoint["channel_values"].get(channel, _MISSING_SENTINEL) + if prev_val is _MISSING_SENTINEL: + break + elif isinstance(prev_val, DeltaValue): + cursor = prev_val + else: + base = prev_val + break + + chain_deltas.reverse() + assembled[channel] = DeltaChainValue(base=base, deltas=chain_deltas) + + return assembled + def empty_checkpoint() -> Checkpoint: return Checkpoint(