mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 22:15:11 +02:00
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.
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user