diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py index 149254702..54844c824 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py @@ -8,12 +8,12 @@ from typing import Any from langchain_core.runnables import RunnableConfig from langgraph.checkpoint.base import ( - DELTA_SENTINEL, WRITES_IDX_MAP, ChannelVersions, Checkpoint, CheckpointMetadata, CheckpointTuple, + _ChannelWritesHistory, get_checkpoint_id, get_serializable_checkpoint_metadata, ) @@ -185,7 +185,7 @@ class PostgresSaver(BasePostgresSaver): value["channel_values"], ) for value in values: - yield self._load_checkpoint_tuple(value, cur) + yield self._load_checkpoint_tuple(value) def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None: """Get a checkpoint tuple from the database. @@ -256,7 +256,7 @@ class PostgresSaver(BasePostgresSaver): value["channel_values"], ) - return self._load_checkpoint_tuple(value, cur) + return self._load_checkpoint_tuple(value) def put( self, @@ -436,45 +436,40 @@ class PostgresSaver(BasePostgresSaver): with conn.cursor(binary=True, row_factory=dict_row) as cur: yield cur - def _reconstruct_delta_channel( - self, - *, - thread_id: str, - checkpoint_ns: str, - channel: str, - target_id: str, - cur: Cursor[DictRow], - ) -> Any: - """Run the three reconstruction SELECTs and assemble `DeltaChannelWrites`. + def _get_channel_writes_history( + self, config: RunnableConfig, channel: str + ) -> _ChannelWritesHistory: + """Fast-path override of `BaseCheckpointSaver._get_channel_writes_history`. Three indexed roundtrips (`checkpoints`, `checkpoint_writes`, `checkpoint_blobs`) each filtered by `(thread_id, checkpoint_ns)` and the per-table key. Plain SELECTs let the planner pick straight index scans; rationale + benchmark in `notes/delta_channel_query_bench.md`. """ - cur.execute(SELECT_DELTA_PARENTS_SQL, (channel, thread_id, checkpoint_ns)) - parents_rows = cur.fetchall() - cur.execute(SELECT_DELTA_WRITES_SQL, (thread_id, checkpoint_ns, channel)) - writes_rows = cur.fetchall() - cur.execute(SELECT_DELTA_BLOBS_SQL, (thread_id, checkpoint_ns, channel)) - blobs_rows = cur.fetchall() - return self._build_delta_channel_writes( - target_id=target_id, + thread_id = config["configurable"]["thread_id"] + checkpoint_ns = config["configurable"].get("checkpoint_ns", "") + checkpoint_id = config["configurable"]["checkpoint_id"] + with self._cursor() as cur: + cur.execute(SELECT_DELTA_PARENTS_SQL, (channel, thread_id, checkpoint_ns)) + parents_rows = cur.fetchall() + cur.execute(SELECT_DELTA_WRITES_SQL, (thread_id, checkpoint_ns, channel)) + writes_rows = cur.fetchall() + cur.execute(SELECT_DELTA_BLOBS_SQL, (thread_id, checkpoint_ns, channel)) + blobs_rows = cur.fetchall() + return self._build_delta_channel_writes_history( + channel=channel, + target_id=checkpoint_id, parents_rows=parents_rows, writes_rows=writes_rows, blobs_rows=blobs_rows, ) - def _load_checkpoint_tuple( - self, value: DictRow, cur: Cursor[DictRow] - ) -> CheckpointTuple: + def _load_checkpoint_tuple(self, value: DictRow) -> CheckpointTuple: """ Convert a database row into a CheckpointTuple object. Args: value: A row from the database containing checkpoint data. - cur: The cursor used by the caller; reused for DeltaChannel - reconstruction to avoid acquiring `self.lock` a second time. Returns: CheckpointTuple: A structured representation of the checkpoint, @@ -482,15 +477,6 @@ class PostgresSaver(BasePostgresSaver): and pending writes. """ channel_values = self._load_blobs(value["channel_values"]) - delta_channels = [ch for ch, v in channel_values.items() if v is DELTA_SENTINEL] - for channel in delta_channels: - channel_values[channel] = self._reconstruct_delta_channel( - thread_id=value["thread_id"], - checkpoint_ns=value["checkpoint_ns"], - channel=channel, - target_id=value["checkpoint_id"], - cur=cur, - ) return CheckpointTuple( { "configurable": { diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py index 06f809087..dc5a655e1 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py @@ -8,13 +8,12 @@ from typing import Any from langchain_core.runnables import RunnableConfig from langgraph.checkpoint.base import ( - DELTA_SENTINEL, WRITES_IDX_MAP, ChannelVersions, Checkpoint, CheckpointMetadata, CheckpointTuple, - DeltaChannelWrites, + _ChannelWritesHistory, get_checkpoint_id, get_serializable_checkpoint_metadata, ) @@ -175,7 +174,7 @@ class AsyncPostgresSaver(BasePostgresSaver): value["channel_values"], ) for value in values: - yield await self._load_checkpoint_tuple(value, cur) + yield await self._load_checkpoint_tuple(value) async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None: """Get a checkpoint tuple from the database asynchronously. @@ -226,7 +225,7 @@ class AsyncPostgresSaver(BasePostgresSaver): value["channel_values"], ) - return await self._load_checkpoint_tuple(value, cur) + return await self._load_checkpoint_tuple(value) async def aput( self, @@ -398,61 +397,46 @@ class AsyncPostgresSaver(BasePostgresSaver): async with conn.cursor(binary=True, row_factory=dict_row) as cur: yield cur - async def _areconstruct_delta_channel( - self, - *, - thread_id: str, - checkpoint_ns: str, - channel: str, - target_id: str, - cur: AsyncCursor[DictRow], - ) -> Any: - """Async mirror of `PostgresSaver._reconstruct_delta_channel`. + async def _aget_channel_writes_history( + self, config: RunnableConfig, channel: str + ) -> _ChannelWritesHistory: + """Fast-path override of `BaseCheckpointSaver._aget_channel_writes_history`. Three indexed roundtrips (`checkpoints`, `checkpoint_writes`, `checkpoint_blobs`); rows assembled by the shared pure helper on `BasePostgresSaver`. Rationale + benchmark in `notes/delta_channel_query_bench.md`. """ - await cur.execute(SELECT_DELTA_PARENTS_SQL, (channel, thread_id, checkpoint_ns)) - parents_rows = await cur.fetchall() - await cur.execute(SELECT_DELTA_WRITES_SQL, (thread_id, checkpoint_ns, channel)) - writes_rows = await cur.fetchall() - await cur.execute(SELECT_DELTA_BLOBS_SQL, (thread_id, checkpoint_ns, channel)) - blobs_rows = await cur.fetchall() - return self._build_delta_channel_writes( - target_id=target_id, + thread_id = config["configurable"]["thread_id"] + checkpoint_ns = config["configurable"].get("checkpoint_ns", "") + checkpoint_id = config["configurable"]["checkpoint_id"] + async with self._cursor() as cur: + await cur.execute( + SELECT_DELTA_PARENTS_SQL, (channel, thread_id, checkpoint_ns) + ) + parents_rows = await cur.fetchall() + await cur.execute( + SELECT_DELTA_WRITES_SQL, (thread_id, checkpoint_ns, channel) + ) + writes_rows = await cur.fetchall() + await cur.execute( + SELECT_DELTA_BLOBS_SQL, (thread_id, checkpoint_ns, channel) + ) + blobs_rows = await cur.fetchall() + return self._build_delta_channel_writes_history( + channel=channel, + target_id=checkpoint_id, parents_rows=parents_rows, writes_rows=writes_rows, blobs_rows=blobs_rows, ) - async def aget_channel_writes( - self, config: RunnableConfig, channel: str - ) -> DeltaChannelWrites: - thread_id = config["configurable"]["thread_id"] - checkpoint_ns = config["configurable"].get("checkpoint_ns", "") - checkpoint_id = config["configurable"]["checkpoint_id"] - async with self._cursor() as cur: - return await self._areconstruct_delta_channel( - thread_id=thread_id, - checkpoint_ns=checkpoint_ns, - channel=channel, - target_id=checkpoint_id, - cur=cur, - ) - - async def _load_checkpoint_tuple( - self, value: DictRow, cur: AsyncCursor[DictRow] - ) -> CheckpointTuple: + async def _load_checkpoint_tuple(self, value: DictRow) -> CheckpointTuple: """ Convert a database row into a CheckpointTuple object. Args: value: A row from the database containing checkpoint data. - cur: The cursor used by the caller; reused for DeltaChannel - reconstruction to avoid re-entering `self.lock`, which is - `asyncio.Lock` and would deadlock. Returns: CheckpointTuple: A structured representation of the checkpoint, @@ -462,21 +446,9 @@ class AsyncPostgresSaver(BasePostgresSaver): thread_id = value["thread_id"] checkpoint_ns = value["checkpoint_ns"] blob_values = value["channel_values"] - channel_values: dict[str, Any] = {} if blob_values: channel_values = self._load_blobs(blob_values) - delta_channels = [ - ch for ch, v in channel_values.items() if v is DELTA_SENTINEL - ] - for channel in delta_channels: - channel_values[channel] = await self._areconstruct_delta_channel( - thread_id=thread_id, - checkpoint_ns=checkpoint_ns, - channel=channel, - target_id=value["checkpoint_id"], - cur=cur, - ) return CheckpointTuple( { diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py index 6d727b1ba..b56c503fd 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py @@ -12,7 +12,8 @@ from langgraph.checkpoint.base import ( WRITES_IDX_MAP, BaseCheckpointSaver, ChannelVersions, - DeltaChannelWrites, + PendingWrite, + _ChannelWritesHistory, get_checkpoint_id, ) from langgraph.checkpoint.serde.types import TASKS @@ -223,15 +224,16 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): result[k.decode()] = self.serde.loads_typed((type_tag, v)) return result - def _build_delta_channel_writes( + def _build_delta_channel_writes_history( self, *, + channel: str, target_id: str, parents_rows: Sequence[Any], writes_rows: Sequence[Any], blobs_rows: Sequence[Any], - ) -> DeltaChannelWrites: - """Reconstruct one delta channel from rows of the three SELECTs. + ) -> _ChannelWritesHistory: + """Reconstruct one delta channel's history from rows of the three SELECTs. Pure data transform shared by sync (`PostgresSaver`) and async (`AsyncPostgresSaver`); both paths run the queries themselves and @@ -239,8 +241,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): Walk is newest → oldest from the target's parent. A non-sentinel blob in `checkpoint_blobs` (a pre-delta snapshot) terminates the - walk and is bound as `DeltaChannelWrites.seed` so replay starts - from it. + walk and is returned as the seed so replay starts from it. Writes stored at `target_id` itself are pending writes for the next step and are excluded — the walk begins at the target's parent. @@ -258,7 +259,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): ancestors.append(cid) cid = parent_of.get(cid) if not ancestors: - return DeltaChannelWrites(writes=[]) + return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=[]) ancestor_set = set(ancestors) # Group writes by ancestor cid; sort within (task_id DESC, idx DESC) @@ -278,9 +279,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): r["version"]: (r["type"], r["blob"]) for r in blobs_rows } - collected: list[Any] = [] # newest first; reversed at the end - seed: Any = None - found_seed = False + collected: list[PendingWrite] = [] # newest first; reversed at the end for cid in ancestors: # Pre-delta blob terminator: subsumes any writes at this ancestor. ver = ver_of.get(cid) @@ -289,17 +288,14 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): if seed_blob is not None and seed_blob[0] != "empty": blob_value = self.serde.loads_typed(seed_blob) if blob_value is not DELTA_SENTINEL: - seed = blob_value - found_seed = True - break - for type_tag, write_blob, _task_id, _idx in writes_by_cid.get(cid, []): + collected.reverse() + return _ChannelWritesHistory(seed=blob_value, writes=collected) + for type_tag, write_blob, task_id, _idx in writes_by_cid.get(cid, []): val = self.serde.loads_typed((type_tag, write_blob)) - collected.append(val) + collected.append((task_id, channel, val)) collected.reverse() # oldest → newest - if found_seed: - return DeltaChannelWrites(writes=collected, seed=seed) - return DeltaChannelWrites(writes=collected) + return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=collected) def _dump_blobs( self, diff --git a/libs/checkpoint/langgraph/checkpoint/base/__init__.py b/libs/checkpoint/langgraph/checkpoint/base/__init__.py index af65db1c7..af4ab3dc4 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/base/__init__.py @@ -29,12 +29,6 @@ from langgraph.checkpoint.serde.types import ( SCHEDULED, ChannelProtocol, ) -from langgraph.checkpoint.serde.types import ( - SEED_UNSET as SEED_UNSET, -) -from langgraph.checkpoint.serde.types import ( - DeltaChannelWrites as DeltaChannelWrites, -) V = TypeVar("V", int, float, str) PendingWrite = tuple[str, str, Any] @@ -48,26 +42,6 @@ _DELTA_RECONSTRUCTION: contextvars.ContextVar[bool] = contextvars.ContextVar( ) -def _split_list_config( - config: RunnableConfig, -) -> tuple[RunnableConfig, RunnableConfig | None]: - """Split a `get_channel_writes` config into `(list_config, before_config)`. - - Most savers collapse `list(config)` to a single row when `config` carries a - `checkpoint_id`. For ancestor traversal we need the opposite: every tuple - strictly earlier than the target. Drop `checkpoint_id` from `list_config` - and pass the original as `before=` (which filters `checkpoint_id < target`). - """ - configurable = config.get("configurable", {}) or {} - target_id = configurable.get("checkpoint_id") - if target_id is None: - return config, None - list_config: RunnableConfig = { - "configurable": {k: v for k, v in configurable.items() if k != "checkpoint_id"} - } - return list_config, config - - logger = logging.getLogger(__name__) @@ -159,6 +133,30 @@ class CheckpointTuple(NamedTuple): pending_writes: list[PendingWrite] | None = None +class _ChannelWritesHistory(NamedTuple): + """Result of `BaseCheckpointSaver._get_channel_writes_history`. + + Storage-level view of what one channel wrote across the ancestor chain + of a target checkpoint: + + * `seed` — the nearest ancestor's stored blob value for this channel, + or `DELTA_SENTINEL` if the walk reached the root without finding a + stored value. A non-sentinel seed typically indicates a pre-delta + snapshot preserved across a channel-type migration (e.g. + `BinaryOperatorAggregate` storage extended under `DeltaChannel`). + * `writes` — on-path deltas oldest→newest, one `PendingWrite` per + step that wrote to this channel. Writes stored at the target + checkpoint itself are pending for the next super-step and are + excluded. + + Experimental: method surface may change; the NamedTuple shape is the + contract. + """ + + seed: Any + writes: list[PendingWrite] + + class BaseCheckpointSaver(Generic[V]): """Base class for creating a graph checkpointer. @@ -497,44 +495,43 @@ class BaseCheckpointSaver(Generic[V]): """ raise NotImplementedError - def get_channel_writes( + def _get_channel_writes_history( self, config: RunnableConfig, channel: str - ) -> DeltaChannelWrites: - """Reconstruct a `DeltaChannel`'s write history at this checkpoint. + ) -> _ChannelWritesHistory: + """**Experimental.** Query one channel's writes along the parent chain. - Returns a `DeltaChannelWrites` carrying: + Storage-level query, not channel semantics: returns `(seed, writes)` + reflecting what storage knows about a single channel across the + ancestor chain of the target checkpoint identified by `config`. - * `writes` — per-step deltas from ancestors, oldest→newest, ready - to be replayed through the reducer in `DeltaChannel.from_checkpoint`. - * `seed` — when the ancestor walk hits a pre-delta blob (a value - stored before `DeltaChannel` was enabled for this field), replay - starts from that snapshot instead of the channel's empty value. - Default `SEED_UNSET` means no seed. + * `writes` — on-path deltas oldest→newest as `PendingWrite` tuples. + Writes stored at the target `checkpoint_id` itself are pending + for the next super-step and are excluded. + * `seed` — the nearest ancestor's stored blob value for this + channel; `DELTA_SENTINEL` if the walk reached the root without + finding a stored value. A non-sentinel seed typically indicates + a pre-delta snapshot preserved across a channel-type migration. - Walks the **parent chain** (not `list(before=...)`): for a thread with - forks, only on-path ancestors contribute. Writes are returned - oldest→newest. + Walks the **parent chain** (not `list(before=...)`): for forked + threads, only on-path ancestors contribute. - Writes stored at the target `checkpoint_id` itself are pending writes - for the next step and are excluded — pregel applies them separately - via `apply_writes`. + Reference implementation walks `get_tuple` + `parent_config`, + inspecting each ancestor's `channel_values[channel]` for the seed + terminator. Savers with direct storage access (`InMemorySaver`, + `PostgresSaver`) override for performance; the return contract is + fixed here. - The base implementation uses `get_tuple` and `pending_writes`; it - never sees blobs, so it never sets `seed`. Savers that can read the - blob table directly (`InMemorySaver`, `PostgresSaver`) override this - method to set `seed` when appropriate, which both shortens the walk - and recovers state from pre-delta threads after migration. + Underscore-prefixed because the method surface is experimental. """ # Guard against re-entrant calls: when get_tuple() triggers # reconstruction which calls get_tuple() again, the inner call - # returns tuples with DELTA_SENTINEL in channel_values (which this - # method ignores — it only reads pending_writes). + # short-circuits here. if _DELTA_RECONSTRUCTION.get(): - return DeltaChannelWrites(writes=[]) + return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=[]) token = _DELTA_RECONSTRUCTION.set(True) try: - collected: list[Any] = [] # newest first + collected: list[PendingWrite] = [] # newest first; reversed at the end target_tuple = self.get_tuple(config) cursor_config: RunnableConfig | None = ( target_tuple.parent_config if target_tuple else None @@ -543,29 +540,36 @@ class BaseCheckpointSaver(Generic[V]): tup = self.get_tuple(cursor_config) if tup is None: break + # Pre-delta seed terminator: if the ancestor has a stored + # (non-sentinel) value for this channel, that snapshot + # subsumes any earlier writes on the chain. Stop here. + ancestor_value = tup.checkpoint["channel_values"].get(channel) + if ancestor_value is not None and ancestor_value is not DELTA_SENTINEL: + collected.reverse() + return _ChannelWritesHistory(seed=ancestor_value, writes=collected) if tup.pending_writes: # Within a superstep, pending_writes are oldest→newest; # reverse to scan newest-first. - for _, ch, value in reversed(tup.pending_writes): - if ch != channel: + for write in reversed(tup.pending_writes): + if write[1] != channel: continue - collected.append(value) + collected.append(write) cursor_config = tup.parent_config collected.reverse() - return DeltaChannelWrites(writes=collected) + return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=collected) finally: _DELTA_RECONSTRUCTION.reset(token) - async def aget_channel_writes( + async def _aget_channel_writes_history( self, config: RunnableConfig, channel: str - ) -> DeltaChannelWrites: - """Async version of `get_channel_writes`. See docstring there.""" + ) -> _ChannelWritesHistory: + """Async version of `_get_channel_writes_history`. See docstring there.""" if _DELTA_RECONSTRUCTION.get(): - return DeltaChannelWrites(writes=[]) + return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=[]) token = _DELTA_RECONSTRUCTION.set(True) try: - collected: list[Any] = [] + collected: list[PendingWrite] = [] target_tuple = await self.aget_tuple(config) cursor_config: RunnableConfig | None = ( target_tuple.parent_config if target_tuple else None @@ -574,14 +578,18 @@ class BaseCheckpointSaver(Generic[V]): tup = await self.aget_tuple(cursor_config) if tup is None: break + ancestor_value = tup.checkpoint["channel_values"].get(channel) + if ancestor_value is not None and ancestor_value is not DELTA_SENTINEL: + collected.reverse() + return _ChannelWritesHistory(seed=ancestor_value, writes=collected) if tup.pending_writes: - for _, ch, value in reversed(tup.pending_writes): - if ch != channel: + for write in reversed(tup.pending_writes): + if write[1] != channel: continue - collected.append(value) + collected.append(write) cursor_config = tup.parent_config collected.reverse() - return DeltaChannelWrites(writes=collected) + return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=collected) finally: _DELTA_RECONSTRUCTION.reset(token) diff --git a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py index b02d16af1..39bb133dc 100644 --- a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py @@ -21,8 +21,9 @@ from langgraph.checkpoint.base import ( Checkpoint, CheckpointMetadata, CheckpointTuple, - DeltaChannelWrites, + PendingWrite, SerializerProtocol, + _ChannelWritesHistory, get_checkpoint_id, get_checkpoint_metadata, ) @@ -139,21 +140,9 @@ class InMemorySaver( result[k] = self.serde.loads_typed(vv) return result - def _resolve_delta_channels( - self, - config: RunnableConfig, - channel_values: dict[str, Any], - ) -> None: - """Replace DELTA_SENTINEL entries with DeltaChannelWrites so - `DeltaChannel.from_checkpoint` can distinguish reconstructed writes - from a pre-DeltaChannel accumulated value.""" - for channel, value in channel_values.items(): - if value is DELTA_SENTINEL: - channel_values[channel] = self.get_channel_writes(config, channel) - - def get_channel_writes( + def _get_channel_writes_history( self, config: RunnableConfig, channel: str - ) -> DeltaChannelWrites: + ) -> _ChannelWritesHistory: thread_id = config["configurable"]["thread_id"] checkpoint_ns = config["configurable"].get("checkpoint_ns", "") checkpoint_id = config["configurable"].get("checkpoint_id", "") @@ -183,7 +172,7 @@ class InMemorySaver( # which already subsumes any writes stored under it. Processing # those writes first would fold them into the reconstructed value # twice (once via the blob, once via replay). - collected: list[Any] = [] # newest first + collected: list[PendingWrite] = [] # newest first for cp_id in chain: # newest → oldest entry = ns_storage.get(cp_id) if entry is not None: @@ -199,25 +188,27 @@ class InMemorySaver( # Pre-delta snapshot terminator. Skip this # ancestor's writes — the blob subsumes them. collected.reverse() - return DeltaChannelWrites(writes=collected, seed=blob_value) + return _ChannelWritesHistory( + seed=blob_value, writes=collected + ) step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {}) # Within a superstep, sorted by (task_id, idx) = oldest → newest; # reverse for newest-first scan. - for (_task_id, _idx), (_, ch, serialized, _) in sorted( + for (_task_id, _idx), (tid, ch, serialized, _) in sorted( step_writes.items(), reverse=True ): if ch != channel: continue val = self.serde.loads_typed(serialized) - collected.append(val) + collected.append((tid, ch, val)) collected.reverse() - return DeltaChannelWrites(writes=collected) + return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=collected) - async def aget_channel_writes( + async def _aget_channel_writes_history( self, config: RunnableConfig, channel: str - ) -> DeltaChannelWrites: - return self.get_channel_writes(config, channel) + ) -> _ChannelWritesHistory: + return self._get_channel_writes_history(config, channel) def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None: """Get a checkpoint tuple from the in-memory storage. @@ -245,7 +236,6 @@ class InMemorySaver( checkpoint_ns, checkpoint_["channel_versions"], ) - self._resolve_delta_channels(config, channel_values) return CheckpointTuple( config=config, checkpoint={ @@ -289,7 +279,6 @@ class InMemorySaver( checkpoint_ns, checkpoint_["channel_versions"], ) - self._resolve_delta_channels(resolved_config, channel_values) return CheckpointTuple( config=resolved_config, checkpoint={ @@ -404,7 +393,6 @@ class InMemorySaver( checkpoint_ns, checkpoint_["channel_versions"], ) - self._resolve_delta_channels(list_config, channel_values) yield CheckpointTuple( config=list_config, diff --git a/libs/checkpoint/langgraph/checkpoint/serde/types.py b/libs/checkpoint/langgraph/checkpoint/serde/types.py index 29da8f67b..ea4cd8180 100644 --- a/libs/checkpoint/langgraph/checkpoint/serde/types.py +++ b/libs/checkpoint/langgraph/checkpoint/serde/types.py @@ -1,4 +1,3 @@ -import dataclasses from collections.abc import Sequence from typing import ( Any, @@ -34,39 +33,6 @@ class _DeltaSentinel: DELTA_SENTINEL = _DeltaSentinel() -class _SeedUnset: - """Marker used as the default for `DeltaChannelWrites.seed`. - - Distinct from `None`, which is a legitimate pre-delta value - (e.g. an `Optional` field whose accumulated value really was `None`). - """ - - __slots__ = () - - def __repr__(self) -> str: - return "SEED_UNSET" - - -SEED_UNSET = _SeedUnset() - - -@dataclasses.dataclass -class DeltaChannelWrites: - """In-memory wrapper around per-step writes reconstructed by a saver. - Consumed by `DeltaChannel.from_checkpoint`. Never serialized — if this - reaches the wire, something upstream forgot to unwrap it. - - `seed` is the value from which chain replay should begin. When the saver - encounters a pre-delta blob during the ancestor walk, it uses that blob - as the seed and stops walking further back (the older chain is - represented by the seed). `SEED_UNSET` means "no seed — replay from the - channel's empty value". - """ - - writes: list[Any] - seed: Any = SEED_UNSET - - Value = TypeVar("Value", covariant=True) Update = TypeVar("Update", contravariant=True) C = TypeVar("C") diff --git a/libs/checkpoint/tests/test_memory.py b/libs/checkpoint/tests/test_memory.py index e32fbb5b6..617b23f68 100644 --- a/libs/checkpoint/tests/test_memory.py +++ b/libs/checkpoint/tests/test_memory.py @@ -7,10 +7,8 @@ from pydantic import BaseModel from langgraph.checkpoint.base import ( DELTA_SENTINEL, - SEED_UNSET, Checkpoint, CheckpointMetadata, - DeltaChannelWrites, create_checkpoint, empty_checkpoint, ) @@ -346,8 +344,8 @@ class TestInMemorySaverDeltaChannel: assert result[channel] is DELTA_SENTINEL def test_get_channel_writes_collects_ancestor_writes_only(self) -> None: - """get_channel_writes collects ancestor writes oldest→newest, and - excludes writes stored at the target checkpoint itself (those are + """_get_channel_writes_history collects ancestor writes oldest→newest, + and excludes writes stored at the target checkpoint itself (those are pending writes for the next step, applied separately by pregel).""" saver = InMemorySaver() serde = JsonPlusSerializer() @@ -386,9 +384,10 @@ class TestInMemorySaverDeltaChannel: "checkpoint_id": "cp2", } } - result = saver.get_channel_writes(config, channel) - assert result == DeltaChannelWrites(writes=[{"content": "hi"}]) - assert result.seed is SEED_UNSET + result = saver._get_channel_writes_history(config, channel) + assert result.seed is DELTA_SENTINEL + values = [v for _, _, v in result.writes] + assert values == [{"content": "hi"}] def test_get_channel_writes_at_root_returns_empty(self) -> None: """Reconstructing the root checkpoint's state: no ancestors → [].""" @@ -415,15 +414,15 @@ class TestInMemorySaverDeltaChannel: "checkpoint_id": "cp1", } } - assert saver.get_channel_writes(config, channel) == DeltaChannelWrites( - writes=[] - ) + result = saver._get_channel_writes_history(config, channel) + assert result.seed is DELTA_SENTINEL + assert result.writes == [] class TestBaseFallbackGetChannelWrites: - """Exercises the `BaseCheckpointSaver.get_channel_writes` default + """Exercises the `BaseCheckpointSaver._get_channel_writes_history` default implementation — the path third-party savers inherit when they don't - override `get_channel_writes` themselves. + override `_get_channel_writes_history` themselves. Regression guard for a bug where the fallback passed the caller's config (with `checkpoint_id`) straight to `self.list()`, which most savers @@ -439,11 +438,11 @@ class TestBaseFallbackGetChannelWrites: """ class _ThirdPartyStyleSaver(InMemorySaver): - get_channel_writes = ( - InMemorySaver.__mro__[1].get_channel_writes # type: ignore[attr-defined] + _get_channel_writes_history = ( + InMemorySaver.__mro__[1]._get_channel_writes_history # type: ignore[attr-defined] ) - aget_channel_writes = ( - InMemorySaver.__mro__[1].aget_channel_writes # type: ignore[attr-defined] + _aget_channel_writes_history = ( + InMemorySaver.__mro__[1]._aget_channel_writes_history # type: ignore[attr-defined] ) saver = _ThirdPartyStyleSaver() @@ -487,12 +486,11 @@ class TestBaseFallbackGetChannelWrites: } } - result = saver.get_channel_writes(config, "messages") + result = saver._get_channel_writes_history(config, "messages") - assert result == DeltaChannelWrites( - writes=[{"content": "first"}, {"content": "second"}] - ) - assert result.seed is SEED_UNSET + assert result.seed is DELTA_SENTINEL + values = [v for _, _, v in result.writes] + assert values == [{"content": "first"}, {"content": "second"}] async def test_async_fallback_returns_ancestor_writes_oldest_first(self) -> None: saver, thread_id, ns = self._build_saver_with_chain() @@ -505,18 +503,17 @@ class TestBaseFallbackGetChannelWrites: } } - result = await saver.aget_channel_writes(config, "messages") + result = await saver._aget_channel_writes_history(config, "messages") - assert result == DeltaChannelWrites( - writes=[{"content": "first"}, {"content": "second"}] - ) - assert result.seed is SEED_UNSET + assert result.seed is DELTA_SENTINEL + values = [v for _, _, v in result.writes] + assert values == [{"content": "first"}, {"content": "second"}] async def test_async_fallback_concurrent_tasks_do_not_interfere(self) -> None: """Regression: the re-entrancy guard must be task-local, not thread-local. - Two concurrent `aget_channel_writes` calls on the same event-loop - thread must each see their full reconstructed writes. A + Two concurrent `_aget_channel_writes_history` calls on the same + event-loop thread must each see their full reconstructed writes. A `threading.local()` guard would let whichever task set it first short-circuit the other to `writes=[]`. """ @@ -545,15 +542,16 @@ class TestBaseFallbackGetChannelWrites: } results = await asyncio.gather( - saver.aget_channel_writes(config, "messages"), - saver.aget_channel_writes(config, "messages"), + saver._aget_channel_writes_history(config, "messages"), + saver._aget_channel_writes_history(config, "messages"), ) - expected = DeltaChannelWrites( - writes=[{"content": "first"}, {"content": "second"}] - ) - assert results[0] == expected - assert results[1] == expected + expected_values = [{"content": "first"}, {"content": "second"}] + for result in results: + assert result.seed is DELTA_SENTINEL + values = [v for _, _, v in result.writes] + assert values == expected_values + class TestPreDeltaBlobTerminator: """Verify the pre-delta blob terminator: when the ancestor walk hits a @@ -636,14 +634,15 @@ class TestPreDeltaBlobTerminator: } } - result = saver.get_channel_writes(config, channel) + result = saver._get_channel_writes_history(config, channel) # Seed came from the pre-delta blob at cp1. assert result.seed == ["A"] # Delta-era writes from cp2 replay through the reducer on top of seed. # cp3 is the target — its own write is pending for the NEXT step and # must be excluded. - assert result.writes == ["B"] + values = [v for _, _, v in result.writes] + assert values == ["B"] def test_pre_delta_blob_terminates_walk_before_older_writes(self) -> None: """Writes stored at the pre-delta ancestor itself must not be replayed @@ -657,9 +656,10 @@ class TestPreDeltaBlobTerminator: } } - result = saver.get_channel_writes(config, channel) + result = saver._get_channel_writes_history(config, channel) + values = [v for _, _, v in result.writes] # The pre-delta write under cp1 must not appear (the blob subsumes it). - assert "PRE-DELTA-WRITE" not in result.writes + assert "PRE-DELTA-WRITE" not in values # And the pending write at the target is never folded in. - assert "PENDING-AT-TARGET" not in result.writes + assert "PENDING-AT-TARGET" not in values diff --git a/libs/langgraph/langgraph/channels/delta.py b/libs/langgraph/langgraph/channels/delta.py index 28e223a61..e8dd955c2 100644 --- a/libs/langgraph/langgraph/channels/delta.py +++ b/libs/langgraph/langgraph/channels/delta.py @@ -4,7 +4,7 @@ import copy as _copy from collections.abc import Callable, Sequence from typing import Any, Generic -from langgraph.checkpoint.base import DELTA_SENTINEL, SEED_UNSET, DeltaChannelWrites +from langgraph.checkpoint.base import DELTA_SENTINEL, PendingWrite from typing_extensions import Self from langgraph._internal._typing import MISSING @@ -96,28 +96,41 @@ class DeltaChannel(Generic[Value], BaseChannel[Any, Any, Any]): return self.operator(base, write) def from_checkpoint(self, checkpoint: Any) -> Self: + """Initialize from a seed value. + + Pregel's hydration path calls this with the `seed` returned by + `saver.get_channel_history`: + + * `MISSING` / `DELTA_SENTINEL` → channel starts empty. The walk + either reached the root (fresh delta thread) or found nothing + to seed from. + * any other value → use as the base value. Typically a pre-delta + blob preserved across a channel-type migration; `replay_writes` + folds subsequent deltas on top. + """ new: DeltaChannel[Value] = DeltaChannel(self.operator) new.typ = self.typ new.key = self.key - if checkpoint is MISSING: + if checkpoint is MISSING or checkpoint is DELTA_SENTINEL: new.value = _empty(new.typ) - elif isinstance(checkpoint, DeltaChannelWrites): - # Saver reconstructed per-step writes; replay through the operator. - # `seed` (if set) is a pre-delta accumulated value that terminates - # the ancestor walk on the saver side: replay starts from it - # instead of the channel's empty value. - value: Any = ( - _empty(new.typ) if checkpoint.seed is SEED_UNSET else checkpoint.seed - ) - for write in checkpoint.writes: - value = new._apply_write(value, write) - new.value = value else: - # Backward compat: a pre-DeltaChannel thread stored the accumulated - # value directly (no saver-side reconstruction happened). Trust it. new.value = checkpoint return new + def replay_writes(self, writes: Sequence[PendingWrite]) -> None: + """Fold a sequence of `PendingWrite` tuples into the current value. + + Called after `from_checkpoint` during pregel hydration to replay + per-step deltas from on-path ancestors through the reducer. Writes + are oldest→newest. `Overwrite` values inside the stream reset the + reducer state at that point, same as during a live super-step. + The `task_id` and `channel` fields of each `PendingWrite` are + ignored — `_get_channel_writes_history` has already filtered to + this channel. + """ + for _, _, value in writes: + self.value = self._apply_write(self.value, value) + def update(self, values: Sequence[Any]) -> bool: if not values: return False diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index 3d510ac7d..f3c7f1102 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -3,11 +3,13 @@ from __future__ import annotations from collections.abc import Mapping from datetime import datetime, timezone -from langgraph.checkpoint.base import Checkpoint +from langchain_core.runnables import RunnableConfig +from langgraph.checkpoint.base import DELTA_SENTINEL, BaseCheckpointSaver, Checkpoint from langgraph.checkpoint.base.id import uuid6 from langgraph._internal._typing import MISSING from langgraph.channels.base import BaseChannel +from langgraph.channels.delta import DeltaChannel from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec LATEST_VERSION = 4 @@ -58,8 +60,22 @@ def create_checkpoint( def channels_from_checkpoint( specs: Mapping[str, BaseChannel | ManagedValueSpec], checkpoint: Checkpoint, + *, + saver: BaseCheckpointSaver | None = None, + config: RunnableConfig | None = None, ) -> tuple[Mapping[str, BaseChannel], ManagedValueMapping]: - """Get channels from a checkpoint.""" + """Hydrate channels from a checkpoint. + + For most channels, `spec.from_checkpoint(checkpoint["channel_values"][k])` + is sufficient — the stored value IS the reconstructed state. + + `DeltaChannel` is the exception: its stored value is a sentinel; the + full state is spread across `checkpoint_writes` along the ancestor + chain. When `saver` and `config` are provided, this function fetches + that history via `saver._get_channel_writes_history` and folds it + through the channel's reducer. Without them (static contexts — graph + drawing, unit tests), delta channels fall back to empty. + """ channel_specs: dict[str, BaseChannel] = {} managed_specs: dict[str, ManagedValueSpec] = {} for k, v in specs.items(): @@ -67,13 +83,67 @@ def channels_from_checkpoint( channel_specs[k] = v else: managed_specs[k] = v - return ( - { - k: v.from_checkpoint(checkpoint["channel_values"].get(k, MISSING)) - for k, v in channel_specs.items() - }, - managed_specs, - ) + + channels: dict[str, BaseChannel] = {} + for k, spec in channel_specs.items(): + ch: BaseChannel + stored = checkpoint["channel_values"].get(k, MISSING) + if ( + isinstance(spec, DeltaChannel) + and saver is not None + and config is not None + and (stored is MISSING or stored is DELTA_SENTINEL) + ): + # Target's own blob is empty/sentinel — walk ancestors for + # seed + writes. Skipping this when `stored` is a real value + # preserves state written via `update_state` or sitting at the + # tip of a pre-migration thread: the saver's ancestor walk + # intentionally excludes the target's own blob, so without + # this short-circuit we'd lose it. + history = saver._get_channel_writes_history(config, k) + delta_ch = spec.from_checkpoint(history.seed) + delta_ch.replay_writes(history.writes) + ch = delta_ch + else: + ch = spec.from_checkpoint(stored) + channels[k] = ch + return channels, managed_specs + + +async def achannels_from_checkpoint( + specs: Mapping[str, BaseChannel | ManagedValueSpec], + checkpoint: Checkpoint, + *, + saver: BaseCheckpointSaver | None = None, + config: RunnableConfig | None = None, +) -> tuple[Mapping[str, BaseChannel], ManagedValueMapping]: + """Async version of `channels_from_checkpoint`. See docstring there.""" + channel_specs: dict[str, BaseChannel] = {} + managed_specs: dict[str, ManagedValueSpec] = {} + for k, v in specs.items(): + if isinstance(v, BaseChannel): + channel_specs[k] = v + else: + managed_specs[k] = v + + channels: dict[str, BaseChannel] = {} + for k, spec in channel_specs.items(): + ch: BaseChannel + stored = checkpoint["channel_values"].get(k, MISSING) + if ( + isinstance(spec, DeltaChannel) + and saver is not None + and config is not None + and (stored is MISSING or stored is DELTA_SENTINEL) + ): + history = await saver._aget_channel_writes_history(config, k) + delta_ch = spec.from_checkpoint(history.seed) + delta_ch.replay_writes(history.writes) + ch = delta_ch + else: + ch = spec.from_checkpoint(stored) + channels[k] = ch + return channels, managed_specs def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint: diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 976c1382b..b2ce1d5f9 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -92,6 +92,7 @@ from langgraph.pregel._algo import ( task_path_str, ) from langgraph.pregel._checkpoint import ( + achannels_from_checkpoint, channels_from_checkpoint, copy_checkpoint, create_checkpoint, @@ -1273,7 +1274,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): ) self.submit = self.stack.enter_context(BackgroundExecutor(self.config)) self.channels, self.managed = channels_from_checkpoint( - self.specs, self.checkpoint + self.specs, + self.checkpoint, + saver=self.checkpointer, + config=self.checkpoint_config, ) self.stack.push(self._suppress_interrupt) self.status = "input" @@ -1476,8 +1480,11 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): self.submit = await self.stack.enter_async_context( AsyncBackgroundExecutor(self.config) ) - self.channels, self.managed = channels_from_checkpoint( - self.specs, self.checkpoint + self.channels, self.managed = await achannels_from_checkpoint( + self.specs, + self.checkpoint, + saver=self.checkpointer, + config=self.checkpoint_config, ) self.stack.push(self._suppress_interrupt) self.status = "input" diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index c440e77f9..6d0455055 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -122,6 +122,7 @@ from langgraph.pregel._algo import ( ) from langgraph.pregel._call import identifier from langgraph.pregel._checkpoint import ( + achannels_from_checkpoint, channels_from_checkpoint, copy_checkpoint, create_checkpoint, @@ -1052,6 +1053,10 @@ class Pregel( channels, managed = channels_from_checkpoint( self.channels, saved.checkpoint, + saver=self.checkpointer + if isinstance(self.checkpointer, BaseCheckpointSaver) + else None, + config=saved.config, ) # tasks for this checkpoint next_tasks = prepare_next_tasks( @@ -1168,9 +1173,13 @@ class Pregel( step = saved.metadata.get("step", -1) + 1 stop = step + 2 - channels, managed = channels_from_checkpoint( + channels, managed = await achannels_from_checkpoint( self.channels, saved.checkpoint, + saver=self.checkpointer + if isinstance(self.checkpointer, BaseCheckpointSaver) + else None, + config=saved.config, ) # tasks for this checkpoint next_tasks = prepare_next_tasks( @@ -1541,6 +1550,11 @@ class Pregel( channels, managed = channels_from_checkpoint( self.channels, checkpoint, + saver=self.checkpointer + if saved is not None + and isinstance(self.checkpointer, BaseCheckpointSaver) + else None, + config=saved.config if saved is not None else None, ) values, as_node = updates[0][:2] @@ -1984,9 +1998,14 @@ class Pregel( ) if saved: checkpoint_config = patch_configurable(config, saved.config[CONF]) - channels, managed = channels_from_checkpoint( + channels, managed = await achannels_from_checkpoint( self.channels, checkpoint, + saver=self.checkpointer + if saved is not None + and isinstance(self.checkpointer, BaseCheckpointSaver) + else None, + config=saved.config if saved is not None else None, ) values, as_node = updates[0][:2] # no values, just clear all tasks diff --git a/libs/langgraph/tests/test_channels.py b/libs/langgraph/tests/test_channels.py index de0cea7a6..4fd3bb9d3 100644 --- a/libs/langgraph/tests/test_channels.py +++ b/libs/langgraph/tests/test_channels.py @@ -3,7 +3,7 @@ from collections.abc import Sequence import pytest from langchain_core.messages import AIMessage, HumanMessage -from langgraph.checkpoint.base import SEED_UNSET, DeltaChannelWrites +from langgraph.checkpoint.base import DELTA_SENTINEL from langgraph._internal._typing import MISSING from langgraph.channels.binop import BinaryOperatorAggregate @@ -149,22 +149,21 @@ def test_delta_channel_basic_two_steps() -> None: def test_delta_channel_from_checkpoint_writes_list() -> None: - """from_checkpoint given DeltaChannelWrites replays through the operator.""" + """replay_writes on a fresh channel replays through the operator.""" from langchain_core.messages import AIMessage, HumanMessage - from langgraph.checkpoint.base import DeltaChannelWrites from langgraph.channels.delta import DeltaChannel from langgraph.graph.message import add_messages spec = DeltaChannel(add_messages) - writes = DeltaChannelWrites( + ch = spec.from_checkpoint(DELTA_SENTINEL) + ch.replay_writes( [ - HumanMessage(content="hi", id="h1"), - AIMessage(content="hello", id="a1"), - HumanMessage(content="bye", id="h2"), + ("t0", "messages", HumanMessage(content="hi", id="h1")), + ("t1", "messages", AIMessage(content="hello", id="a1")), + ("t2", "messages", HumanMessage(content="bye", id="h2")), ] ) - ch = spec.from_checkpoint(writes) msgs = ch.get() assert len(msgs) == 3 assert msgs[0].content == "hi" @@ -227,16 +226,14 @@ def test_delta_channel_remove_message_and_replay() -> None: assert ch.get() == [HumanMessage(content="hi", id="h1")] # Replay the writes list from scratch — must reproduce the post-remove state - from langgraph.checkpoint.base import DeltaChannelWrites - - writes = DeltaChannelWrites( + ch2 = spec.from_checkpoint(DELTA_SENTINEL) + ch2.replay_writes( [ - HumanMessage(content="hi", id="h1"), - AIMessage(content="hello", id="a1"), - RemoveMessage(id="a1"), + ("t0", "messages", HumanMessage(content="hi", id="h1")), + ("t1", "messages", AIMessage(content="hello", id="a1")), + ("t2", "messages", RemoveMessage(id="a1")), ] ) - ch2 = spec.from_checkpoint(writes) assert ch2.get() == [HumanMessage(content="hi", id="h1")] @@ -258,15 +255,13 @@ def test_delta_channel_update_by_id_and_replay() -> None: assert ch.get() == [HumanMessage(content="updated", id="h1")] # Replay writes — must produce the updated message, not the original - from langgraph.checkpoint.base import DeltaChannelWrites - - writes = DeltaChannelWrites( + ch2 = spec.from_checkpoint(DELTA_SENTINEL) + ch2.replay_writes( [ - HumanMessage(content="original", id="h1"), - HumanMessage(content="updated", id="h1"), + ("t0", "messages", HumanMessage(content="original", id="h1")), + ("t1", "messages", HumanMessage(content="updated", id="h1")), ] ) - ch2 = spec.from_checkpoint(writes) assert len(ch2.get()) == 1 assert ch2.get()[0].content == "updated" @@ -318,16 +313,13 @@ def test_delta_channel_inmemory_saver_assembles_writes() -> None: graph.invoke({"messages": [HumanMessage(content="hi", id="h1")]}, config) graph.invoke({"messages": [HumanMessage(content="bye", id="h2")]}, config) - # get_tuple must return a DeltaChannelWrites wrapper (not the raw sentinel) - from langgraph.checkpoint.base import DELTA_SENTINEL, DeltaChannelWrites - + # get_tuple returns raw storage shape — channel_values stores DELTA_SENTINEL + # for delta channels; the reconstructed writes flow separately via + # saver._get_channel_writes_history. saved = saver.get_tuple(config) assert saved is not None assert "messages" in saved.checkpoint["channel_values"] - assert saved.checkpoint["channel_values"]["messages"] is not DELTA_SENTINEL - assert isinstance( - saved.checkpoint["channel_values"]["messages"], DeltaChannelWrites - ) + assert saved.checkpoint["channel_values"]["messages"] is DELTA_SENTINEL state = graph.get_state(config) assert len(state.values["messages"]) == 4 # 2 human + 2 AI @@ -376,15 +368,20 @@ def test_delta_channel_dict_reducer_basic_updates() -> None: def test_delta_channel_dict_reducer_writes_reconstruction() -> None: - """from_checkpoint given DeltaChannelWrites replays through a dict merge reducer.""" - from langgraph.checkpoint.base import DeltaChannelWrites + """replay_writes on a fresh channel replays through a dict merge reducer.""" def merge_dicts(left: dict, right: dict) -> dict: return {**left, **right} spec = _delta_channel_with_type(merge_dicts, dict) - writes = DeltaChannelWrites([{"a": 1}, {"b": 2}, {"c": 3}]) - ch = spec.from_checkpoint(writes) + ch = spec.from_checkpoint(DELTA_SENTINEL) + ch.replay_writes( + [ + ("t0", "files", {"a": 1}), + ("t1", "files", {"b": 2}), + ("t2", "files", {"c": 3}), + ] + ) assert ch.get() == {"a": 1, "b": 2, "c": 3} @@ -412,16 +409,14 @@ def test_delta_channel_dict_reducer_with_deletions() -> None: assert ch.get() == {"file2.py": "content2", "file3.py": "content3"} # Confirm writes reconstruction produces the same result - from langgraph.checkpoint.base import DeltaChannelWrites - - writes = DeltaChannelWrites( + spec = _delta_channel_with_type(merge_files, dict) + ch2 = spec.from_checkpoint(DELTA_SENTINEL) + ch2.replay_writes( [ - {"file1.py": "content1", "file2.py": "content2"}, - {"file1.py": None, "file3.py": "content3"}, + ("t0", "files", {"file1.py": "content1", "file2.py": "content2"}), + ("t1", "files", {"file1.py": None, "file3.py": "content3"}), ] ) - spec = _delta_channel_with_type(merge_files, dict) - ch2 = spec.from_checkpoint(writes) assert ch2.get() == {"file2.py": "content2", "file3.py": "content3"} @@ -440,23 +435,21 @@ def test_delta_channel_dict_reducer_overwrite_in_update() -> None: def test_delta_channel_dict_reducer_overwrite_in_writes_replay() -> None: - """Overwrite(dict) embedded in DeltaChannelWrites must reconstruct as dict.""" - from langgraph.checkpoint.base import DeltaChannelWrites - + """Overwrite(dict) embedded in replayed writes must reconstruct as dict.""" from langgraph.types import Overwrite def merge_dicts(left: dict, right: dict) -> dict: return {**left, **right} spec = _delta_channel_with_type(merge_dicts, dict) - writes = DeltaChannelWrites( + ch = spec.from_checkpoint(DELTA_SENTINEL) + ch.replay_writes( [ - {"a": 1}, - Overwrite({"x": 10, "y": 20}), - {"z": 30}, + ("t0", "files", {"a": 1}), + ("t1", "files", Overwrite({"x": 10, "y": 20})), + ("t2", "files", {"z": 30}), ] ) - ch = spec.from_checkpoint(writes) assert ch.get() == {"x": 10, "y": 20, "z": 30} @@ -498,7 +491,6 @@ def test_delta_channel_dict_reducer_end_to_end_filesystem() -> None: """ from typing import Annotated - from langgraph.checkpoint.base import DELTA_SENTINEL, DeltaChannelWrites from langgraph.checkpoint.memory import InMemorySaver from typing_extensions import TypedDict @@ -540,8 +532,7 @@ def test_delta_channel_dict_reducer_end_to_end_filesystem() -> None: saved = saver.get_tuple(config) assert saved is not None cv = saved.checkpoint["channel_values"]["files"] - assert cv is not DELTA_SENTINEL - assert isinstance(cv, DeltaChannelWrites) + assert cv is DELTA_SENTINEL state = graph.get_state(config) assert state.values["files"] == { @@ -586,7 +577,7 @@ def test_delta_channel_dict_reducer_backwards_compat() -> None: def test_delta_channel_from_checkpoint_honors_seed() -> None: - """DeltaChannelWrites(seed=...) starts replay from that snapshot. + """A non-sentinel value to from_checkpoint is used as the pre-delta seed. Guards the pre-delta migration path: when the saver's ancestor walk hits a pre-DeltaChannel blob it passes it as `seed` so replay reconstructs @@ -594,14 +585,13 @@ def test_delta_channel_from_checkpoint_honors_seed() -> None: """ spec = DeltaChannel(add_messages) seed = [HumanMessage(content="pre-delta", id="p1")] - writes = DeltaChannelWrites( - writes=[ - AIMessage(content="delta-1", id="d1"), - HumanMessage(content="delta-2", id="d2"), - ], - seed=seed, + ch = spec.from_checkpoint(seed) + ch.replay_writes( + [ + ("t0", "messages", AIMessage(content="delta-1", id="d1")), + ("t1", "messages", HumanMessage(content="delta-2", id="d2")), + ] ) - ch = spec.from_checkpoint(writes) msgs = ch.get() assert [m.content for m in msgs] == ["pre-delta", "delta-1", "delta-2"] @@ -611,23 +601,23 @@ def test_delta_channel_from_checkpoint_seed_without_writes() -> None: just the seed — the saver's terminator fired immediately.""" spec = DeltaChannel(add_messages) seed = [HumanMessage(content="only-snap", id="s1")] - ch = spec.from_checkpoint(DeltaChannelWrites(writes=[], seed=seed)) + ch = spec.from_checkpoint(seed) + ch.replay_writes([]) assert ch.get() == seed -def test_delta_channel_from_checkpoint_seed_none_is_distinct_from_unset() -> None: - """`seed=None` must start replay from None, not from the channel's empty - value. `SEED_UNSET` is the sentinel meaning 'no seed'.""" +def test_delta_channel_from_checkpoint_seed_none_is_distinct_from_sentinel() -> None: + """`seed=None` must start replay from None, not from an empty channel. + + The DELTA_SENTINEL / MISSING sentinels mean 'no seed'; passing `None` + explicitly should feed None to the reducer as the left operand. + """ def replace(left, right): return right spec = DeltaChannel(replace) - ch = spec.from_checkpoint(DeltaChannelWrites(writes=["after"], seed=None)) + ch = spec.from_checkpoint(None) + ch.replay_writes([("t0", "x", "after")]) # Reducer replaces; seed=None → first write produces "after". assert ch.get() == "after" - # And the default (unset) is distinct. - unset = DeltaChannelWrites(writes=["after"]) - assert unset.seed is SEED_UNSET - - diff --git a/libs/langgraph/tests/test_delta_channel_benchmark.py b/libs/langgraph/tests/test_delta_channel_benchmark.py index bc8b9d441..ab710ebcd 100644 --- a/libs/langgraph/tests/test_delta_channel_benchmark.py +++ b/libs/langgraph/tests/test_delta_channel_benchmark.py @@ -331,9 +331,7 @@ def _run_benchmark_for_checkpointer(cp_hint: Any) -> None: # ── Table 2: Read latency ───────────────────────────────────────────────── print("Read latency (avg of 5 get_state calls)") print("=" * W) - print( - f"{'turns':>6} {'ctx size':>10} {'add_msgs':>12} {'delta':>12}" - ) + print(f"{'turns':>6} {'ctx size':>10} {'add_msgs':>12} {'delta':>12}") print("-" * W) for turns, b_bytes, d_bytes, b_rt, d_rt, b_wt, d_wt in rows: print( @@ -346,15 +344,15 @@ def _run_benchmark_for_checkpointer(cp_hint: Any) -> None: # ── Table 3: Per-invoke latency (total write_elapsed / turns) ───────────── print("Per-invoke latency (total graph.invoke time / turns)") print("=" * W) - print( - f"{'turns':>6} {'ctx size':>10} {'add_msgs':>12} {'delta':>12}" - ) + print(f"{'turns':>6} {'ctx size':>10} {'add_msgs':>12} {'delta':>12}") print("-" * W) for turns, b_bytes, d_bytes, b_rt, d_rt, b_wt, d_wt in rows: + def _per(wt: Any) -> str: if wt is None: return "n/a" return f"{(wt / turns) * 1000:.1f}ms" + print( f"{turns:>6} {_approx_tokens(turns):>10} " f"{_per(b_wt):>12} {_per(d_wt):>12}" diff --git a/libs/langgraph/tests/test_delta_channel_migration.py b/libs/langgraph/tests/test_delta_channel_migration.py new file mode 100644 index 000000000..72286522f --- /dev/null +++ b/libs/langgraph/tests/test_delta_channel_migration.py @@ -0,0 +1,504 @@ +"""Tests for the BinaryOperatorAggregate -> DeltaChannel migration path. + +A thread written under `BinaryOperatorAggregate(...)` must keep working +after its annotation is swapped to `DeltaChannel(...)` on the same +checkpointer — pre-migration state visible at each *settled* ancestor +checkpoint is preserved, and post-migration writes fold on top through +the reducer. + +Mechanism under test: the saver's `_get_channel_writes_history(config, +channel)` walks the parent chain; when it encounters an ancestor whose +`channel_values[channel]` is a real value (not `DELTA_SENTINEL`), it +returns that as the `seed`. `DeltaChannel.from_checkpoint(seed)` uses +it as the base value, and `replay_writes(writes)` folds on-path deltas. + +Scenarios covered: + +1. **Basic migration (sync + async)**: build pre-migration state with + `BinaryOperatorAggregate`, swap the annotation to `DeltaChannel` on + the same checkpointer, and verify that every settled pre-migration + super-step boundary (`next=('__start__',)`) round-trips exactly + under the delta-channel view. +2. **Time travel into a pre-migration checkpoint** after migration — + `graph.get_state(pre_migration_config)` at a settled ancestor + returns the same state as under the binop channel. +3. **Continuing a migrated thread**: driving one more super-step after + migration produces a state that includes the pre-migration settled + prefix plus the new delta write — proving `from_checkpoint(seed)` + + `replay_writes` correctly fold post-migration deltas onto the + pre-migration seed. +4. **Base-saver fallback path**: a third-party-style subclass that + removes the optimized `InMemorySaver` override and falls back to + `BaseCheckpointSaver._get_channel_writes_history` must produce the + same result as the optimized path. +5. **Channel-type isolation across threads**: two threads on the same + checkpointer under the delta-channel graph — one freshly-started, + one migrated from pre-migration state — don't cross-contaminate. + The parent-chain walk is scoped to the thread. + +TODO: add postgres variants in the existing `libs/checkpoint-postgres` +test files (different fixture setup; not this file). +""" + +from __future__ import annotations + +import operator +from typing import Annotated, Any + +import pytest +from langgraph.checkpoint.memory import InMemorySaver +from typing_extensions import TypedDict + +from langgraph.channels.binop import BinaryOperatorAggregate +from langgraph.channels.delta import DeltaChannel +from langgraph.graph import END, START, StateGraph + +pytestmark = pytest.mark.anyio + + +# --------------------------------------------------------------------------- +# Graph factories +# +# A minimal reducer (`operator.add` on lists of str) with a noop node keeps +# state change localized to the HumanMessage-like payload passed through +# `invoke`. That isolates the pre/post-migration parity assertions to +# channel-hydration semantics. +# --------------------------------------------------------------------------- + + +def _noop(_state: Any) -> dict: + return {} + + +def _binop_graph(checkpointer: Any) -> Any: + class BinopState(TypedDict): + items: Annotated[list, BinaryOperatorAggregate(list, operator.add)] + + return ( + StateGraph(BinopState) + .add_node("noop", _noop) + .add_edge(START, "noop") + .add_edge("noop", END) + .compile(checkpointer=checkpointer) + ) + + +def _delta_graph(checkpointer: Any) -> Any: + class DeltaState(TypedDict): + items: Annotated[list, DeltaChannel(operator.add)] + + return ( + StateGraph(DeltaState) + .add_node("noop", _noop) + .add_edge(START, "noop") + .add_edge("noop", END) + .compile(checkpointer=checkpointer) + ) + + +def _drive(graph: Any, config: dict, tag: str, n: int) -> None: + for i in range(n): + graph.invoke({"items": [f"{tag}{i}"]}, config) + + +async def _adrive(graph: Any, config: dict, tag: str, n: int) -> None: + for i in range(n): + await graph.ainvoke({"items": [f"{tag}{i}"]}, config) + + +def _settled_boundaries(history: list) -> list[tuple[dict, list]]: + """Return `[(config, items), ...]` for every checkpoint in `history` + whose `next == ('__start__',)` — the stable boundaries between invokes. + """ + return [ + (s.config, list(s.values.get("items", []))) + for s in history + if s.next == ("__start__",) + ] + + +# --------------------------------------------------------------------------- +# 1. Basic migration (sync + async) +# --------------------------------------------------------------------------- + + +def test_basic_migration_preserves_pre_migration_state() -> None: + """Build state under `BinaryOperatorAggregate`, migrate to + `DeltaChannel` on the same checkpointer, and verify that every + settled pre-migration super-step boundary round-trips exactly. + + Settled boundaries (`next=('__start__',)`) are the stable hydration + targets for the migration path: writes that produced the NEXT + super-step are kept as `pending_writes` on the ancestor, so walking + from a descendant finds the ancestor's blob as the seed and + reconstructs the correct state. + """ + + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "basic-sync"}} + + # Pre-migration: accumulate items across 3 invokes. + binop = _binop_graph(checkpointer) + _drive(binop, config, "u", 3) + + pre_boundaries = _settled_boundaries(list(binop.get_state_history(config))) + assert len(pre_boundaries) >= 2, "expected multiple settled boundaries" + + # Migrate: swap the annotation on the same checkpointer. + delta = _delta_graph(checkpointer) + + for cfg, items in pre_boundaries: + snap = delta.get_state(cfg) + assert list(snap.values.get("items", [])) == items, ( + f"snapshot mismatch at {cfg['configurable']['checkpoint_id']}: " + f"expected {items}, got {snap.values.get('items', [])}" + ) + + +async def test_basic_migration_preserves_pre_migration_state_async() -> None: + """Async variant of the basic migration scenario.""" + + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "basic-async"}} + + binop = _binop_graph(checkpointer) + await _adrive(binop, config, "u", 3) + + pre_history = [s async for s in binop.aget_state_history(config)] + pre_boundaries = _settled_boundaries(pre_history) + assert len(pre_boundaries) >= 2 + + delta = _delta_graph(checkpointer) + + for cfg, items in pre_boundaries: + snap = await delta.aget_state(cfg) + assert list(snap.values.get("items", [])) == items, ( + f"async snapshot mismatch at {cfg['configurable']['checkpoint_id']}" + ) + + +# --------------------------------------------------------------------------- +# 2. Time travel into a pre-migration checkpoint after migration +# --------------------------------------------------------------------------- + + +def test_time_travel_into_pre_migration_checkpoint() -> None: + """After migration, `graph.get_state(pre_migration_config)` at a + settled ancestor returns the state as stored at that point.""" + + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "time-travel"}} + + binop = _binop_graph(checkpointer) + _drive(binop, config, "u", 3) + + pre_boundaries = _settled_boundaries(list(binop.get_state_history(config))) + assert pre_boundaries, "no settled ancestors to time-travel to" + + delta = _delta_graph(checkpointer) + + # Pick the oldest non-empty boundary — a long distance to walk back. + non_empty = [(cfg, items) for cfg, items in pre_boundaries if items] + assert non_empty, "expected at least one non-empty boundary" + target_cfg, expected_items = non_empty[-1] + + snap = delta.get_state(target_cfg) + assert list(snap.values.get("items", [])) == expected_items + + +# --------------------------------------------------------------------------- +# 3. Continuing a migrated thread: deltas fold onto pre-migration seed +# --------------------------------------------------------------------------- + + +def test_continuing_migrated_thread_folds_deltas_on_seed() -> None: + """Resume a pre-migration settled ancestor via `invoke(None, cfg)` + under the delta-channel graph. Since the pre-migration checkpoint + has an existing `pending_writes` entry (the input for the NEXT + super-step), re-running from that ancestor reproduces the same + post-ancestor state as the original binop run. + + This proves the seed-terminator + write-replay pipeline works + end-to-end across the migration boundary. + """ + + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "continue"}} + + binop = _binop_graph(checkpointer) + _drive(binop, config, "u", 2) + + # Pick the oldest settled boundary with non-empty state. + pre_boundaries = _settled_boundaries(list(binop.get_state_history(config))) + target_cfg, seed_items = next( + (cfg, items) for cfg, items in reversed(pre_boundaries) if items + ) + assert seed_items, "need a non-empty seed boundary" + + # Migrate and resume from the pre-migration ancestor. `invoke(None, + # cfg)` replays the pending writes staged at `cfg` under the new + # channel; the reducer folds those deltas onto the seed. + delta = _delta_graph(checkpointer) + result = delta.invoke(None, target_cfg) + + # The resumed state must include the pre-migration seed items in order. + result_items = list(result.get("items", [])) + for idx, prefix_item in enumerate(seed_items): + assert result_items[idx] == prefix_item, ( + f"pre-migration seed item at {idx} not preserved: " + f"got {result_items[: idx + 1]}, expected {seed_items}" + ) + + +# --------------------------------------------------------------------------- +# 4. Base-saver fallback path +# --------------------------------------------------------------------------- + + +class _ThirdPartyStyleSaver(InMemorySaver): + """Simulates a third-party saver that inherits the reference + `_get_channel_writes_history` implementation from + `BaseCheckpointSaver` rather than overriding it. + + We rebind the two methods to the base-class versions (via MRO) so + the fallback path is exercised even though the storage layer is + still the in-memory one. + """ + + # MRO: [_ThirdPartyStyleSaver, InMemorySaver, BaseCheckpointSaver, ...] + _get_channel_writes_history = ( # type: ignore[assignment] + InMemorySaver.__mro__[1]._get_channel_writes_history # type: ignore[attr-defined] + ) + _aget_channel_writes_history = ( # type: ignore[assignment] + InMemorySaver.__mro__[1]._aget_channel_writes_history # type: ignore[attr-defined] + ) + + +def test_base_saver_fallback_matches_optimized_override() -> None: + """The reference `BaseCheckpointSaver` implementation must produce + the same migration behavior as the optimized `InMemorySaver` + override. We drive the same migration scenario through both savers + and assert per-snapshot parity in the delta-channel view.""" + + # Fast path: optimized InMemorySaver override. + fast_saver = InMemorySaver() + fast_config = {"configurable": {"thread_id": "fast"}} + fast_binop = _binop_graph(fast_saver) + _drive(fast_binop, fast_config, "u", 3) + fast_delta = _delta_graph(fast_saver) + fast_history = [ + (s.next, list(s.values.get("items", []))) + for s in fast_delta.get_state_history(fast_config) + ] + + # Slow path: base-class fallback. + slow_saver = _ThirdPartyStyleSaver() + slow_config = {"configurable": {"thread_id": "slow"}} + slow_binop = _binop_graph(slow_saver) + _drive(slow_binop, slow_config, "u", 3) + slow_delta = _delta_graph(slow_saver) + slow_history = [ + (s.next, list(s.values.get("items", []))) + for s in slow_delta.get_state_history(slow_config) + ] + + assert slow_history == fast_history, ( + "base-saver fallback should match optimized-override behavior; " + f"fast={fast_history}, slow={slow_history}" + ) + + +# --------------------------------------------------------------------------- +# 5. Thread isolation under mixed-generation storage +# --------------------------------------------------------------------------- + + +def test_delta_and_migrated_threads_do_not_cross_contaminate() -> None: + """Two threads sharing a checkpointer — one migrated from + pre-migration state, one freshly-started under DeltaChannel — must + maintain independent state. The parent-chain walk in + `_get_channel_writes_history` must be scoped to the target thread. + """ + + checkpointer = InMemorySaver() + migrated_cfg = {"configurable": {"thread_id": "migrated"}} + fresh_cfg = {"configurable": {"thread_id": "fresh"}} + + # Thread A: pre-migration build-up. + binop = _binop_graph(checkpointer) + _drive(binop, migrated_cfg, "m", 2) + + # Thread B: fresh delta-channel run. + delta = _delta_graph(checkpointer) + _drive(delta, fresh_cfg, "f", 2) + + # Thread A: migrate and confirm its state is anchored in its own + # thread's pre-migration history (tag 'm'), never mixing in tag 'f'. + migrated_boundaries = _settled_boundaries( + list(delta.get_state_history(migrated_cfg)) + ) + assert migrated_boundaries, "migrated thread has no settled boundaries" + for _, items in migrated_boundaries: + for it in items: + assert it.startswith("m"), ( + f"migrated thread leaked item from other thread: {it}" + ) + + # Thread B: settled boundaries must only contain 'f' tags. + fresh_boundaries = _settled_boundaries(list(delta.get_state_history(fresh_cfg))) + assert fresh_boundaries, "fresh thread has no settled boundaries" + for _, items in fresh_boundaries: + for it in items: + assert it.startswith("f"), ( + f"fresh thread leaked item from migrated thread: {it}" + ) + + +# --------------------------------------------------------------------------- +# 6. Tip-of-pre-migration hydration: the latest checkpoint from a binop-run +# thread has a real accumulated value in its own `channel_values["items"]`. +# When hydrated under the delta-channel graph via `get_state(config)` with no +# `checkpoint_id`, the short-circuit must use that value directly instead of +# walking ancestors (which would skip the tip's own blob). +# --------------------------------------------------------------------------- + + +def test_tip_of_pre_migration_hydrates_directly() -> None: + """`graph.get_state(config)` at the latest (pre-migration) checkpoint + returns the full accumulated list stored in that checkpoint's own + `channel_values`. The hydration must not walk ancestors past it.""" + + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "tip-sync"}} + + binop = _binop_graph(checkpointer) + _drive(binop, config, "u", 3) + + binop_tip = binop.get_state(config) + expected_items = list(binop_tip.values.get("items", [])) + assert expected_items == ["u0", "u1", "u2"], ( + f"sanity: pre-migration tip should accumulate all 3 items, got {expected_items}" + ) + + delta = _delta_graph(checkpointer) + + snap = delta.get_state(config) + assert list(snap.values.get("items", [])) == expected_items, ( + f"tip hydration mismatch: expected {expected_items}, " + f"got {snap.values.get('items', [])}" + ) + + +async def test_tip_of_pre_migration_hydrates_directly_async() -> None: + """Async variant of the tip-of-pre-migration hydration scenario.""" + + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "tip-async"}} + + binop = _binop_graph(checkpointer) + await _adrive(binop, config, "u", 3) + + binop_tip = await binop.aget_state(config) + expected_items = list(binop_tip.values.get("items", [])) + assert expected_items == ["u0", "u1", "u2"] + + delta = _delta_graph(checkpointer) + + snap = await delta.aget_state(config) + assert list(snap.values.get("items", [])) == expected_items, ( + f"async tip hydration mismatch: expected {expected_items}, " + f"got {snap.values.get('items', [])}" + ) + + +# --------------------------------------------------------------------------- +# 7. `update_state` after migration writes a real value to the new +# checkpoint's `channel_values` (not a sentinel). Hydration must use it +# directly — the ancestor walk would skip this blob and return stale state. +# --------------------------------------------------------------------------- + + +def test_update_state_after_migration_uses_written_value() -> None: + """After migrating and running at least one post-migration super-step + (so the thread's tip has a `DELTA_SENTINEL`), `update_state` writes a + concrete value to a new checkpoint's `channel_values`. `get_state` + must reflect that concrete value.""" + + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "update-state"}} + + # Pre-migration: accumulate a little state. + binop = _binop_graph(checkpointer) + _drive(binop, config, "u", 2) + + # Migrate and run one more super-step so the tip is a post-migration + # checkpoint with `DELTA_SENTINEL` in its own `channel_values`. + delta = _delta_graph(checkpointer) + delta.invoke({"items": ["post"]}, config) + + # `update_state` writes a concrete value into a new checkpoint's blob + # via the reducer against the hydrated prior state. + delta.update_state(config, {"items": ["x", "y"]}) + + snap = delta.get_state(config) + updated_items = list(snap.values.get("items", [])) + # Must include the "x","y" update; without the hydration fix, the + # update_state-written blob would be skipped in favor of an ancestor + # walk, and the update values would disappear. + assert "x" in updated_items and "y" in updated_items, ( + f"update_state values missing from snapshot: {updated_items}" + ) + # The "x","y" items should be folded onto the prior accumulated state, + # not stand alone. This verifies the update-written blob is used + # directly by `get_state` (no ancestor walk past it). + assert len(updated_items) >= 4, ( + f"update_state snapshot should preserve pre-update state, got {updated_items}" + ) + assert updated_items[-2:] == ["x", "y"], ( + f"update_state deltas should be at the tail, got {updated_items}" + ) + + +# --------------------------------------------------------------------------- +# 8. Fork from an `update_state` checkpoint: a new run branched off the +# update_state-produced checkpoint must see that checkpoint's concrete +# `channel_values` as its base, with new deltas folded on top. +# --------------------------------------------------------------------------- + + +def test_fork_from_update_state_checkpoint() -> None: + """Branching a new run from the checkpoint produced by `update_state` + must use that checkpoint's concrete blob as the base. Additional + deltas from the forked run fold onto it through the reducer.""" + + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "fork"}} + + # Pre-migration build-up, then migrate and add one post-migration step. + binop = _binop_graph(checkpointer) + _drive(binop, config, "u", 2) + delta = _delta_graph(checkpointer) + delta.invoke({"items": ["post"]}, config) + + # Apply `update_state` and capture the returned config (references + # the new checkpoint produced by the update). + update_cfg = delta.update_state(config, {"items": ["x", "y"]}) + + update_snap = delta.get_state(update_cfg) + base_items = list(update_snap.values.get("items", [])) + assert "x" in base_items and "y" in base_items, ( + f"update_state values missing from snapshot: {base_items}" + ) + assert base_items[-2:] == ["x", "y"], ( + f"sanity: update_state deltas should be at the tail, got {base_items}" + ) + + # Fork: invoke from the update_state checkpoint with a new delta. + forked = delta.invoke({"items": ["fork0"]}, update_cfg) + forked_items = list(forked.get("items", [])) + # The fork must see the update_state-written blob as its base (not + # walk past it), and the new delta must fold on top of it. + assert forked_items[: len(base_items)] == base_items, ( + f"fork lost update_state base: base={base_items}, forked={forked_items}" + ) + assert forked_items[-1] == "fork0", f"fork delta not appended: {forked_items}"