diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index cb28563c9..2ba1c9d68 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -102,6 +102,11 @@ def create_checkpoint( continue ch = channels[k] if k in channels_to_snapshot: + # Callers force a full snapshot blob here: exit mode when a + # delta channel reaches its snapshot cadence, and update_state + # on a fresh thread (no ancestor to replay writes from). The + # manual version-bump below only applies to the exit-mode case. + # # In exit mode, the snapshot decision is deferred to exit # time (intermediate steps have do_checkpoint=False). The # channel's count may have reached snapshot_frequency over diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index cf99b82a1..ea0c540fe 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -1997,33 +1997,11 @@ class Pregel( ), ) # save task writes - has_delta_writes = any( - isinstance(channels.get(c), DeltaChannel) - for task in run_tasks - for c, _ in task.writes - ) - 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": {}}, - {}, - ) - 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: + if saved and channel_writes: checkpointer.put_writes(checkpoint_config, channel_writes, task_id) - # apply to checkpoint and save apply_writes( checkpoint, @@ -2032,7 +2010,21 @@ class Pregel( checkpointer.get_next_version, self.trigger_to_nodes, ) - checkpoint = create_checkpoint(checkpoint, channels, step + 1) + # On a fresh thread there is no ancestor to replay DeltaChannel + # writes from, so force a self-contained snapshot in the first + # checkpoint instead of relying on ancestor write-replay. + delta_snapshot = ( + { + k + for k, ch in channels.items() + if isinstance(ch, DeltaChannel) and ch.is_available() + } + if saved is None + else None + ) + checkpoint = create_checkpoint( + checkpoint, channels, step + 1, channels_to_snapshot=delta_snapshot + ) next_config = checkpointer.put( checkpoint_config, checkpoint, @@ -2464,31 +2456,10 @@ class Pregel( ), ) # save task writes - has_delta_writes = any( - isinstance(channels.get(c), DeltaChannel) - for task in run_tasks - for c, _ in task.writes - ) - 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": {}}, - {}, - ) - 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: + if saved and channel_writes: await checkpointer.aput_writes( checkpoint_config, channel_writes, task_id ) @@ -2500,7 +2471,21 @@ class Pregel( checkpointer.get_next_version, self.trigger_to_nodes, ) - checkpoint = create_checkpoint(checkpoint, channels, step + 1) + # On a fresh thread there is no ancestor to replay DeltaChannel + # writes from, so force a self-contained snapshot in the first + # checkpoint instead of relying on ancestor write-replay. + delta_snapshot = ( + { + k + for k, ch in channels.items() + if isinstance(ch, DeltaChannel) and ch.is_available() + } + if saved is None + else None + ) + checkpoint = create_checkpoint( + checkpoint, channels, step + 1, channels_to_snapshot=delta_snapshot + ) # save checkpoint, after applying writes next_config = await checkpointer.aput( checkpoint_config, diff --git a/libs/langgraph/tests/test_delta_channel_update_state.py b/libs/langgraph/tests/test_delta_channel_update_state.py index 177e24feb..e7b3b6b3b 100644 --- a/libs/langgraph/tests/test_delta_channel_update_state.py +++ b/libs/langgraph/tests/test_delta_channel_update_state.py @@ -2,10 +2,13 @@ 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`). +because channel writes were only persisted when a previous checkpoint existed +and no snapshot was written either, so the checkpoint reconstructed to empty. + +Fixed by forcing a self-contained `_DeltaSnapshot` blob into the first +checkpoint on a fresh thread (`saved is None`), so the value is stored inline +and no ancestor write-replay is required. This keeps the read/replay path +untouched. Coverage: @@ -13,8 +16,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 chain shape on a fresh thread (single self-contained update + checkpoint with the snapshot inline and no parent) """ from typing import Annotated, Any @@ -94,7 +97,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.""" + down so the forced-snapshot change for fresh threads doesn't regress it.""" saver = InMemorySaver() graph = _build_graph(saver) config = {"configurable": {"thread_id": "after-invoke-sync"}} @@ -133,9 +136,9 @@ 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.""" + """First update_state forces a self-contained snapshot seed; the second + sees a real parent (`saved is not None`) and anchors its writes under that + seed. Both messages must round-trip in chronological order.""" saver = InMemorySaver() graph = _build_graph(saver) config = {"configurable": {"thread_id": "consecutive-sync"}} @@ -252,14 +255,14 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None: # --------------------------------------------------------------------------- -# Public-API observation of the lazy-stub mechanism +# Public-API observation of the forced-snapshot mechanism # --------------------------------------------------------------------------- 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'`.""" + """A fresh-thread `update_state` should produce a single self-contained + checkpoint visible via `get_state_history`: step=0, `source='update'`, + no parent, with the DeltaChannel value snapshotted inline.""" saver = InMemorySaver() graph = _build_graph(saver) config = {"configurable": {"thread_id": "history-chain"}} @@ -270,25 +273,12 @@ 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 - - update_snapshot, stub_snapshot = history + assert len(history) == 1 + (update_snapshot,) = history 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"] - )