mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
858e55f232 | ||
|
|
e59ecbc23d | ||
|
|
719a4d71bc | ||
|
|
ddaf708cd0 | ||
|
|
9d5b0f1991 | ||
|
|
9d16b52955 | ||
|
|
2de0c47c1f | ||
|
|
377083220e | ||
|
|
9c5914861b | ||
|
|
24cf33f348 |
@@ -8,13 +8,15 @@ from typing import Any, cast
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langgraph.checkpoint.base import (
|
||||
BaseCheckpointSaver,
|
||||
ChannelVersions,
|
||||
Checkpoint,
|
||||
PendingWrite,
|
||||
)
|
||||
from langgraph.checkpoint.base.id import uuid6
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
|
||||
from langgraph._internal._config import DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT
|
||||
from langgraph._internal._constants import PUSH
|
||||
from langgraph._internal._constants import INTERRUPT, PUSH
|
||||
from langgraph._internal._typing import MISSING
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
@@ -89,6 +91,23 @@ def get_delta_channels_from_all_channels(
|
||||
}
|
||||
|
||||
|
||||
def delta_channels_with_pending_writes(
|
||||
specs: Mapping[str, Any],
|
||||
pending_writes: Iterable[PendingWrite] | None,
|
||||
) -> set[str]:
|
||||
"""DeltaChannels a branch starting from this checkpoint must snapshot.
|
||||
|
||||
A checkpoint's pending writes belong to the child that consumed them, and
|
||||
nothing records which child that was. A new branch snapshots every delta
|
||||
channel they touch, so its ancestor walk never replays them.
|
||||
"""
|
||||
return {
|
||||
ch
|
||||
for _, ch, _ in pending_writes or ()
|
||||
if isinstance(specs.get(ch), DeltaChannel)
|
||||
}
|
||||
|
||||
|
||||
def create_metadata_for_update_state_api(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
updated_channels: set[str],
|
||||
@@ -122,6 +141,7 @@ def create_checkpoint_plan_for_update_state_api(
|
||||
parents: dict[str, Any],
|
||||
saved_metadata: Mapping[str, Any] | None,
|
||||
is_fresh_thread: bool,
|
||||
fork_channels: set[str],
|
||||
) -> tuple[set[str], dict[str, Any]]:
|
||||
"""Return ``(channels_to_snapshot, metadata)`` for an update_state head."""
|
||||
metadata: dict[str, Any] = {
|
||||
@@ -137,7 +157,9 @@ def create_checkpoint_plan_for_update_state_api(
|
||||
updated_channels,
|
||||
prev_metadata=saved_metadata,
|
||||
)
|
||||
channels_to_snapshot = delta_channels_to_snapshot(channels, new_counters)
|
||||
channels_to_snapshot = (
|
||||
delta_channels_to_snapshot(channels, new_counters) | fork_channels
|
||||
)
|
||||
for k in channels_to_snapshot:
|
||||
new_counters[k] = (0, 0)
|
||||
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
|
||||
@@ -167,6 +189,7 @@ def create_checkpoint(
|
||||
"""
|
||||
ts = datetime.now(timezone.utc).isoformat()
|
||||
channels_to_snapshot = channels_to_snapshot or set()
|
||||
bumped: dict[str, tuple[Any, Any]] = {}
|
||||
if channels is None:
|
||||
values = checkpoint["channel_values"]
|
||||
channel_versions = checkpoint["channel_versions"]
|
||||
@@ -174,30 +197,29 @@ def create_checkpoint(
|
||||
values = {}
|
||||
channel_versions = dict(checkpoint["channel_versions"])
|
||||
for k in channels:
|
||||
if k not in channel_versions:
|
||||
continue
|
||||
ch = channels[k]
|
||||
if k not in channel_versions:
|
||||
# A forced snapshot of a never-written channel still has to
|
||||
# land to stop the ancestor walk, and `put` only stores blobs
|
||||
# for versioned channels.
|
||||
if k in channels_to_snapshot and get_next_version is not None:
|
||||
channel_versions[k] = get_next_version(None, None)
|
||||
bumped[k] = (None, channel_versions[k])
|
||||
values[k] = _DeltaSnapshot(
|
||||
ch.get() if ch.is_available() else ch.typ()
|
||||
)
|
||||
continue
|
||||
if k in channels_to_snapshot:
|
||||
# Callers force a full snapshot blob here: exit mode when a
|
||||
# delta channel reaches its snapshot cadence, and update_state
|
||||
# on a fresh thread (no ancestor to replay writes from). The
|
||||
# manual version-bump below only applies to the exit-mode case.
|
||||
#
|
||||
# In exit mode, the snapshot decision is deferred to exit
|
||||
# time (intermediate steps have do_checkpoint=False). The
|
||||
# channel's count may have reached snapshot_frequency over
|
||||
# several supersteps, but the LAST superstep may not have
|
||||
# written to this channel. In that case apply_writes()
|
||||
# (in _algo.py) didn't bump this channel's version, so
|
||||
# saver.put() wouldn't include it in new_versions and
|
||||
# the snapshot blob would be silently dropped. The manual
|
||||
# bump below closes the gap. In sync/async durability this
|
||||
# branch is effectively dead code (the step that pushes
|
||||
# the count to freq always writes the channel).
|
||||
# `put` only stores a blob for a channel whose version moved,
|
||||
# so snapshotting a channel this step did not write needs a
|
||||
# bump: exit mode reaching the cadence on a superstep that
|
||||
# skipped the channel, and a fork's first checkpoint.
|
||||
if get_next_version is not None and (
|
||||
updated_channels is None or k not in updated_channels
|
||||
):
|
||||
channel_versions[k] = get_next_version(channel_versions[k], None)
|
||||
old = channel_versions[k]
|
||||
channel_versions[k] = get_next_version(old, None)
|
||||
bumped[k] = (old, channel_versions[k])
|
||||
values[k] = _DeltaSnapshot(ch.get())
|
||||
else:
|
||||
v = ch.checkpoint()
|
||||
@@ -209,11 +231,30 @@ def create_checkpoint(
|
||||
id=id or str(uuid6(clock_seq=step)),
|
||||
channel_values=values,
|
||||
channel_versions=channel_versions,
|
||||
versions_seen=checkpoint["versions_seen"],
|
||||
versions_seen=_mark_bumps_seen(checkpoint["versions_seen"], bumped),
|
||||
updated_channels=None if updated_channels is None else sorted(updated_channels),
|
||||
)
|
||||
|
||||
|
||||
def _mark_bumps_seen(
|
||||
versions_seen: dict[str, ChannelVersions],
|
||||
bumped: Mapping[str, tuple[Any, Any]],
|
||||
) -> dict[str, ChannelVersions]:
|
||||
"""Advance whoever had seen a bumped channel's old version to the new one.
|
||||
|
||||
A bump that only stores a snapshot is not a write. Left unseen, it would
|
||||
re-fire `interrupt_before` and rerun the channel's subscribers.
|
||||
"""
|
||||
if not bumped:
|
||||
return versions_seen
|
||||
out: dict[str, ChannelVersions] = {}
|
||||
for node, seen in {INTERRUPT: {}, **versions_seen}.items():
|
||||
advanced = {k: new for k, (old, new) in bumped.items() if seen.get(k) == old}
|
||||
if advanced or node in versions_seen:
|
||||
out[node] = {**seen, **advanced}
|
||||
return out
|
||||
|
||||
|
||||
def _needs_replay(spec: BaseChannel, stored: object) -> bool:
|
||||
"""True if `spec` is a `DeltaChannel` and no value is stored at this
|
||||
checkpoint, requiring an ancestor walk to reconstruct.
|
||||
|
||||
@@ -102,6 +102,7 @@ from langgraph.pregel._checkpoint import (
|
||||
copy_checkpoint,
|
||||
create_checkpoint,
|
||||
delta_channels_to_snapshot,
|
||||
delta_channels_with_pending_writes,
|
||||
empty_checkpoint,
|
||||
exit_delta_task_id,
|
||||
)
|
||||
@@ -222,10 +223,13 @@ class PregelLoop:
|
||||
# under the saver's `ORDER BY task_id, idx` sorting.
|
||||
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
|
||||
|
||||
# Delta channels that saw an Overwrite since the last checkpoint. These
|
||||
# channels must snapshot after live update applies overwrite semantics so
|
||||
# sparse replay starts from the same post-overwrite value.
|
||||
_delta_channels_with_overwrite: set[str]
|
||||
# Delta channels that must snapshot at the next checkpoint, whatever their
|
||||
# cadence counters say:
|
||||
# * an Overwrite arrived since the last checkpoint, so sparse replay has to
|
||||
# start from the post-overwrite value;
|
||||
# * the checkpoint this run starts from has pending writes to them; see
|
||||
# `delta_channels_with_pending_writes`.
|
||||
_delta_channels_forced_snapshot: set[str]
|
||||
|
||||
# The checkpoint_config that points at the parent loaded at `__enter__`
|
||||
# (or the synthetic-empty checkpoint, on first run). We capture it
|
||||
@@ -683,7 +687,7 @@ class PregelLoop:
|
||||
def after_tick(self) -> None:
|
||||
# finish superstep
|
||||
writes = [w for t in self.tasks.values() for w in t.writes]
|
||||
self._delta_channels_with_overwrite.update(
|
||||
self._delta_channels_forced_snapshot.update(
|
||||
ch
|
||||
for ch, v in writes
|
||||
if isinstance(self.specs.get(ch), DeltaChannel) and _get_overwrite(v)[0]
|
||||
@@ -898,6 +902,15 @@ class PregelLoop:
|
||||
self.checkpoint_pending_writes = [
|
||||
w for w in self.checkpoint_pending_writes if w[1] != RESUME
|
||||
]
|
||||
# A resume that is not replaying reuses the head's pending writes
|
||||
# instead of rerunning their tasks, so none of them can leak.
|
||||
self._delta_channels_forced_snapshot = (
|
||||
set()
|
||||
if is_resuming and not self.is_replaying
|
||||
else delta_channels_with_pending_writes(
|
||||
self.specs, self.checkpoint_pending_writes
|
||||
)
|
||||
)
|
||||
|
||||
# map command to writes
|
||||
if input_is_command:
|
||||
@@ -991,7 +1004,7 @@ class PregelLoop:
|
||||
manager=None,
|
||||
updated_channels=updated_channels,
|
||||
)
|
||||
self._delta_channels_with_overwrite.update(
|
||||
self._delta_channels_forced_snapshot.update(
|
||||
c
|
||||
for c, v in input_writes
|
||||
if isinstance(self.specs.get(c), DeltaChannel) and _get_overwrite(v)[0]
|
||||
@@ -1136,7 +1149,7 @@ class PregelLoop:
|
||||
# create new checkpoint
|
||||
channels_to_snapshot = (
|
||||
delta_channels_to_snapshot(self.channels, new_counters)
|
||||
| self._delta_channels_with_overwrite
|
||||
| self._delta_channels_forced_snapshot
|
||||
if do_checkpoint
|
||||
else set()
|
||||
)
|
||||
@@ -1154,7 +1167,7 @@ class PregelLoop:
|
||||
for k in channels_to_snapshot:
|
||||
new_counters[k] = (0, 0)
|
||||
if do_checkpoint:
|
||||
self._delta_channels_with_overwrite.difference_update(channels_to_snapshot)
|
||||
self._delta_channels_forced_snapshot.difference_update(channels_to_snapshot)
|
||||
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
|
||||
if non_zero:
|
||||
self.checkpoint_metadata["counters_since_delta_snapshot"] = non_zero
|
||||
@@ -1239,7 +1252,7 @@ class PregelLoop:
|
||||
)
|
||||
channels_to_snapshot = (
|
||||
delta_channels_to_snapshot(self.channels, counters)
|
||||
| self._delta_channels_with_overwrite
|
||||
| self._delta_channels_forced_snapshot
|
||||
)
|
||||
|
||||
pending = [
|
||||
@@ -1684,7 +1697,6 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
)
|
||||
self._delta_write_futs = []
|
||||
self._error_handler_write_futs = []
|
||||
self._delta_channels_with_overwrite = set()
|
||||
self._exit_delta_writes = (
|
||||
[] if self.durability == "exit" and self.checkpointer is not None else None
|
||||
)
|
||||
@@ -1942,7 +1954,6 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
)
|
||||
self._delta_write_futs = []
|
||||
self._error_handler_write_futs = []
|
||||
self._delta_channels_with_overwrite = set()
|
||||
self._exit_delta_writes = (
|
||||
[] if self.durability == "exit" and self.checkpointer is not None else None
|
||||
)
|
||||
|
||||
@@ -133,6 +133,7 @@ from langgraph.pregel._checkpoint import (
|
||||
copy_checkpoint,
|
||||
create_checkpoint,
|
||||
create_checkpoint_plan_for_update_state_api,
|
||||
delta_channels_with_pending_writes,
|
||||
empty_checkpoint,
|
||||
get_updated_channels_from_tasks,
|
||||
)
|
||||
@@ -1637,12 +1638,21 @@ class Pregel(
|
||||
else:
|
||||
raise ValueError(f"Subgraph {recast} not found")
|
||||
|
||||
# Taken from the first superstep's base, and cleared by the first
|
||||
# checkpoint that carries the snapshots, which `__copy__` does not write.
|
||||
fork_pending: set[str] | None = None
|
||||
|
||||
def perform_superstep(
|
||||
input_config: RunnableConfig, updates: Sequence[StateUpdate]
|
||||
) -> RunnableConfig:
|
||||
nonlocal fork_pending
|
||||
# get last checkpoint
|
||||
config = ensure_config(self.config, input_config)
|
||||
saved = checkpointer.get_tuple(config)
|
||||
if fork_pending is None:
|
||||
fork_pending = delta_channels_with_pending_writes(
|
||||
self.channels, saved.pending_writes if saved else None
|
||||
)
|
||||
if saved is not None:
|
||||
self._migrate_checkpoint(saved.checkpoint)
|
||||
checkpoint = (
|
||||
@@ -1726,9 +1736,17 @@ class Pregel(
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
# save checkpoint
|
||||
next_checkpoint = create_checkpoint(
|
||||
checkpoint,
|
||||
channels,
|
||||
step,
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=fork_pending,
|
||||
)
|
||||
fork_pending.difference_update(next_checkpoint["channel_values"])
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
create_checkpoint(checkpoint, channels, step),
|
||||
next_checkpoint,
|
||||
{
|
||||
"source": "update",
|
||||
"step": step + 1,
|
||||
@@ -1736,7 +1754,7 @@ class Pregel(
|
||||
},
|
||||
get_new_channel_versions(
|
||||
checkpoint_previous_versions,
|
||||
checkpoint["channel_versions"],
|
||||
next_checkpoint["channel_versions"],
|
||||
),
|
||||
)
|
||||
return patch_checkpoint_map(
|
||||
@@ -1765,9 +1783,17 @@ class Pregel(
|
||||
if saved and saved.metadata.get("step") is not None
|
||||
else -1
|
||||
)
|
||||
next_checkpoint = create_checkpoint(
|
||||
checkpoint,
|
||||
channels,
|
||||
next_step,
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=fork_pending,
|
||||
)
|
||||
fork_pending.difference_update(next_checkpoint["channel_values"])
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
create_checkpoint(checkpoint, channels, next_step),
|
||||
next_checkpoint,
|
||||
{
|
||||
"source": "input",
|
||||
"step": next_step,
|
||||
@@ -1777,7 +1803,7 @@ class Pregel(
|
||||
},
|
||||
get_new_channel_versions(
|
||||
checkpoint_previous_versions,
|
||||
checkpoint["channel_versions"],
|
||||
next_checkpoint["channel_versions"],
|
||||
),
|
||||
)
|
||||
|
||||
@@ -2020,18 +2046,19 @@ class Pregel(
|
||||
parents=saved.metadata.get("parents", {}) if saved else {},
|
||||
saved_metadata=saved.metadata if saved else None,
|
||||
is_fresh_thread=saved is None,
|
||||
fork_channels=fork_pending,
|
||||
)
|
||||
)
|
||||
checkpoint = create_checkpoint(
|
||||
checkpoint,
|
||||
channels,
|
||||
step + 1,
|
||||
updated_channels=updated_channels if channels_to_snapshot else None,
|
||||
get_next_version=checkpointer.get_next_version
|
||||
if channels_to_snapshot
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
)
|
||||
fork_pending.difference_update(checkpoint["channel_values"])
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
checkpoint,
|
||||
@@ -2103,12 +2130,21 @@ class Pregel(
|
||||
else:
|
||||
raise ValueError(f"Subgraph {recast} not found")
|
||||
|
||||
# Taken from the first superstep's base, and cleared by the first
|
||||
# checkpoint that carries the snapshots, which `__copy__` does not write.
|
||||
fork_pending: set[str] | None = None
|
||||
|
||||
async def aperform_superstep(
|
||||
input_config: RunnableConfig, updates: Sequence[StateUpdate]
|
||||
) -> RunnableConfig:
|
||||
nonlocal fork_pending
|
||||
# get last checkpoint
|
||||
config = ensure_config(self.config, input_config)
|
||||
saved = await checkpointer.aget_tuple(config)
|
||||
if fork_pending is None:
|
||||
fork_pending = delta_channels_with_pending_writes(
|
||||
self.channels, saved.pending_writes if saved else None
|
||||
)
|
||||
if saved is not None:
|
||||
self._migrate_checkpoint(saved.checkpoint)
|
||||
checkpoint = (
|
||||
@@ -2190,16 +2226,25 @@ class Pregel(
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
# save checkpoint
|
||||
next_checkpoint = create_checkpoint(
|
||||
checkpoint,
|
||||
channels,
|
||||
step,
|
||||
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(
|
||||
checkpoint_config,
|
||||
create_checkpoint(checkpoint, channels, step),
|
||||
next_checkpoint,
|
||||
{
|
||||
"source": "update",
|
||||
"step": step + 1,
|
||||
"parents": saved.metadata.get("parents", {}) if saved else {},
|
||||
},
|
||||
get_new_channel_versions(
|
||||
checkpoint_previous_versions, checkpoint["channel_versions"]
|
||||
checkpoint_previous_versions,
|
||||
next_checkpoint["channel_versions"],
|
||||
),
|
||||
)
|
||||
return patch_checkpoint_map(
|
||||
@@ -2228,9 +2273,17 @@ class Pregel(
|
||||
if saved and saved.metadata.get("step") is not None
|
||||
else -1
|
||||
)
|
||||
next_checkpoint = create_checkpoint(
|
||||
checkpoint,
|
||||
channels,
|
||||
next_step,
|
||||
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(
|
||||
checkpoint_config,
|
||||
create_checkpoint(checkpoint, channels, next_step),
|
||||
next_checkpoint,
|
||||
{
|
||||
"source": "input",
|
||||
"step": next_step,
|
||||
@@ -2240,7 +2293,7 @@ class Pregel(
|
||||
},
|
||||
get_new_channel_versions(
|
||||
checkpoint_previous_versions,
|
||||
checkpoint["channel_versions"],
|
||||
next_checkpoint["channel_versions"],
|
||||
),
|
||||
)
|
||||
|
||||
@@ -2480,18 +2533,19 @@ class Pregel(
|
||||
parents=saved.metadata.get("parents", {}) if saved else {},
|
||||
saved_metadata=saved.metadata if saved else None,
|
||||
is_fresh_thread=saved is None,
|
||||
fork_channels=fork_pending,
|
||||
)
|
||||
)
|
||||
checkpoint = create_checkpoint(
|
||||
checkpoint,
|
||||
channels,
|
||||
step + 1,
|
||||
updated_channels=updated_channels if channels_to_snapshot else None,
|
||||
get_next_version=checkpointer.get_next_version
|
||||
if channels_to_snapshot
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
)
|
||||
fork_pending.difference_update(checkpoint["channel_values"])
|
||||
next_config = await checkpointer.aput(
|
||||
checkpoint_config,
|
||||
checkpoint,
|
||||
|
||||
@@ -85,11 +85,13 @@ class MemorySaverAssertImmutable(InMemorySaver):
|
||||
)
|
||||
== saved
|
||||
), config["configurable"]["checkpoint_ns"]
|
||||
next_config = super().put(config, checkpoint, metadata, new_versions)
|
||||
# Read back, not the object handed in: a DeltaChannel a step did not
|
||||
# write is refilled on read from the blob its inherited version points at.
|
||||
self.storage_for_copies[thread_id][checkpoint_ns][checkpoint["id"]] = (
|
||||
self.serde.dumps_typed(checkpoint)
|
||||
self.serde.dumps_typed(super().get(next_config))
|
||||
)
|
||||
# call super to write checkpoint
|
||||
return super().put(config, checkpoint, metadata, new_versions)
|
||||
return next_config
|
||||
|
||||
|
||||
class MemorySaverNoPending(InMemorySaver):
|
||||
|
||||
@@ -0,0 +1,533 @@
|
||||
"""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
|
||||
values; the plain channel needs no replay, so it is the oracle.
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from operator import add
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph._internal._constants import INPUT
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
from langgraph.types import Command, Durability, StateSnapshot, StateUpdate, interrupt
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
def _append(current: list | None, writes: Sequence[Any]) -> list:
|
||||
out = list(current or [])
|
||||
for write in writes:
|
||||
out.extend(write if isinstance(write, list) else [write])
|
||||
return out
|
||||
|
||||
|
||||
class _State(TypedDict):
|
||||
log: Annotated[list, DeltaChannel(_append, snapshot_frequency=1000)]
|
||||
plain: Annotated[list, add]
|
||||
other: Annotated[list, add]
|
||||
|
||||
|
||||
def _build(checkpointer: BaseCheckpointSaver, tag: str) -> Any:
|
||||
def node(state: _State) -> dict:
|
||||
return {"log": [f"{tag}-out"], "plain": [f"{tag}-out"]}
|
||||
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("n", node)
|
||||
builder.set_entry_point("n")
|
||||
builder.set_finish_point("n")
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
|
||||
|
||||
def _build_without_delta_writes(checkpointer: BaseCheckpointSaver, tag: str) -> Any:
|
||||
def node(state: _State) -> dict:
|
||||
return {"other": [f"{tag}-other"]}
|
||||
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("n", node)
|
||||
builder.set_entry_point("n")
|
||||
builder.set_finish_point("n")
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
|
||||
|
||||
def _thread(thread_id: str) -> RunnableConfig:
|
||||
return {"configurable": {"thread_id": thread_id}}
|
||||
|
||||
|
||||
def _at(config: RunnableConfig, snapshot: StateSnapshot) -> RunnableConfig:
|
||||
return {
|
||||
"configurable": {
|
||||
**config["configurable"],
|
||||
"checkpoint_ns": "",
|
||||
"checkpoint_id": snapshot.config["configurable"]["checkpoint_id"],
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def _both(marker: str) -> dict:
|
||||
return {"log": [marker], "plain": [marker]}
|
||||
|
||||
|
||||
def _snapshotted_checkpoints(
|
||||
checkpointer: BaseCheckpointSaver, config: RunnableConfig
|
||||
) -> list[str]:
|
||||
return [
|
||||
tuple_.config["configurable"]["checkpoint_id"]
|
||||
for tuple_ in checkpointer.list(config)
|
||||
if isinstance(tuple_.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
|
||||
]
|
||||
|
||||
|
||||
def _assert_fork_is_clean(state: StateSnapshot, abandoned: str) -> None:
|
||||
assert state.values["log"] == state.values["plain"], (
|
||||
f"delta channel diverged from the plain channel: "
|
||||
f"{state.values['log']} != {state.values['plain']}"
|
||||
)
|
||||
assert abandoned not in state.values["log"], (
|
||||
f"{abandoned!r} belongs to the branch the fork replaced, "
|
||||
f"but was replayed into {state.values['log']}"
|
||||
)
|
||||
|
||||
|
||||
def test_fork_by_invoke(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
_build(sync_checkpointer, "first").invoke(
|
||||
_both("in-1"), config, durability=durability
|
||||
)
|
||||
graph = _build(sync_checkpointer, "second")
|
||||
graph.invoke(_both("in-2"), config, durability=durability)
|
||||
abandoned_head = graph.get_state(config)
|
||||
|
||||
base = next(
|
||||
snapshot
|
||||
for snapshot in graph.get_state_history(config)
|
||||
if "in-2" not in snapshot.values["log"]
|
||||
)
|
||||
_build(sync_checkpointer, "third").invoke(
|
||||
_both("in-3"), _at(config, base), durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
_assert_fork_is_clean(state, "in-2")
|
||||
assert state.values["log"] == [*base.values["log"], "in-3", "third-out"]
|
||||
|
||||
abandoned = graph.get_state(abandoned_head.config).values
|
||||
assert abandoned["log"] == abandoned["plain"] == abandoned_head.values["log"]
|
||||
|
||||
|
||||
async def test_afork_by_invoke(
|
||||
async_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
await _build(async_checkpointer, "first").ainvoke(
|
||||
_both("in-1"), config, durability=durability
|
||||
)
|
||||
graph = _build(async_checkpointer, "second")
|
||||
await graph.ainvoke(_both("in-2"), config, durability=durability)
|
||||
abandoned_head = 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"]
|
||||
)
|
||||
await _build(async_checkpointer, "third").ainvoke(
|
||||
_both("in-3"), _at(config, base), durability=durability
|
||||
)
|
||||
|
||||
state = await graph.aget_state(config)
|
||||
_assert_fork_is_clean(state, "in-2")
|
||||
assert state.values["log"] == [*base.values["log"], "in-3", "third-out"]
|
||||
|
||||
abandoned = (await graph.aget_state(abandoned_head.config)).values
|
||||
assert abandoned["log"] == abandoned["plain"] == abandoned_head.values["log"]
|
||||
|
||||
|
||||
def test_fork_off_checkpoint_before_first_input(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(sync_checkpointer, "first")
|
||||
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(
|
||||
_both("in-9"), _at(config, root), durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
_assert_fork_is_clean(state, "in-1")
|
||||
assert state.values["log"] == ["in-9", "third-out"]
|
||||
|
||||
|
||||
async def test_afork_off_checkpoint_before_first_input(
|
||||
async_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(async_checkpointer, "first")
|
||||
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(
|
||||
_both("in-9"), _at(config, root), durability=durability
|
||||
)
|
||||
|
||||
state = await graph.aget_state(config)
|
||||
_assert_fork_is_clean(state, "in-1")
|
||||
assert state.values["log"] == ["in-9", "third-out"]
|
||||
|
||||
|
||||
def test_fork_by_update_state(sync_checkpointer: BaseCheckpointSaver) -> None:
|
||||
config = _thread("t")
|
||||
_build(sync_checkpointer, "first").invoke(_both("in-1"), config)
|
||||
graph = _build(sync_checkpointer, "second")
|
||||
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), _both("patched"))
|
||||
|
||||
state = graph.get_state(forked)
|
||||
_assert_fork_is_clean(state, "in-2")
|
||||
assert state.values["log"] == [*base.values["log"], "patched"]
|
||||
|
||||
|
||||
async def test_afork_by_update_state(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
await _build(async_checkpointer, "first").ainvoke(_both("in-1"), config)
|
||||
graph = _build(async_checkpointer, "second")
|
||||
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), _both("patched"))
|
||||
|
||||
state = await graph.aget_state(forked)
|
||||
_assert_fork_is_clean(state, "in-2")
|
||||
assert state.values["log"] == [*base.values["log"], "patched"]
|
||||
|
||||
|
||||
def test_unaddressed_run_keeps_snapshot_cadence(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(sync_checkpointer, "first")
|
||||
graph.invoke(_both("in-1"), config, durability=durability)
|
||||
graph.invoke(_both("in-2"), config, durability=durability)
|
||||
|
||||
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||
|
||||
|
||||
def test_fork_before_first_value_when_fork_never_writes_the_channel(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(sync_checkpointer, "first")
|
||||
graph.invoke(_both("in-1"), config, durability=durability)
|
||||
|
||||
root = list(graph.get_state_history(config))[-1]
|
||||
assert root.values["log"] == []
|
||||
|
||||
_build_without_delta_writes(sync_checkpointer, "third").invoke(
|
||||
{"other": ["in-9"]}, _at(config, root), durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
_assert_fork_is_clean(state, "in-1")
|
||||
assert state.values["log"] == []
|
||||
|
||||
|
||||
async def test_afork_before_first_value_when_fork_never_writes_the_channel(
|
||||
async_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(async_checkpointer, "first")
|
||||
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_without_delta_writes(async_checkpointer, "third").ainvoke(
|
||||
{"other": ["in-9"]}, _at(config, root), durability=durability
|
||||
)
|
||||
|
||||
state = await graph.aget_state(config)
|
||||
_assert_fork_is_clean(state, "in-1")
|
||||
assert state.values["log"] == []
|
||||
|
||||
|
||||
def test_fork_before_first_value_by_bulk_update(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(sync_checkpointer, "first")
|
||||
graph.invoke(_both("in-1"), config)
|
||||
|
||||
root = list(graph.get_state_history(config))[-1]
|
||||
assert root.values["log"] == []
|
||||
|
||||
forked = graph.bulk_update_state(
|
||||
_at(config, root),
|
||||
[
|
||||
[StateUpdate({"other": ["s1"]}, "n")],
|
||||
[StateUpdate(_both("s2"), "n")],
|
||||
],
|
||||
)
|
||||
|
||||
state = graph.get_state(forked)
|
||||
_assert_fork_is_clean(state, "in-1")
|
||||
assert state.values["log"] == ["s2"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("first_as_node", [INPUT, END, "__copy__"])
|
||||
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(_both("in-1"), config)
|
||||
graph = _build(sync_checkpointer, "second")
|
||||
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"]
|
||||
)
|
||||
first = (
|
||||
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(_both("second-step"), "n")]],
|
||||
)
|
||||
|
||||
state = graph.get_state(forked)
|
||||
assert state.values["log"] == state.values["plain"], (
|
||||
f"delta channel diverged from the plain channel: "
|
||||
f"{state.values['log']} != {state.values['plain']}"
|
||||
)
|
||||
|
||||
|
||||
def test_unaddressed_bulk_update_keeps_snapshot_cadence(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(sync_checkpointer, "first")
|
||||
graph.invoke(_both("in-1"), config)
|
||||
|
||||
graph.bulk_update_state(
|
||||
config,
|
||||
[[StateUpdate(_both(f"u{i}"), "n")] for i in range(4)],
|
||||
)
|
||||
|
||||
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||
|
||||
|
||||
def _build_paused_before_b(checkpointer: BaseCheckpointSaver) -> Any:
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("a", lambda state: _both("a"))
|
||||
builder.add_node("b", lambda state: _both("b"))
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge("a", "b")
|
||||
builder.add_edge("b", END)
|
||||
return builder.compile(checkpointer=checkpointer, interrupt_before=["b"])
|
||||
|
||||
|
||||
def _build_parallel_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
|
||||
def ask(state: _State) -> dict:
|
||||
interrupt("approve?")
|
||||
return {"other": ["q"]}
|
||||
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("p", lambda state: _both("p"))
|
||||
builder.add_node("q", ask)
|
||||
builder.add_edge(START, "p")
|
||||
builder.add_edge(START, "q")
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
|
||||
|
||||
def test_resume_at_interrupt_before_with_the_head_checkpoint_id_runs_the_node(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build_paused_before_b(sync_checkpointer)
|
||||
graph.invoke(_both("in"), config, durability=durability)
|
||||
|
||||
graph.invoke(None, graph.get_state(config).config, durability=durability)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.next == (), f"resume paused again before {state.next}"
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
|
||||
|
||||
|
||||
async def test_aresume_at_interrupt_before_with_the_head_checkpoint_id_runs_the_node(
|
||||
async_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build_paused_before_b(async_checkpointer)
|
||||
await graph.ainvoke(_both("in"), config, durability=durability)
|
||||
|
||||
await graph.ainvoke(
|
||||
None, (await graph.aget_state(config)).config, durability=durability
|
||||
)
|
||||
|
||||
state = await graph.aget_state(config)
|
||||
assert state.next == (), f"resume paused again before {state.next}"
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
|
||||
|
||||
|
||||
def test_replay_from_a_paused_checkpoint_runs_the_node_once(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build_paused_before_b(sync_checkpointer)
|
||||
graph.invoke(_both("in"), config)
|
||||
paused = graph.get_state(config).config
|
||||
graph.invoke(None, config)
|
||||
|
||||
graph.invoke(None, paused)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.next == (), f"replay paused again before {state.next}"
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("addressed", [False, True])
|
||||
def test_new_input_on_an_interrupted_head_does_not_replay_its_pending_writes(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability, addressed: bool
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build_parallel_interrupt(sync_checkpointer)
|
||||
graph.invoke(_both("in-1"), config, durability=durability)
|
||||
head = graph.get_state(config).config
|
||||
|
||||
graph.invoke(_both("in-2"), head if addressed else config, durability=durability)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.values["log"] == state.values["plain"] == ["in-1", "in-2", "p"]
|
||||
|
||||
|
||||
def _build_deferred_after_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("a", lambda state: _both("a"))
|
||||
builder.add_node("b", lambda state: _both("b"), defer=True)
|
||||
builder.add_node("c", lambda state: {})
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge("a", "b")
|
||||
builder.add_edge("a", "c")
|
||||
return builder.compile(checkpointer=checkpointer, interrupt_after=["a"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"durability",
|
||||
[
|
||||
"sync",
|
||||
"async",
|
||||
pytest.param(
|
||||
"exit",
|
||||
marks=pytest.mark.xfail(
|
||||
reason="exit durability stores a resumed run's loaded writes twice",
|
||||
strict=True,
|
||||
),
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build_parallel_interrupt(sync_checkpointer)
|
||||
graph.invoke(_both("in-1"), config, durability=durability)
|
||||
|
||||
graph.invoke(Command(resume="yes"), config, durability=durability)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.next == ()
|
||||
assert state.values["log"] == state.values["plain"] == ["in-1", "p"]
|
||||
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||
|
||||
|
||||
def test_resume_addressed_at_an_interrupted_head_reruns_its_tasks_once(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build_parallel_interrupt(sync_checkpointer)
|
||||
graph.invoke(_both("in-1"), config, durability=durability)
|
||||
|
||||
graph.invoke(
|
||||
Command(resume="yes"), graph.get_state(config).config, durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.next == ()
|
||||
assert state.values["log"] == state.values["plain"] == ["in-1", "p"]
|
||||
|
||||
|
||||
def test_update_state_with_the_head_checkpoint_id_keeps_a_deferred_node(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_deferred_after_interrupt(sync_checkpointer)
|
||||
config = _thread("t")
|
||||
graph.invoke(_both("in"), config)
|
||||
|
||||
graph.update_state(graph.get_state(config).config, _both("u"), as_node="c")
|
||||
graph.invoke(None, config)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.next == (), f"deferred node never ran, still pending: {state.next}"
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "a", "u", "b"]
|
||||
|
||||
|
||||
async def test_aupdate_state_with_the_head_checkpoint_id_keeps_a_deferred_node(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_deferred_after_interrupt(async_checkpointer)
|
||||
config = _thread("t")
|
||||
await graph.ainvoke(_both("in"), config)
|
||||
|
||||
await graph.aupdate_state(
|
||||
(await graph.aget_state(config)).config, _both("u"), as_node="c"
|
||||
)
|
||||
await graph.ainvoke(None, config)
|
||||
|
||||
state = await graph.aget_state(config)
|
||||
assert state.next == (), f"deferred node never ran, still pending: {state.next}"
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "a", "u", "b"]
|
||||
|
||||
|
||||
def test_turns_addressed_at_the_head_store_no_snapshot(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
config = _thread("t")
|
||||
graph = _build(sync_checkpointer, "turn")
|
||||
graph.invoke(_both("in-1"), config)
|
||||
for turn in range(2, 5):
|
||||
graph.invoke(_both(f"in-{turn}"), graph.get_state(config).config)
|
||||
|
||||
assert not _snapshotted_checkpoints(sync_checkpointer, config)
|
||||
assert (
|
||||
graph.get_state(config).values["log"] == graph.get_state(config).values["plain"]
|
||||
)
|
||||
@@ -338,3 +338,29 @@ def test_state_history_chain_after_fresh_update_state_delta_channel() -> None:
|
||||
assert update_snapshot.metadata["step"] == 0
|
||||
assert update_snapshot.parent_config is None
|
||||
assert [m.content for m in update_snapshot.values["messages"]] == ["hello"]
|
||||
|
||||
|
||||
def test_update_state_that_snapshots_keeps_a_deferred_node_pending() -> None:
|
||||
channel = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1)
|
||||
|
||||
class State(TypedDict):
|
||||
messages: Annotated[list, channel]
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("a", lambda state: {"messages": [HumanMessage("a", id="a")]})
|
||||
builder.add_node(
|
||||
"b", lambda state: {"messages": [HumanMessage("b", id="b")]}, defer=True
|
||||
)
|
||||
builder.add_node("c", lambda state: {})
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge("a", "b")
|
||||
builder.add_edge("a", "c")
|
||||
graph = builder.compile(checkpointer=InMemorySaver(), interrupt_after=["a"])
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"messages": [HumanMessage("s", id="s")]}, config)
|
||||
|
||||
graph.update_state(config, {"messages": [HumanMessage("u", id="u")]}, as_node="c")
|
||||
final = graph.invoke(None, config)
|
||||
|
||||
assert [m.content for m in final["messages"]] == ["s", "a", "u", "b"]
|
||||
assert graph.get_state(config).next == ()
|
||||
|
||||
Reference in New Issue
Block a user