diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index 25a7cd331..1ea3dd258 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -168,32 +168,6 @@ def create_checkpoint_plan_for_update_state_api( return channels_to_snapshot, metadata -def create_fork_checkpoint( - checkpoint: Checkpoint, - channels: Mapping[str, BaseChannel], - step: int, - *, - fork_channels: set[str], - get_next_version: GetNextVersion, -) -> Checkpoint: - """``create_checkpoint`` for the update_state paths that skip the plan. - - The fork has to be sealed by its first checkpoint: any later superstep - has already rebuilt its delta channels through the shared base. These - paths never write the delta channel, so its version must be bumped here - or ``put`` drops the blob; derive ``new_versions`` from the result. - """ - if not fork_channels: - return create_checkpoint(checkpoint, channels, step) - return create_checkpoint( - checkpoint, - channels, - step, - get_next_version=get_next_version, - channels_to_snapshot=fork_channels, - ) - - 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 1bde8f00f..1cf14932c 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -133,7 +133,6 @@ from langgraph.pregel._checkpoint import ( copy_checkpoint, create_checkpoint, create_checkpoint_plan_for_update_state_api, - create_fork_checkpoint, delta_channels_with_pending_writes, empty_checkpoint, get_updated_channels_from_tasks, @@ -1737,12 +1736,12 @@ class Pregel( self.trigger_to_nodes, ) # save checkpoint - next_checkpoint = create_fork_checkpoint( + next_checkpoint = create_checkpoint( checkpoint, channels, step, - fork_channels=fork_pending, get_next_version=checkpointer.get_next_version, + channels_to_snapshot=fork_pending, ) fork_pending.difference_update(next_checkpoint["channel_values"]) next_config = checkpointer.put( @@ -1784,12 +1783,12 @@ class Pregel( if saved and saved.metadata.get("step") is not None else -1 ) - next_checkpoint = create_fork_checkpoint( + next_checkpoint = create_checkpoint( checkpoint, channels, next_step, - fork_channels=fork_pending, get_next_version=checkpointer.get_next_version, + channels_to_snapshot=fork_pending, ) fork_pending.difference_update(next_checkpoint["channel_values"]) next_config = checkpointer.put( @@ -2227,12 +2226,12 @@ class Pregel( self.trigger_to_nodes, ) # save checkpoint - next_checkpoint = create_fork_checkpoint( + next_checkpoint = create_checkpoint( checkpoint, channels, step, - fork_channels=fork_pending, get_next_version=checkpointer.get_next_version, + channels_to_snapshot=fork_pending, ) fork_pending.difference_update(next_checkpoint["channel_values"]) next_config = await checkpointer.aput( @@ -2274,12 +2273,12 @@ class Pregel( if saved and saved.metadata.get("step") is not None else -1 ) - next_checkpoint = create_fork_checkpoint( + next_checkpoint = create_checkpoint( checkpoint, channels, next_step, - fork_channels=fork_pending, get_next_version=checkpointer.get_next_version, + channels_to_snapshot=fork_pending, ) fork_pending.difference_update(next_checkpoint["channel_values"]) next_config = await checkpointer.aput( diff --git a/libs/langgraph/tests/test_delta_channel_fork.py b/libs/langgraph/tests/test_delta_channel_fork.py index 87a84f6df..8d017c1eb 100644 --- a/libs/langgraph/tests/test_delta_channel_fork.py +++ b/libs/langgraph/tests/test_delta_channel_fork.py @@ -1,6 +1,6 @@ """Forking a thread must not replay the abandoned branch into the fork. -Every graph carries a ``DeltaChannel`` and a plain reducer channel fed the same +Every graph carries a `DeltaChannel` and a plain reducer channel fed the same values; the plain channel needs no replay, so it is the oracle. """ @@ -71,7 +71,7 @@ def _at(config: RunnableConfig, snapshot: StateSnapshot) -> RunnableConfig: } -def _input(marker: str) -> dict: +def _both(marker: str) -> dict: return {"log": [marker], "plain": [marker]} @@ -101,10 +101,10 @@ def test_fork_by_invoke( ) -> None: config = _thread("t") _build(sync_checkpointer, "first").invoke( - _input("in-1"), config, durability=durability + _both("in-1"), config, durability=durability ) graph = _build(sync_checkpointer, "second") - graph.invoke(_input("in-2"), config, durability=durability) + graph.invoke(_both("in-2"), config, durability=durability) abandoned_head = graph.get_state(config) base = next( @@ -113,7 +113,7 @@ def test_fork_by_invoke( if "in-2" not in snapshot.values["log"] ) _build(sync_checkpointer, "third").invoke( - _input("in-3"), _at(config, base), durability=durability + _both("in-3"), _at(config, base), durability=durability ) state = graph.get_state(config) @@ -129,10 +129,10 @@ async def test_afork_by_invoke( ) -> None: config = _thread("t") await _build(async_checkpointer, "first").ainvoke( - _input("in-1"), config, durability=durability + _both("in-1"), config, durability=durability ) graph = _build(async_checkpointer, "second") - await graph.ainvoke(_input("in-2"), config, durability=durability) + await graph.ainvoke(_both("in-2"), config, durability=durability) abandoned_head = await graph.aget_state(config) base = await anext( @@ -141,7 +141,7 @@ async def test_afork_by_invoke( if "in-2" not in snapshot.values["log"] ) await _build(async_checkpointer, "third").ainvoke( - _input("in-3"), _at(config, base), durability=durability + _both("in-3"), _at(config, base), durability=durability ) state = await graph.aget_state(config) @@ -157,13 +157,13 @@ def test_fork_off_checkpoint_before_first_input( ) -> None: config = _thread("t") graph = _build(sync_checkpointer, "first") - graph.invoke(_input("in-1"), config, durability=durability) + graph.invoke(_both("in-1"), config, durability=durability) root = list(graph.get_state_history(config))[-1] assert root.values["log"] == [] _build(sync_checkpointer, "third").invoke( - _input("in-9"), _at(config, root), durability=durability + _both("in-9"), _at(config, root), durability=durability ) state = graph.get_state(config) @@ -176,13 +176,13 @@ async def test_afork_off_checkpoint_before_first_input( ) -> None: config = _thread("t") graph = _build(async_checkpointer, "first") - await graph.ainvoke(_input("in-1"), config, durability=durability) + await graph.ainvoke(_both("in-1"), config, durability=durability) root = [snapshot async for snapshot in graph.aget_state_history(config)][-1] assert root.values["log"] == [] await _build(async_checkpointer, "third").ainvoke( - _input("in-9"), _at(config, root), durability=durability + _both("in-9"), _at(config, root), durability=durability ) state = await graph.aget_state(config) @@ -192,16 +192,16 @@ async def test_afork_off_checkpoint_before_first_input( def test_fork_by_update_state(sync_checkpointer: BaseCheckpointSaver) -> None: config = _thread("t") - _build(sync_checkpointer, "first").invoke(_input("in-1"), config) + _build(sync_checkpointer, "first").invoke(_both("in-1"), config) graph = _build(sync_checkpointer, "second") - graph.invoke(_input("in-2"), config) + graph.invoke(_both("in-2"), config) base = next( snapshot for snapshot in graph.get_state_history(config) if "in-2" not in snapshot.values["log"] ) - forked = graph.update_state(_at(config, base), _input("patched")) + forked = graph.update_state(_at(config, base), _both("patched")) state = graph.get_state(forked) _assert_fork_is_clean(state, "in-2") @@ -212,16 +212,16 @@ async def test_afork_by_update_state( async_checkpointer: BaseCheckpointSaver, ) -> None: config = _thread("t") - await _build(async_checkpointer, "first").ainvoke(_input("in-1"), config) + await _build(async_checkpointer, "first").ainvoke(_both("in-1"), config) graph = _build(async_checkpointer, "second") - await graph.ainvoke(_input("in-2"), config) + await graph.ainvoke(_both("in-2"), config) base = await anext( snapshot async for snapshot in graph.aget_state_history(config) if "in-2" not in snapshot.values["log"] ) - forked = await graph.aupdate_state(_at(config, base), _input("patched")) + forked = await graph.aupdate_state(_at(config, base), _both("patched")) state = await graph.aget_state(forked) _assert_fork_is_clean(state, "in-2") @@ -233,8 +233,8 @@ def test_unaddressed_run_keeps_snapshot_cadence( ) -> None: config = _thread("t") graph = _build(sync_checkpointer, "first") - graph.invoke(_input("in-1"), config, durability=durability) - graph.invoke(_input("in-2"), config, durability=durability) + graph.invoke(_both("in-1"), config, durability=durability) + graph.invoke(_both("in-2"), config, durability=durability) assert not _snapshotted_checkpoints(sync_checkpointer, config) @@ -244,7 +244,7 @@ def test_fork_before_first_value_when_fork_never_writes_the_channel( ) -> None: config = _thread("t") graph = _build(sync_checkpointer, "first") - graph.invoke(_input("in-1"), config, durability=durability) + graph.invoke(_both("in-1"), config, durability=durability) root = list(graph.get_state_history(config))[-1] assert root.values["log"] == [] @@ -263,7 +263,7 @@ async def test_afork_before_first_value_when_fork_never_writes_the_channel( ) -> None: config = _thread("t") graph = _build(async_checkpointer, "first") - await graph.ainvoke(_input("in-1"), config, durability=durability) + await graph.ainvoke(_both("in-1"), config, durability=durability) root = [snapshot async for snapshot in graph.aget_state_history(config)][-1] assert root.values["log"] == [] @@ -282,7 +282,7 @@ def test_fork_before_first_value_by_bulk_update( ) -> None: config = _thread("t") graph = _build(sync_checkpointer, "first") - graph.invoke(_input("in-1"), config) + graph.invoke(_both("in-1"), config) root = list(graph.get_state_history(config))[-1] assert root.values["log"] == [] @@ -291,7 +291,7 @@ def test_fork_before_first_value_by_bulk_update( _at(config, root), [ [StateUpdate({"other": ["s1"]}, "n")], - [StateUpdate(_input("s2"), "n")], + [StateUpdate(_both("s2"), "n")], ], ) @@ -305,9 +305,9 @@ def test_fork_by_bulk_update_whose_first_superstep_skips_the_plan( sync_checkpointer: BaseCheckpointSaver, first_as_node: str ) -> None: config = _thread("t") - _build(sync_checkpointer, "first").invoke(_input("in-1"), config) + _build(sync_checkpointer, "first").invoke(_both("in-1"), config) graph = _build(sync_checkpointer, "second") - graph.invoke(_input("in-2"), config) + graph.invoke(_both("in-2"), config) base = next( snapshot @@ -315,13 +315,13 @@ def test_fork_by_bulk_update_whose_first_superstep_skips_the_plan( if "in-2" not in snapshot.values["log"] ) first = ( - StateUpdate(_input("first-step"), first_as_node) + StateUpdate(_both("first-step"), first_as_node) if first_as_node == INPUT else StateUpdate(None, first_as_node) ) forked = graph.bulk_update_state( _at(config, base), - [[first], [StateUpdate(_input("second-step"), "n")]], + [[first], [StateUpdate(_both("second-step"), "n")]], ) state = graph.get_state(forked) @@ -336,20 +336,16 @@ def test_unaddressed_bulk_update_keeps_snapshot_cadence( ) -> None: config = _thread("t") graph = _build(sync_checkpointer, "first") - graph.invoke(_input("in-1"), config) + graph.invoke(_both("in-1"), config) graph.bulk_update_state( config, - [[StateUpdate(_input(f"u{i}"), "n")] for i in range(4)], + [[StateUpdate(_both(f"u{i}"), "n")] for i in range(4)], ) assert not _snapshotted_checkpoints(sync_checkpointer, config) -def _both(marker: str) -> dict: - return {"log": [marker], "plain": [marker]} - - def _build_paused_before_b(checkpointer: BaseCheckpointSaver) -> Any: builder = StateGraph(_State) builder.add_node("a", lambda state: _both("a")) @@ -527,9 +523,9 @@ def test_turns_addressed_at_the_head_store_no_snapshot( ) -> None: config = _thread("t") graph = _build(sync_checkpointer, "turn") - graph.invoke(_input("in-1"), config) + graph.invoke(_both("in-1"), config) for turn in range(2, 5): - graph.invoke(_input(f"in-{turn}"), graph.get_state(config).config) + graph.invoke(_both(f"in-{turn}"), graph.get_state(config).config) assert not _snapshotted_checkpoints(sync_checkpointer, config) assert (