refactor(langgraph): call create_checkpoint directly for the fork seal

create_fork_checkpoint only forwarded to create_checkpoint, and an empty
fork set bumps nothing there either, so the four update_state call sites
pass the set straight through. The fork tests also had two identical
helpers; keep one.
This commit is contained in:
Elior Nataf Lackritz
2026-09-30 15:07:17 -04:00
parent e59ecbc23d
commit 858e55f232
3 changed files with 40 additions and 71 deletions
@@ -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,
+8 -9
View File
@@ -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(
+32 -36
View File
@@ -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 (