diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index 1cf14932c..79f9556a8 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -108,6 +108,7 @@ 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 @@ -1649,6 +1650,7 @@ class Pregel( # get last checkpoint config = ensure_config(self.config, input_config) saved = checkpointer.get_tuple(config) + first_superstep = fork_pending is None if fork_pending is None: fork_pending = delta_channels_with_pending_writes( self.channels, saved.pending_writes if saved else None @@ -2024,13 +2026,21 @@ class Pregel( ), ) updated_channels = get_updated_channels_from_tasks(run_tasks) - if saved is not None: - for task_id, task in zip(run_task_ids, run_tasks): - channel_writes = [w for w in task.writes if w[0] != PUSH] - if channel_writes: - checkpointer.put_writes( - checkpoint_config, channel_writes, task_id - ) + edited_delta_channels = { + ch + for ch in updated_channels + if isinstance(self.channels.get(ch), DeltaChannel) + } + # The base's other children replay whatever is stored on it, so an + # edit of an older checkpoint snapshots its delta channels here + # instead. Later supersteps address the checkpoint just written. + if ( + first_superstep + and saved is not None + and edited_delta_channels + and _is_older_checkpoint(checkpointer, config, saved) + ): + fork_pending.update(edited_delta_channels) apply_writes( checkpoint, channels, @@ -2058,7 +2068,17 @@ class Pregel( else None, channels_to_snapshot=channels_to_snapshot, ) + sealed = fork_pending.intersection(checkpoint["channel_values"]) fork_pending.difference_update(checkpoint["channel_values"]) + if saved is not None: + for task_id, task in zip(run_task_ids, run_tasks): + channel_writes = [ + w for w in task.writes if w[0] != PUSH and w[0] not in sealed + ] + if channel_writes: + checkpointer.put_writes( + checkpoint_config, channel_writes, task_id + ) next_config = checkpointer.put( checkpoint_config, checkpoint, @@ -2141,6 +2161,7 @@ class Pregel( # get last checkpoint config = ensure_config(self.config, input_config) saved = await checkpointer.aget_tuple(config) + first_superstep = fork_pending is None if fork_pending is None: fork_pending = delta_channels_with_pending_writes( self.channels, saved.pending_writes if saved else None @@ -2511,13 +2532,21 @@ class Pregel( ), ) updated_channels = get_updated_channels_from_tasks(run_tasks) - if saved is not None: - for task_id, task in zip(run_task_ids, run_tasks): - channel_writes = [w for w in task.writes if w[0] != PUSH] - if channel_writes: - await checkpointer.aput_writes( - checkpoint_config, channel_writes, task_id - ) + edited_delta_channels = { + ch + for ch in updated_channels + if isinstance(self.channels.get(ch), DeltaChannel) + } + # The base's other children replay whatever is stored on it, so an + # edit of an older checkpoint snapshots its delta channels here + # instead. Later supersteps address the checkpoint just written. + if ( + first_superstep + and saved is not None + and edited_delta_channels + and await _ais_older_checkpoint(checkpointer, config, saved) + ): + fork_pending.update(edited_delta_channels) apply_writes( checkpoint, channels, @@ -2545,7 +2574,17 @@ class Pregel( else None, channels_to_snapshot=channels_to_snapshot, ) + sealed = fork_pending.intersection(checkpoint["channel_values"]) fork_pending.difference_update(checkpoint["channel_values"]) + if saved is not None: + for task_id, task in zip(run_task_ids, run_tasks): + channel_writes = [ + w for w in task.writes if w[0] != PUSH and w[0] not in sealed + ] + if channel_writes: + await checkpointer.aput_writes( + checkpoint_config, channel_writes, task_id + ) next_config = await checkpointer.aput( checkpoint_config, checkpoint, @@ -4235,6 +4274,30 @@ def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str return dict(trigger_to_nodes) +def _is_older_checkpoint( + checkpointer: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple +) -> bool: + """Whether `config` addressed a checkpoint the thread has moved past.""" + if not config[CONF].get(CONFIG_KEY_CHECKPOINT_ID): + return False + latest = checkpointer.get_tuple( + patch_configurable(config, {CONFIG_KEY_CHECKPOINT_ID: None}) + ) + return latest is not None and latest.checkpoint["id"] != saved.checkpoint["id"] + + +async def _ais_older_checkpoint( + checkpointer: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple +) -> bool: + """Whether `config` addressed a checkpoint the thread has moved past.""" + if not config[CONF].get(CONFIG_KEY_CHECKPOINT_ID): + return False + latest = await checkpointer.aget_tuple( + patch_configurable(config, {CONFIG_KEY_CHECKPOINT_ID: None}) + ) + return latest is not None and latest.checkpoint["id"] != saved.checkpoint["id"] + + def _output( stream_mode: StreamMode | Sequence[StreamMode], print_mode: StreamMode | Sequence[StreamMode], diff --git a/libs/langgraph/tests/test_delta_channel_fork.py b/libs/langgraph/tests/test_delta_channel_fork.py index 8d017c1eb..a805af1c0 100644 --- a/libs/langgraph/tests/test_delta_channel_fork.py +++ b/libs/langgraph/tests/test_delta_channel_fork.py @@ -228,6 +228,118 @@ async def test_afork_by_update_state( assert state.values["log"] == [*base.values["log"], "patched"] +def _assert_branch_unchanged(state: StateSnapshot, expected: list, edit: str) -> None: + assert state.values["log"] == state.values["plain"] == expected, ( + f"{edit!r} was written by an update_state on this branch's base, " + f"but this branch now reads {state.values['log']}" + ) + + +# The old checkpoint is either a finished turn, which saved no writes, or one +# whose next node already ran there, so the edit reuses that task's id. +@pytest.mark.parametrize("next_node_ran", [False, True]) +def test_update_state_on_an_old_checkpoint_leaves_its_other_branch_alone( + sync_checkpointer: BaseCheckpointSaver, next_node_ran: bool +) -> None: + config = _thread("t") + graph = _build(sync_checkpointer, "first") + graph.invoke(_both("in-1"), config) + _build(sync_checkpointer, "second").invoke(_both("in-2"), config) + branch = graph.get_state(config) + base = next( + snapshot + for snapshot in graph.get_state_history(config) + if "in-2" not in snapshot.values["log"] + and snapshot.next == (("n",) if next_node_ran else ()) + ) + + edited = graph.update_state(_at(config, base), _both("edit"), as_node="n") + + _assert_branch_unchanged( + graph.get_state(branch.config), branch.values["log"], "edit" + ) + assert graph.get_state(edited).values["log"] == [*base.values["log"], "edit"] + + _build(sync_checkpointer, "third").invoke(_both("in-3"), branch.config) + _assert_branch_unchanged( + graph.get_state(config), + [*branch.values["log"], "in-3", "third-out"], + "edit", + ) + + +@pytest.mark.parametrize("next_node_ran", [False, True]) +async def test_aupdate_state_on_an_old_checkpoint_leaves_its_other_branch_alone( + async_checkpointer: BaseCheckpointSaver, next_node_ran: bool +) -> None: + config = _thread("t") + graph = _build(async_checkpointer, "first") + await graph.ainvoke(_both("in-1"), config) + await _build(async_checkpointer, "second").ainvoke(_both("in-2"), config) + branch = await graph.aget_state(config) + base = await anext( + snapshot + async for snapshot in graph.aget_state_history(config) + if "in-2" not in snapshot.values["log"] + and snapshot.next == (("n",) if next_node_ran else ()) + ) + + edited = await graph.aupdate_state(_at(config, base), _both("edit"), as_node="n") + + _assert_branch_unchanged( + await graph.aget_state(branch.config), branch.values["log"], "edit" + ) + assert (await graph.aget_state(edited)).values["log"] == [ + *base.values["log"], + "edit", + ] + + await _build(async_checkpointer, "third").ainvoke(_both("in-3"), branch.config) + _assert_branch_unchanged( + await graph.aget_state(config), + [*branch.values["log"], "in-3", "third-out"], + "edit", + ) + + +def test_bulk_update_on_an_old_checkpoint_leaves_its_other_branch_alone( + sync_checkpointer: BaseCheckpointSaver, +) -> None: + config = _thread("t") + graph = _build(sync_checkpointer, "first") + graph.invoke(_both("in-1"), config) + base = graph.get_state(config) + _build(sync_checkpointer, "second").invoke(_both("in-2"), config) + branch = graph.get_state(config) + + edited = graph.bulk_update_state( + _at(config, base), + [[StateUpdate(_both("s1"), "n")], [StateUpdate(_both("s2"), "n")]], + ) + + _assert_branch_unchanged(graph.get_state(branch.config), branch.values["log"], "s1") + assert graph.get_state(edited).values["log"] == [*base.values["log"], "s1", "s2"] + + +def test_update_state_with_the_head_checkpoint_id_stores_no_snapshot( + sync_checkpointer: BaseCheckpointSaver, +) -> None: + config = _thread("t") + graph = _build(sync_checkpointer, "first") + graph.invoke(_both("in-1"), config) + for i in range(3): + graph.update_state(graph.get_state(config).config, _both(f"u{i}")) + + assert not _snapshotted_checkpoints(sync_checkpointer, config) + assert graph.get_state(config).values["log"] == [ + "in-1", + "first-out", + "u0", + "u1", + "u2", + ] + + def test_unaddressed_run_keeps_snapshot_cadence( sync_checkpointer: BaseCheckpointSaver, durability: Durability ) -> None: