Compare commits

..
Author SHA1 Message Date
Elior Nataf Lackritz 22083959f4 fix(langgraph): keep an update_state on an older checkpoint out of its other branches
update_state stores its writes on the checkpoint it addresses. When that
checkpoint already has other children, a DeltaChannel in each of them
replays those writes, so editing an earlier turn leaks the edit into the
branch it forked away from.

When the addressed checkpoint is not the thread's latest, the new
checkpoint snapshots the delta channels the update writes, and those
channels' writes are no longer stored on the addressed checkpoint.
Updates on the latest checkpoint are unchanged.
2026-10-01 12:20:24 -04:00
2 changed files with 189 additions and 14 deletions
+77 -14
View File
@@ -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],
@@ -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: