mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 05:55:14 +02:00
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:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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 (
|
||||
|
||||
Reference in New Issue
Block a user