diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index cb28563c9..8679a5c85 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -1,7 +1,7 @@ from __future__ import annotations import uuid -from collections.abc import Callable, Mapping +from collections.abc import Callable, Iterable, Mapping, Sequence from datetime import datetime, timezone from typing import Any, cast @@ -14,6 +14,7 @@ from langgraph.checkpoint.base.id import uuid6 from langgraph.checkpoint.serde.types import _DeltaSnapshot from langgraph._internal._config import DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT +from langgraph._internal._constants import PUSH from langgraph._internal._typing import MISSING from langgraph.channels.base import BaseChannel from langgraph.channels.delta import DeltaChannel @@ -70,6 +71,26 @@ def delta_channels_to_snapshot( return result +def update_state_channels_plan( + run_tasks: Iterable[Any], + channels: Mapping[str, BaseChannel], +) -> tuple[set[str], set[str]]: + """Return channels written and DeltaChannels to snapshot on update_state.""" + updated_channels = {c for task in run_tasks for c, _ in task.writes if c != PUSH} + channels_to_snapshot = { + c for c in updated_channels if isinstance(channels.get(c), DeltaChannel) + } + return updated_channels, channels_to_snapshot + + +def update_state_channel_writes( + writes: Sequence[tuple[str, Any]], + channels_to_snapshot: set[str], +) -> list[tuple[str, Any]]: + """Channel writes to persist separately from a head snapshot.""" + return [w for w in writes if w[0] != PUSH and w[0] not in channels_to_snapshot] + + def create_checkpoint( checkpoint: Checkpoint, channels: Mapping[str, BaseChannel] | None, diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index cf99b82a1..c8af86c6e 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -108,7 +108,6 @@ from langgraph.callbacks import ( get_sync_graph_callback_manager_for_config, ) from langgraph.channels.base import BaseChannel -from langgraph.channels.delta import DeltaChannel from langgraph.channels.topic import Topic from langgraph.config import get_config from langgraph.constants import END @@ -134,6 +133,8 @@ from langgraph.pregel._checkpoint import ( copy_checkpoint, create_checkpoint, empty_checkpoint, + update_state_channel_writes, + update_state_channels_plan, ) from langgraph.pregel._draw import draw_graph from langgraph.pregel._io import map_input, read_channels @@ -1996,35 +1997,19 @@ class Pregel( }, ), ) - # save task writes - has_delta_writes = any( - isinstance(channels.get(c), DeltaChannel) - for task in run_tasks - for c, _ in task.writes + updated_channels, channels_to_snapshot = update_state_channels_plan( + run_tasks, channels ) - should_put_writes = saved is not None or has_delta_writes - if saved is None and has_delta_writes: - # If there is no previous checkpoint, we need to create a stub checkpoint - # so the first delta writes has a parent to anchor under. - # This is the model of DeltaChannel. - stub = empty_checkpoint() - checkpoint_config = checkpointer.put( - patch_configurable( - checkpoint_config, {CONFIG_KEY_CHECKPOINT_ID: None} - ), - stub, - {"source": "update", "step": -1, "parents": {}}, - {}, - ) + if saved is not None: + for task_id, task in zip(run_task_ids, run_tasks): + if channel_writes := update_state_channel_writes( + task.writes, channels_to_snapshot + ): + checkpointer.put_writes( + checkpoint_config, channel_writes, task_id + ) - for task_id, task in zip(run_task_ids, run_tasks): - # channel writes are saved to current checkpoint - channel_writes = [w for w in task.writes if w[0] != PUSH] - if should_put_writes and channel_writes: - checkpointer.put_writes(checkpoint_config, channel_writes, task_id) - - # apply to checkpoint and save apply_writes( checkpoint, channels, @@ -2032,7 +2017,14 @@ class Pregel( checkpointer.get_next_version, self.trigger_to_nodes, ) - checkpoint = create_checkpoint(checkpoint, channels, step + 1) + checkpoint = create_checkpoint( + checkpoint, + channels, + step + 1, + updated_channels=updated_channels, + get_next_version=checkpointer.get_next_version, + channels_to_snapshot=channels_to_snapshot, + ) next_config = checkpointer.put( checkpoint_config, checkpoint, @@ -2046,7 +2038,11 @@ class Pregel( ), ) for task_id, task in zip(run_task_ids, run_tasks): - # save push writes + if saved is None: + if channel_writes := update_state_channel_writes( + task.writes, channels_to_snapshot + ): + checkpointer.put_writes(next_config, channel_writes, task_id) if push_writes := [w for w in task.writes if w[0] == PUSH]: checkpointer.put_writes(next_config, push_writes, task_id) @@ -2463,36 +2459,19 @@ class Pregel( }, ), ) - # save task writes - has_delta_writes = any( - isinstance(channels.get(c), DeltaChannel) - for task in run_tasks - for c, _ in task.writes + updated_channels, channels_to_snapshot = update_state_channels_plan( + run_tasks, channels ) - should_put_writes = saved is not None or has_delta_writes - if saved is None and has_delta_writes: - # If there is no previous checkpoint, we need to create a stub checkpoint - # so the first delta writes has a parent to anchor under. - # This is the model of DeltaChannel. - stub = empty_checkpoint() - checkpoint_config = await checkpointer.aput( - patch_configurable( - checkpoint_config, {CONFIG_KEY_CHECKPOINT_ID: None} - ), - stub, - {"source": "update", "step": -1, "parents": {}}, - {}, - ) + if saved is not None: + for task_id, task in zip(run_task_ids, run_tasks): + if channel_writes := update_state_channel_writes( + task.writes, channels_to_snapshot + ): + await checkpointer.aput_writes( + checkpoint_config, channel_writes, task_id + ) - for task_id, task in zip(run_task_ids, run_tasks): - # channel writes are saved to current checkpoint - channel_writes = [w for w in task.writes if w[0] != PUSH] - if should_put_writes and channel_writes: - await checkpointer.aput_writes( - checkpoint_config, channel_writes, task_id - ) - # apply to checkpoint and save apply_writes( checkpoint, channels, @@ -2500,8 +2479,14 @@ class Pregel( checkpointer.get_next_version, self.trigger_to_nodes, ) - checkpoint = create_checkpoint(checkpoint, channels, step + 1) - # save checkpoint, after applying writes + checkpoint = create_checkpoint( + checkpoint, + channels, + step + 1, + updated_channels=updated_channels, + get_next_version=checkpointer.get_next_version, + channels_to_snapshot=channels_to_snapshot, + ) next_config = await checkpointer.aput( checkpoint_config, checkpoint, @@ -2515,7 +2500,13 @@ class Pregel( ), ) for task_id, task in zip(run_task_ids, run_tasks): - # save push writes + if saved is None: + if channel_writes := update_state_channel_writes( + task.writes, channels_to_snapshot + ): + await checkpointer.aput_writes( + next_config, channel_writes, task_id + ) if push_writes := [w for w in task.writes if w[0] == PUSH]: await checkpointer.aput_writes(next_config, push_writes, task_id) return patch_checkpoint_map(next_config, saved.metadata if saved else None) diff --git a/libs/langgraph/tests/test_delta_channel_update_state.py b/libs/langgraph/tests/test_delta_channel_update_state.py index 177e24feb..219e1162b 100644 --- a/libs/langgraph/tests/test_delta_channel_update_state.py +++ b/libs/langgraph/tests/test_delta_channel_update_state.py @@ -1,11 +1,12 @@ """Tests for `update_state` / `aupdate_state` against `DeltaChannel`. -Originally a regression suite for deepagents#3774 — `update_state` on a *fresh* -thread silently dropped the first write to a `DeltaChannel`-backed channel -because channel writes were only persisted when a previous checkpoint existed. -Fixed by lazily persisting an empty stub checkpoint on a fresh thread so the -first write has a parent to anchor under (mirrors the exit-mode lazy-stub -pattern in `_loop._put_exit_delta_writes`). +Regression suite for deepagents#3774 and Postgres read-path compatibility: +fresh-thread `update_state` must persist DeltaChannel state correctly. + +Fresh-thread updates snapshot updated DeltaChannels on the head checkpoint +(self-contained for Postgres readers). Delta writes are not persisted via +`put_writes`; non-delta channel writes on a fresh thread are attached to +the head after it is saved. Coverage: @@ -13,8 +14,8 @@ Coverage: * non-fresh thread: `update_state` after `invoke`, after another `update_state`, and `bulk_update_state` with multiple per-superstep updates * update-by-id end-to-end via `update_state` (DeltaChannel reducer semantics) -* state-history chain shape on a fresh thread (lazy stub + update checkpoint - with correct parent linking) +* state-history shape on a fresh thread (single snapshotted head checkpoint) +* head checkpoint snapshots updated DeltaChannels for Postgres read paths """ from typing import Annotated, Any @@ -22,6 +23,7 @@ from typing import Annotated, Any import pytest from langchain_core.messages import HumanMessage from langgraph.checkpoint.memory import InMemorySaver +from langgraph.checkpoint.serde.types import _DeltaSnapshot from typing_extensions import TypedDict from langgraph.channels.delta import DeltaChannel @@ -93,8 +95,7 @@ async def test_aupdate_state_fresh_thread_delta_channel() -> None: def test_update_state_after_invoke_delta_channel() -> None: - """The non-fresh-thread path was already working before the fix; pin it - down so the lazy-stub change for fresh threads doesn't regress it.""" + """The non-fresh-thread path must keep working across snapshot changes.""" saver = InMemorySaver() graph = _build_graph(saver) config = {"configurable": {"thread_id": "after-invoke-sync"}} @@ -133,9 +134,8 @@ async def test_aupdate_state_after_invoke_delta_channel() -> None: def test_consecutive_update_states_delta_channel() -> None: - """First update_state lazily persists a stub; the second sees a real - parent (`saved is not None`) and takes the original write path. Both - messages must round-trip in chronological order.""" + """Two consecutive fresh-thread-style updates: the first creates a + snapshotted head; the second anchors on it. Both messages round-trip.""" saver = InMemorySaver() graph = _build_graph(saver) config = {"configurable": {"thread_id": "consecutive-sync"}} @@ -209,10 +209,8 @@ def test_update_state_replaces_message_by_id_delta_channel() -> None: def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None: - """`bulk_update_state` with N updates in one superstep produces N tasks - that each call `put_writes`. Guards the regression where moving - `put_writes` outside the per-task loop would persist only the last - task's writes. + """`bulk_update_state` with N updates in one superstep must accumulate + all N message writes in the snapshotted head state. Explicit `task_id`s are required to disambiguate writes belonging to different `StateUpdate`s targeting the same node — otherwise both share @@ -252,14 +250,13 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None: # --------------------------------------------------------------------------- -# Public-API observation of the lazy-stub mechanism +# Public-API observation of fresh-thread checkpoint shape # --------------------------------------------------------------------------- def test_state_history_chain_after_fresh_update_state_delta_channel() -> None: - """A fresh-thread `update_state` should produce two checkpoints visible - via `get_state_history`: a stub (step=-1, no parent) and the update - (step=0, parent=stub). Both attributed `source='update'`.""" + """Fresh-thread `update_state` on snapshotted DeltaChannels yields one + self-contained checkpoint (step=0, no parent, source='update').""" saver = InMemorySaver() graph = _build_graph(saver) config = {"configurable": {"thread_id": "history-chain"}} @@ -270,25 +267,35 @@ def test_state_history_chain_after_fresh_update_state_delta_channel() -> None: as_node="model", ) - # Newest first per `get_state_history` ordering. history = list(graph.get_state_history(config)) - assert len(history) == 2 + assert len(history) == 1 - update_snapshot, stub_snapshot = history + update_snapshot = history[0] assert update_snapshot.metadata is not None assert update_snapshot.metadata["source"] == "update" assert update_snapshot.metadata["step"] == 0 + assert update_snapshot.parent_config is None assert [m.content for m in update_snapshot.values["messages"]] == ["hello"] - assert stub_snapshot.metadata is not None - assert stub_snapshot.metadata["source"] == "update" - assert stub_snapshot.metadata["step"] == -1 - assert stub_snapshot.parent_config is None - # The update checkpoint's parent is the stub. - assert update_snapshot.parent_config is not None - assert ( - update_snapshot.parent_config["configurable"]["checkpoint_id"] - == stub_snapshot.config["configurable"]["checkpoint_id"] +def test_fresh_update_state_head_snapshots_delta_channel() -> None: + """Postgres checkpointers skip the ancestor walk when the head checkpoint + has no `counters_since_delta_snapshot` entry. Force-snapshot updated + DeltaChannels on the update checkpoint so the head is self-contained.""" + saver = InMemorySaver() + graph = _build_graph(saver) + config = {"configurable": {"thread_id": "head-snapshot"}} + + graph.update_state( + config, + {"messages": [HumanMessage(content="hello", id="m1")]}, + as_node="model", ) + + head = saver.get_tuple(config) + assert head is not None + assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot) + assert [m.content for m in head.checkpoint["channel_values"]["messages"].value] == [ + "hello" + ]