Compare commits

..
Author SHA1 Message Date
Elior Nataf Lackritz 79781ccfbe 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 16:41:56 -04:00
Elior Nataf Lackritz c0405d246a fix(langgraph): bump a snapshot whose version hasn't moved since the last stored checkpoint
An exit-durability replay that stops in its first tick saves the fork
checkpoint built in `_first`, and its `updated_channels` are still the loaded
checkpoint's. The seal took that as a version move and skipped the bump, so
`put` stored no blob for it: the in-memory saver dropped the snapshot and the
fork replayed the abandoned branch, and the Postgres saver kept its inline
marker, so the channel read back as `True`.

`create_checkpoint` now takes the versions of the last stored checkpoint and
bumps a snapshotted channel whenever its version hasn't moved since then.
2026-10-01 16:15:44 -04:00
Elior Nataf Lackritz e4d77cd222 fix(langgraph): seal what a resume drops, and keep storage-only bumps out of scheduling
- A resume that reapplies the head's pending writes now seals the delta
  channels of the loaded writes no task of the run claimed, decided in
  `after_tick` once the tasks are known. `Command(resume=..., goto=[Send(...)])`
  replaces the fan-out, so the finished task's writes stayed on the head and
  the reload replayed them.
- One predicate, `_reapplies_pending_writes`, answers whether the loaded
  writes go back to their tasks, for every site that reapplies them.
- The snapshot cadence skips delta channels without a version. Past
  DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT, every checkpoint minted a version and
  an empty snapshot for each never-written delta channel.
- Versions minted only to store a snapshot are recorded under
  `SNAPSHOT_BUMPS` in `versions_seen`, and `update_state` skips them when it
  infers `as_node`. Advancing `versions_seen` over a bump made a raw `Pregel`
  subscriber look like the last writer after an exit-durability run.
- `_mark_bumps_seen` no longer synthesizes an `INTERRUPT` entry.
2026-10-01 15:14:22 -04:00
Elior Nataf Lackritz 858e55f232 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.
2026-09-30 15:07:17 -04:00
Elior Nataf Lackritz e59ecbc23d revert(langgraph): leave exit-mode resume writes as they are on main
Skipping the loaded writes in the exit accumulator removed the duplicate
replay but reordered it: the loaded write keeps its task id while the
run's later writes get step-prefixed ids that sort first. That ordering
problem belongs to exit mode, not to forking, and exists on main without
a fork, so it gets its own fix. The resume test marks exit durability as
an expected failure until then.
2026-09-29 11:39:36 -04:00
Elior Nataf Lackritz 719a4d71bc fix(langgraph): skip the seal on a resume that reuses the head's writes
A resume that is not replaying reuses the head's pending writes instead of
rerunning their tasks, so none of them can leak; deciding in `_first`,
where that is known, stops a plain `Command(resume=...)` from storing a
snapshot. A replaying resume reruns the tasks, so it still seals.

Exit mode also re-recorded the writes a resume loaded from the head: they
were already stored there, and every later read replayed them twice. The
exit accumulator now skips them.
2026-09-29 09:42:55 -04:00
Elior Nataf Lackritz ddaf708cd0 fix(langgraph): snapshot only what a fork can leak, and hide the bump
A fork used to snapshot every DeltaChannel whenever the caller passed a
checkpoint_id. That fired on every turn a client addresses the head
(storing a full copy of the channel per turn), missed new input sent to an
interrupted head without an id, and its storage-only version bump read as
a real write: interrupt_before fired again on resume, and a replay paused
at a node it had already passed.

Snapshot the delta channels the base checkpoint has pending writes for,
since only those can leak into a branch that does not consume them, and
advance versions_seen past any bump that only stores a snapshot, including
the interrupt tracker. update_state no longer records its narrower
updated_channels when it snapshots, so a deferred node listed in next still
runs on resume (#9089).
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 9d5b0f1991 refactor(langgraph): drop the deferred fork-snapshot queue
Now that a forced snapshot mints a version for a never-written channel,
the first checkpoint a forked run writes seals every delta channel, so
_delta_channels_awaiting_fork_snapshot never outlived it. Seed
_delta_channels_forced_snapshot directly at loop construction instead.

Also asserts that forking by invoke leaves the abandoned branch's reads
intact, and trims comments and test docstrings.
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 9d16b52955 fix(langgraph): seal a fork whose delta channel has no value yet
A DeltaChannel that was never written on the branch being forked has no
value to snapshot and no entry in channel_versions, so create_checkpoint
skipped it and the fork's first checkpoint recorded no boundary at all.
The walk then ran past the fork into the shared base and collected the
abandoned branch's writes, the same failure this branch already fixes for
channels that do have a value.

Two shapes leaked. A run forking off a checkpoint older than the channel's
first value and never writing that channel returned ['in-1'] where the
plain-channel oracle returned []. A bulk update writing the delta key only
in its second superstep returned ['in-1', 's2'] against ['s2'].

No new blob type is needed. _DeltaSnapshot already carries the value and
is already serialized by every saver, and from_checkpoint turns MISSING
into typ(), so _DeltaSnapshot(typ()) reconstructs to the same empty value
the channel would have had. What was missing is a version: without one,
put drops the blob as not-a-new-version, so mint a first one.

Deferring the seal to a later superstep does not work. That superstep
reconstructs through the still-unsealed checkpoint and would only bake the
corrupted value into its own snapshot.

Checked that minting a version does not fire nodes that subscribe to the
channel: a raw Pregel node subscribed directly to the delta channel stays
silent across the fork.

Reported by the Open SWE review bot on #8548.
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 2de0c47c1f test(langgraph): compare read-back checkpoints in the immutability saver
MemorySaverAssertImmutable recorded the checkpoint object handed to put,
then compared it against one read back through get. Those two are not the
same shape: channel_values are stored per (channel, version), so a channel
a step did not write is refilled from the blob its inherited version still
points at.

Every channel except DeltaChannel writes its value into channel_values on
every checkpoint, so the two agreed by accident. A DeltaChannel stores
nothing except at a snapshot, so once one snapshots and a later step does
not write it, the saver reports a checkpoint that changed after it was
written when nothing was mutated.

Reproducible on main with no fork involved: a delta channel with
snapshot_frequency=1 written by the first node and left alone by the next
two trips the assertion. Existing delta tests miss it only because they
all use snapshot_frequency=1000.

Record what the saver reads back instead. Comparing read-back against
read-back still catches a checkpoint whose stored data really changed.
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 377083220e fix(langgraph): seal a fork on the first checkpoint it writes
The as_node INPUT, END and __copy__ paths write a checkpoint and return
before create_checkpoint_plan_for_update_state_api runs, so a bulk update
whose first superstep took one of them left the branch unsealed. Only
INPUT actually leaked: END absorbs the base's already-run task writes, so
its delta and plain channels agree.

Sealing on a later superstep does not help. By then that superstep has
reconstructed its value by walking through the unsealed checkpoint into
the shared base, so it snapshots an already-corrupted list. The fork's
first checkpoint is the one that has to carry the blob, which is what
create_fork_checkpoint does.

That snapshot was still being dropped by put: these paths apply writes to
the input channel, not the delta channel, so nothing bumped the delta
channel's version and it never entered new_versions. Pass get_next_version
for the manual bump, the same reason exit mode needs it, and derive
new_versions from the returned checkpoint.

fork_pending tracks what is still owed, mirroring
_delta_channels_awaiting_fork_snapshot in _loop.py.

Caught by the Open SWE review bot on #8548.
2026-09-28 19:54:23 -04:00
Elior Nataf Lackritz 9c5914861b fix(langgraph): only fork on the first superstep of a bulk update
perform_superstep returns the config of the checkpoint it just wrote and
bulk_update_state feeds that back in, so from the second superstep on the
incoming config always names a checkpoint whether or not the caller
addressed one. Deriving the fork flag from it made every superstep after
the first force-snapshot every available DeltaChannel and reset its
cadence, storing the whole growing value once per superstep.

Resolve the flag once from the caller's config and pass it explicitly,
true only for the first superstep. The clear-tasks recursion carries it
through, since the checkpoint written there has no delta snapshot and so
leaves a fork unsealed.

Caught by the Open SWE review bot on #8548.
2026-09-28 19:54:23 -04:00
24cf33f348 fix(langgraph): don't replay an abandoned branch into a DeltaChannel fork
Addressing an older checkpoint creates a fork: the shared base ends up
with two children and keeps the checkpoint_writes of the branch the fork
abandons. Nothing records which child consumed which write, so the
DeltaChannel ancestor walk collected the abandoned branch's writes too.
Live execution was correct; only the reconstruction after a reload was
wrong, and it was wrong on every saver.

Fixed on the write side, so no saver changes are needed. A run launched
against an explicitly addressed checkpoint forces every DeltaChannel to
snapshot into its first checkpoint, terminating the walk inside the fork
instead of at the shared base. This mirrors the existing force-snapshot
for Overwrite writes, hence the rename to _delta_channels_forced_snapshot.
update_state against an older checkpoint takes the same path, for the
same reason is_fresh_thread already does.

A channel with no value at the fork base cannot carry a snapshot blob
yet, so the request stays queued until the first superstep that gives it
one. Cost is one snapshot per addressed run, not per superstep.

Fixes #8443

Co-Authored-By: AnnaSuSu <64579968+AnnaSuSu@users.noreply.github.com>
Co-Authored-By: UditDewan <194863456+UditDewan@users.noreply.github.com>
2026-09-28 19:54:23 -04:00
14 changed files with 1279 additions and 821 deletions
@@ -23,6 +23,8 @@ RETURN = sys.intern("__return__")
# for writes of a task where we simply record the return value
PREVIOUS = sys.intern("__previous__")
# the implicit branch that handles each node's Control values
SNAPSHOT_BUMPS = sys.intern("__snapshot_bumps__")
# `versions_seen` key for channel versions minted only to store a snapshot
# --- Reserved cache namespaces ---
@@ -116,6 +118,7 @@ RESERVED = {
ERROR,
ERROR_SOURCE_NODE,
NO_WRITES,
SNAPSHOT_BUMPS,
# reserved config.configurable keys
CONFIG_KEY_SEND,
CONFIG_KEY_READ,
+105 -28
View File
@@ -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 PUSH, SNAPSHOT_BUMPS
from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel
from langgraph.channels.delta import DeltaChannel
@@ -50,17 +52,23 @@ def exit_delta_task_id(step: int, task_id: str) -> str:
def delta_channels_to_snapshot(
channels: Mapping[str, BaseChannel],
counters_since_delta_snapshot: Mapping[str, tuple[int, int]],
channel_versions: ChannelVersions,
) -> set[str]:
"""Return the set of DeltaChannel names that should snapshot now.
A channel snapshots when EITHER its accumulated update count reaches
`snapshot_frequency` OR the total supersteps since its last snapshot
reaches `DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT`. This is a pure
predicate — no mutation.
reaches `DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT`. A channel without a version
was never written on this branch, so it has nothing to snapshot. This is a
pure predicate — no mutation.
"""
result: set[str] = set()
for name, ch in channels.items():
if not isinstance(ch, DeltaChannel) or not ch.is_available():
if (
not isinstance(ch, DeltaChannel)
or not ch.is_available()
or name not in channel_versions
):
continue
updates, supersteps = counters_since_delta_snapshot.get(name, (0, 0))
if (
@@ -89,6 +97,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 +147,8 @@ 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],
channel_versions: ChannelVersions,
) -> tuple[set[str], dict[str, Any]]:
"""Return ``(channels_to_snapshot, metadata)`` for an update_state head."""
metadata: dict[str, Any] = {
@@ -137,7 +164,10 @@ 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, channel_versions)
| 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)}
@@ -155,6 +185,7 @@ def create_checkpoint(
updated_channels: set[str] | None = None,
get_next_version: GetNextVersion | None = None,
channels_to_snapshot: set[str] | None = None,
stored_versions: ChannelVersions | None = None,
) -> Checkpoint:
"""Build a new Checkpoint from the previous one and live channel state.
@@ -164,9 +195,15 @@ def create_checkpoint(
from `checkpoint_writes`. Callers compute the set via
`delta_channels_to_snapshot(channels, counters)`; defaults to empty
(no snapshots) when not provided.
`stored_versions` are the channel versions of the last checkpoint the
saver stored. When given, a snapshotted channel whose version has not
moved since then is bumped; otherwise `updated_channels` stands in for the
channels whose version moved.
"""
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 +211,30 @@ 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())
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).
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)
# `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.
unmoved = (
channel_versions[k] == stored_versions.get(k)
if stored_versions is not None
else updated_channels is None or k not in updated_channels
)
if get_next_version is not None and unmoved:
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 +246,51 @@ 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. The bumped
versions are also kept under `SNAPSHOT_BUMPS`, so inferring which node
wrote last can skip them.
"""
if not bumped:
return versions_seen
out = {
node: {
**seen,
**{k: new for k, (old, new) in bumped.items() if seen.get(k) == old},
}
for node, seen in versions_seen.items()
}
out[SNAPSHOT_BUMPS] = {
**versions_seen.get(SNAPSHOT_BUMPS, {}),
**{k: new for k, (_, new) in bumped.items()},
}
return out
def versions_seen_without_bumps(
versions_seen: dict[str, ChannelVersions],
) -> dict[str, ChannelVersions]:
"""`versions_seen` without the versions minted only to store a snapshot."""
if not (bumps := versions_seen.get(SNAPSHOT_BUMPS)):
return versions_seen
return {
node: {k: v for k, v in seen.items() if bumps.get(k) != v}
for node, seen in versions_seen.items()
if node != SNAPSHOT_BUMPS
}
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.
+95 -28
View File
@@ -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,
)
@@ -119,7 +120,6 @@ from langgraph.pregel._io import (
)
from langgraph.pregel._messages import ensure_message_ids
from langgraph.pregel._read import PregelNode
from langgraph.pregel._task_status import read_task_statuses
from langgraph.pregel._utils import get_new_channel_versions, is_xxh3_128_hexdigest
from langgraph.pregel.debug import (
map_debug_checkpoint,
@@ -223,10 +223,16 @@ 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]
# Set by `_first` for a resume whose loaded writes `after_tick` still has
# to check against the tasks that ran.
_seal_unclaimed_writes: bool = False
# The checkpoint_config that points at the parent loaded at `__enter__`
# (or the synthetic-empty checkpoint, on first run). We capture it
@@ -582,7 +588,7 @@ class PregelLoop:
# save the new task
self.tasks[pushed.id] = pushed
# match any pending writes to the new task
if not self.is_replaying:
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({pushed.id: pushed})
# return the new task, to be started if not run before
return pushed
@@ -660,7 +666,7 @@ class PregelLoop:
return False
# if there are pending writes from a previous loop, apply them
if not self.is_replaying and self.checkpoint_pending_writes:
if self._reapplies_pending_writes and self.checkpoint_pending_writes:
self._reapply_writes_to_succeeded_nodes(self.tasks)
self._resume_error_handlers_if_applicable()
@@ -684,11 +690,25 @@ 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]
)
if self._seal_unclaimed_writes:
# A loaded write no task of this run claimed belongs to a task the
# resume dropped, such as a `Send` that `Command(goto=...)` replaced.
self._delta_channels_forced_snapshot.update(
delta_channels_with_pending_writes(
self.specs,
[
w
for w in self.checkpoint_pending_writes
if w[0] != NULL_TASK_ID and w[0] not in self.tasks
],
)
)
self._seal_unclaimed_writes = False
# all tasks have finished
self.updated_channels = apply_writes(
self.checkpoint,
@@ -734,17 +754,26 @@ class PregelLoop:
# private
@property
def _reapplies_pending_writes(self) -> bool:
"""Whether the writes loaded with the checkpoint go back to the tasks
that made them, instead of those tasks rerunning."""
return not self.is_replaying
def _reapply_writes_to_succeeded_nodes(
self, tasks: Mapping[str, PregelExecutableTask]
) -> None:
"""Restore the output of finished tasks from checkpoint to in-memory tasks.
"""Restore successful channel writes from checkpoint to in-memory tasks.
Unfinished (failed or interrupted) tasks keep empty writes, so the
runner re-executes them or routes them to error handlers.
Skips control signals (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME)
so that failed/interrupted tasks remain with empty writes and will be
re-executed (or routed to error handlers) by the runner.
"""
for tid, status in read_task_statuses(self.checkpoint_pending_writes).items():
for tid, k, v in self.checkpoint_pending_writes:
if k in (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME):
continue
if task := tasks.get(tid):
task.writes.extend(status.output)
task.writes.append((k, v))
def _resume_error_handlers_if_applicable(self) -> None:
"""On resume, schedule error handlers for tasks that failed in a prior run.
@@ -814,13 +843,35 @@ class PregelLoop:
self.tasks[handler_task.id] = handler_task
def _pending_interrupts(self) -> set[str]:
"""Return the ids of interrupts that are still waiting for an answer."""
return {
interrupt.id
for status in read_task_statuses(self.checkpoint_pending_writes).values()
for interrupt in status.pending_interrupts
"""Return the set of interrupt ids that are pending without corresponding resume values."""
# mapping of task ids to interrupt ids
pending_interrupts: dict[str, str] = {}
# set of resume task ids
pending_resumes: set[str] = set()
for task_id, write_type, value in self.checkpoint_pending_writes:
if write_type == INTERRUPT:
# interrupts is always a list, but there should only be one element
pending_interrupts[task_id] = value[0].id
elif write_type == RESUME:
pending_resumes.add(task_id)
resumed_interrupt_ids = {
pending_interrupts[task_id]
for task_id in pending_resumes
if task_id in pending_interrupts
}
# Keep only interrupts whose interrupt_id is not resumed
hanging_interrupts: set[str] = {
interrupt_id
for interrupt_id in pending_interrupts.values()
if interrupt_id not in resumed_interrupt_ids
}
return hanging_interrupts
def _first(
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
) -> set[str] | None:
@@ -874,6 +925,17 @@ class PregelLoop:
self.checkpoint_pending_writes = [
w for w in self.checkpoint_pending_writes if w[1] != RESUME
]
# A resume that reapplies the head's pending writes only learns which
# of them its tasks claim once they are scheduled, so `after_tick`
# seals the rest.
self._seal_unclaimed_writes = is_resuming and self._reapplies_pending_writes
self._delta_channels_forced_snapshot = (
set()
if self._seal_unclaimed_writes
else delta_channels_with_pending_writes(
self.specs, self.checkpoint_pending_writes
)
)
# map command to writes
if input_is_command:
@@ -967,7 +1029,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]
@@ -1111,8 +1173,10 @@ class PregelLoop:
)
# create new checkpoint
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, new_counters)
| self._delta_channels_with_overwrite
delta_channels_to_snapshot(
self.channels, new_counters, self.checkpoint["channel_versions"]
)
| self._delta_channels_forced_snapshot
if do_checkpoint
else set()
)
@@ -1126,11 +1190,14 @@ class PregelLoop:
if do_checkpoint
else None,
channels_to_snapshot=channels_to_snapshot,
# An exit-mode run that stops in its first tick still carries the
# loaded checkpoint's `updated_channels`, though nothing moved.
stored_versions=self.checkpoint_previous_versions,
)
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
@@ -1214,8 +1281,10 @@ class PregelLoop:
self.checkpoint_metadata.get("counters_since_delta_snapshot") or {}
)
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, counters)
| self._delta_channels_with_overwrite
delta_channels_to_snapshot(
self.channels, counters, self.checkpoint["channel_versions"]
)
| self._delta_channels_forced_snapshot
)
pending = [
@@ -1576,7 +1645,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if handler_task is None:
return None
self.tasks[handler_task.id] = handler_task
if not self.is_replaying:
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
for task in self.match_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
@@ -1660,7 +1729,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
)
@@ -1831,7 +1899,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
if handler_task is None:
return None
self.tasks[handler_task.id] = handler_task
if not self.is_replaying:
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
for task in await self.amatch_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
@@ -1918,7 +1986,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
)
+2 -4
View File
@@ -45,7 +45,6 @@ from langgraph.errors import GraphBubbleUp, GraphInterrupt
from langgraph.pregel._algo import Call
from langgraph.pregel._executor import Submit
from langgraph.pregel._retry import arun_with_retry, run_with_retry
from langgraph.pregel._task_status import CONTROL_WRITES
from langgraph.types import (
CachePolicy,
PregelExecutableTask,
@@ -607,9 +606,8 @@ class PregelRunner:
task.config is None or TAG_HIDDEN not in task.config.get("tags", [])
):
self.node_finished(task.name)
if all(chan in CONTROL_WRITES for chan, _ in task.writes):
# record that the task finished, even if it produced no output
# (see `langgraph.pregel._task_status`)
if not task.writes:
# add no writes marker
task.writes.append((NO_WRITES, None))
# save task writes to checkpointer
self.put_writes()(task.id, task.writes) # type: ignore[misc]
@@ -1,127 +0,0 @@
"""Read the status of each task from the writes recorded for a superstep.
While a superstep is open, the checkpointer keeps a log of writes for each
task in that step. Entries are added as tasks run and are only discarded when
the whole superstep finishes and a new checkpoint is saved. When a task runs
again, for example after being resumed, its earlier entries stay in the log.
This module is the single place that turns that log into task status. Code that
needs to know whether a task finished, which interrupts it raised, which of them
are still waiting for an answer, or which output it produced must use
`read_task_statuses` instead of inspecting the writes directly.
The log uses two kinds of writes:
- Control writes describe what happened to a task: `INTERRUPT` (the task asked
a question), `RESUME` (answers the task has received), `ERROR`, and
`ERROR_SOURCE_NODE`. `INTERRUPT`, `RESUME` and `ERROR` each have a fixed slot
per task (`WRITES_IDX_MAP`), so a newer write of the same kind can replace an
older one.
- Every other write is output: channel writes, `RETURN` for functional tasks,
and the `NO_WRITES` marker.
The rules are:
1. When a task that ran finishes successfully, `PregelRunner.commit` records at
least one output write, adding `NO_WRITES` if the task produced no other
output.
2. A task that pauses at an interrupt records only control writes.
3. A task is therefore treated as finished if and only if it has an output
write.
4. Because `INTERRUPT` is stored in a fixed slot, its recorded value is the most
recent question the task asked. That question is waiting for an answer only
while the task is unfinished.
A `RESUME` write never means a task is finished: it can hold the answer to an
earlier question while the task waits on a later one.
What these rules cannot see:
- A task whose result came from the cache does not go through
`PregelRunner.commit`, so nothing is recorded for it. It reads as not
finished.
- A task that fails can record partial output writes along with its error. It
reads as finished, which is how the executor has always treated it.
- Writes recorded before rule 1 existed may describe a finished task with no
output using only control writes. Those tasks read as unfinished, which
matches how they were treated before.
"""
from __future__ import annotations
from collections.abc import Iterable, Sequence
from dataclasses import dataclass
from typing import Any
from langgraph.checkpoint.base import PendingWrite
from langgraph._internal._constants import (
ERROR,
ERROR_SOURCE_NODE,
INTERRUPT,
NULL_TASK_ID,
RESUME,
)
from langgraph.types import Interrupt
__all__ = ("CONTROL_WRITES", "TaskStatus", "read_task_statuses")
CONTROL_WRITES = frozenset((ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME))
"""Channels that describe what happened to a task rather than what it produced."""
@dataclass(frozen=True, slots=True)
class TaskStatus:
"""The status of one task, read from the writes recorded for its superstep."""
output: tuple[tuple[str, Any], ...] = ()
"""Output writes in recorded order. Empty if the task has not finished."""
interrupts: tuple[Interrupt, ...] = ()
"""The most recent interrupts the task raised, whether or not they were answered."""
error: BaseException | None = None
"""The recorded error, if any."""
@property
def finished(self) -> bool:
"""Whether the task ran to completion."""
return bool(self.output)
@property
def pending_interrupts(self) -> tuple[Interrupt, ...]:
"""Interrupts waiting for an answer. Always empty for a finished task."""
return () if self.finished else self.interrupts
def read_task_statuses(
pending_writes: Iterable[PendingWrite],
) -> dict[str, TaskStatus]:
"""Return the status of every task that has recorded writes, keyed by task id.
Writes from `NULL_TASK_ID` are input to the superstep, not task activity, so
they are not included.
"""
output: dict[str, list[tuple[str, Any]]] = {}
interrupts: dict[str, list[Interrupt]] = {}
errors: dict[str, BaseException] = {}
for task_id, channel, value in pending_writes:
if task_id == NULL_TASK_ID:
continue
output.setdefault(task_id, [])
if channel == INTERRUPT:
interrupts.setdefault(task_id, []).extend(
value if isinstance(value, Sequence) else [value]
)
elif channel == ERROR:
errors.setdefault(task_id, value)
elif channel not in CONTROL_WRITES:
output[task_id].append((channel, value))
return {
task_id: TaskStatus(
output=tuple(task_output),
interrupts=tuple(interrupts.get(task_id, ())),
error=errors.get(task_id),
)
for task_id, task_output in output.items()
}
+33 -18
View File
@@ -26,7 +26,6 @@ from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel
from langgraph.constants import TAG_HIDDEN
from langgraph.pregel._io import read_channels
from langgraph.pregel._task_status import TaskStatus, read_task_statuses
from langgraph.types import (
CheckpointPayload,
PregelExecutableTask,
@@ -38,8 +37,6 @@ from langgraph.types import (
TASK_NAMESPACE = UUID("6ba7b831-9dad-11d1-80b4-00c04fd430c8")
_NOT_STARTED = TaskStatus()
def map_debug_tasks(tasks: Iterable[PregelExecutableTask]) -> Iterator[TaskPayload]:
"""Produce "task" events for stream_mode=debug."""
@@ -214,21 +211,35 @@ def tasks_w_writes(
pending_writes: list[PendingWrite] | None,
states: dict[str, RunnableConfig | StateSnapshot] | None,
output_keys: str | Sequence[str],
*,
live: bool = False,
) -> tuple[PregelTask, ...]:
"""Apply writes / subgraph states to tasks to be returned in a StateSnapshot.
With `live=True`, tasks report only the interrupts still waiting for an
answer, as of the most recent writes. Otherwise tasks report the interrupts
they raised in the step, including answered ones, as a record of the step.
"""
statuses = read_task_statuses(pending_writes or [])
"""Apply writes / subgraph states to tasks to be returned in a StateSnapshot."""
pending_writes = pending_writes or []
out: list[PregelTask] = []
for task in tasks:
status = statuses.get(task.id, _NOT_STARTED)
rtn = next((val for chan, val in status.output if chan == RETURN), MISSING)
task_writes = [(chan, val) for chan, val in status.output if chan != RETURN]
rtn = next(
(
val
for tid, chan, val in pending_writes
if tid == task.id and chan == RETURN
),
MISSING,
)
task_error = next(
(exc for tid, n, exc in pending_writes if tid == task.id and n == ERROR),
None,
)
task_interrupts = tuple(
v
for tid, n, vv in pending_writes
if tid == task.id and n == INTERRUPT
for v in (vv if isinstance(vv, Sequence) else [vv])
)
task_writes = [
(chan, val)
for tid, chan, val in pending_writes
if tid == task.id and chan not in (ERROR, INTERRUPT, RETURN)
]
if rtn is not MISSING:
task_result = rtn
@@ -250,15 +261,19 @@ def tasks_w_writes(
mapped_writes = map_task_result_writes(filtered_writes)
task_result = mapped_writes if filtered_writes else {}
has_writes = rtn is not MISSING or any(
w[0] == task.id and w[1] not in (ERROR, INTERRUPT) for w in pending_writes
)
out.append(
PregelTask(
task.id,
task.name,
task.path,
status.error,
status.pending_interrupts if live else status.interrupts,
task_error,
task_interrupts,
states.get(task.id) if states else None,
task_result if status.finished else None,
task_result if has_writes else None,
)
)
return tuple(out)
+186 -70
View File
@@ -79,6 +79,7 @@ from langgraph._internal._constants import (
CONFIG_KEY_STREAM_MESSAGES_V2,
CONFIG_KEY_TASK_ID,
CONFIG_KEY_THREAD_ID,
ERROR,
INPUT,
INTERRUPT,
NS_END,
@@ -107,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
@@ -132,8 +134,10 @@ 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,
versions_seen_without_bumps,
)
from langgraph.pregel._draw import draw_graph
from langgraph.pregel._io import map_input, read_channels
@@ -148,7 +152,6 @@ from langgraph.pregel._messages import (
from langgraph.pregel._read import DEFAULT_BOUND, PregelNode
from langgraph.pregel._retry import RetryPolicy
from langgraph.pregel._runner import PregelRunner
from langgraph.pregel._task_status import read_task_statuses
from langgraph.pregel._tools import StreamToolCallHandler
from langgraph.pregel._utils import (
get_new_channel_versions,
@@ -1147,16 +1150,8 @@ class Pregel(
config: RunnableConfig,
saved: CheckpointTuple | None,
recurse: BaseCheckpointSaver | None = None,
live: bool = False,
apply_pending_writes: bool = False,
) -> StateSnapshot:
"""Build a `StateSnapshot` from a saved checkpoint and its pending writes.
With `live=True` the snapshot shows current status: values include the
output of tasks that already finished, `next` lists only tasks that still
need to run, and `interrupts` lists only questions still waiting for an
answer. Otherwise the snapshot is a record of the step: values as of the
start of the step, every task in the step, and the interrupts they raised.
"""
if not saved:
return StateSnapshot(
values={},
@@ -1244,10 +1239,13 @@ class Pregel(
None,
self.trigger_to_nodes,
)
if live and saved.pending_writes:
for tid, status in read_task_statuses(saved.pending_writes).items():
if tid in next_tasks:
next_tasks[tid].writes.extend(status.output)
if apply_pending_writes and saved.pending_writes:
for tid, k, v in saved.pending_writes:
if k in (ERROR, INTERRUPT):
continue
if tid not in next_tasks:
continue
next_tasks[tid].writes.append((k, v))
if tasks := [t for t in next_tasks.values() if t.writes]:
apply_writes(
saved.checkpoint, channels, tasks, None, self.trigger_to_nodes
@@ -1257,7 +1255,6 @@ class Pregel(
saved.pending_writes,
task_states,
self.stream_channels_asis,
live=live,
)
# assemble the state snapshot
return StateSnapshot(
@@ -1276,16 +1273,8 @@ class Pregel(
config: RunnableConfig,
saved: CheckpointTuple | None,
recurse: BaseCheckpointSaver | None = None,
live: bool = False,
apply_pending_writes: bool = False,
) -> StateSnapshot:
"""Build a `StateSnapshot` from a saved checkpoint and its pending writes.
With `live=True` the snapshot shows current status: values include the
output of tasks that already finished, `next` lists only tasks that still
need to run, and `interrupts` lists only questions still waiting for an
answer. Otherwise the snapshot is a record of the step: values as of the
start of the step, every task in the step, and the interrupts they raised.
"""
if not saved:
return StateSnapshot(
values={},
@@ -1373,10 +1362,13 @@ class Pregel(
None,
self.trigger_to_nodes,
)
if live and saved.pending_writes:
for tid, status in read_task_statuses(saved.pending_writes).items():
if tid in next_tasks:
next_tasks[tid].writes.extend(status.output)
if apply_pending_writes and saved.pending_writes:
for tid, k, v in saved.pending_writes:
if k in (ERROR, INTERRUPT):
continue
if tid not in next_tasks:
continue
next_tasks[tid].writes.append((k, v))
if tasks := [t for t in next_tasks.values() if t.writes]:
apply_writes(
saved.checkpoint, channels, tasks, None, self.trigger_to_nodes
@@ -1387,7 +1379,6 @@ class Pregel(
saved.pending_writes,
task_states,
self.stream_channels_asis,
live=live,
)
# assemble the state snapshot
return StateSnapshot(
@@ -1442,7 +1433,7 @@ class Pregel(
config,
saved,
recurse=checkpointer if subgraphs else None,
live=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
)
async def aget_state(
@@ -1486,7 +1477,7 @@ class Pregel(
config,
saved,
recurse=checkpointer if subgraphs else None,
live=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
)
def get_state_history(
@@ -1649,12 +1640,22 @@ 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)
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
)
if saved is not None:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
@@ -1722,12 +1723,13 @@ class Pregel(
checkpointer.get_next_version,
self.trigger_to_nodes,
)
# apply writes from tasks that already finished
for tid, status in read_task_statuses(
saved.pending_writes or []
).items():
if tid in next_tasks:
next_tasks[tid].writes.extend(status.output)
# apply writes from tasks that already ran
for tid, k, v in saved.pending_writes or []:
if k in (ERROR, INTERRUPT):
continue
if tid not in next_tasks:
continue
next_tasks[tid].writes.append((k, v))
# clear all current tasks
apply_writes(
checkpoint,
@@ -1737,9 +1739,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,
@@ -1747,7 +1757,7 @@ class Pregel(
},
get_new_channel_versions(
checkpoint_previous_versions,
checkpoint["channel_versions"],
next_checkpoint["channel_versions"],
),
)
return patch_checkpoint_map(
@@ -1776,9 +1786,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,
@@ -1788,7 +1806,7 @@ class Pregel(
},
get_new_channel_versions(
checkpoint_previous_versions,
checkpoint["channel_versions"],
next_checkpoint["channel_versions"],
),
)
@@ -1921,7 +1939,9 @@ class Pregel(
as_node = tuple(self.nodes)[0]
elif as_node is None and not any(
v
for vv in checkpoint["versions_seen"].values()
for vv in versions_seen_without_bumps(
checkpoint["versions_seen"]
).values()
for v in vv.values()
):
if (
@@ -1932,7 +1952,9 @@ class Pregel(
elif as_node is None:
last_seen_by_node = sorted(
(v, n)
for n, seen in checkpoint["versions_seen"].items()
for n, seen in versions_seen_without_bumps(
checkpoint["versions_seen"]
).items()
if n in self.nodes
for v in seen.values()
)
@@ -2009,13 +2031,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,
@@ -2031,18 +2061,30 @@ 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,
channel_versions=checkpoint["channel_versions"],
)
)
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,
)
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,
@@ -2114,12 +2156,22 @@ 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)
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
)
if saved is not None:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
@@ -2185,12 +2237,13 @@ class Pregel(
checkpointer.get_next_version,
self.trigger_to_nodes,
)
# apply writes from tasks that already finished
for tid, status in read_task_statuses(
saved.pending_writes or []
).items():
if tid in next_tasks:
next_tasks[tid].writes.extend(status.output)
# apply writes from tasks that already ran
for tid, k, v in saved.pending_writes or []:
if k in (ERROR, INTERRUPT):
continue
if tid not in next_tasks:
continue
next_tasks[tid].writes.append((k, v))
# clear all current tasks
apply_writes(
checkpoint,
@@ -2200,16 +2253,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(
@@ -2238,9 +2300,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,
@@ -2250,7 +2320,7 @@ class Pregel(
},
get_new_channel_versions(
checkpoint_previous_versions,
checkpoint["channel_versions"],
next_checkpoint["channel_versions"],
),
)
@@ -2391,7 +2461,9 @@ class Pregel(
elif as_node is None:
last_seen_by_node = sorted(
(v, n)
for n, seen in checkpoint["versions_seen"].items()
for n, seen in versions_seen_without_bumps(
checkpoint["versions_seen"]
).items()
if n in self.nodes
for v in seen.values()
)
@@ -2468,13 +2540,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,
@@ -2490,18 +2570,30 @@ 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,
channel_versions=checkpoint["channel_versions"],
)
)
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,
)
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,
@@ -4191,6 +4283,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],
+1 -7
View File
@@ -726,13 +726,7 @@ class StateSnapshot(NamedTuple):
tasks: tuple[PregelTask, ...]
"""Tasks to execute in this step. If already attempted, may contain an error."""
interrupts: tuple[Interrupt, ...]
"""Interrupts that occurred in this step.
When reading the latest state (`get_state` without a `checkpoint_id`), this
contains only interrupts still waiting for an answer. When reading a specific
checkpoint or state history, it contains the most recent interrupt each task
raised in that step, including ones answered later in the same step.
"""
"""Interrupts that occurred in this step that are pending resolution."""
class Send:
+5 -3
View File
@@ -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,740 @@
"""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,
Send,
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 _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:
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 _build_send_fan_out(checkpointer: BaseCheckpointSaver) -> Any:
def q(state: _State) -> dict:
interrupt("continue?")
return _both("q")
builder = StateGraph(_State)
builder.add_node("p", lambda state: _both("p"))
builder.add_node("q", q)
builder.add_conditional_edges(
START, lambda state: [Send("p", state), Send("q", state)], ["p", "q"]
)
return builder.compile(checkpointer=checkpointer)
def test_resume_that_replaces_the_pending_sends_drops_the_finished_task(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
graph = _build_send_fan_out(sync_checkpointer)
config = _thread("t")
graph.invoke(_both("in"), config, durability=durability)
live = graph.invoke(
Command(resume="yes", goto=[Send("q", _both("unused"))]),
config,
durability=durability,
)
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"] == live["log"], (
f"'p' ran in the fan-out the resume replaced, but the reload reads "
f"{state.values['log']} against the live {live['log']}"
)
async def test_aresume_that_replaces_the_pending_sends_drops_the_finished_task(
async_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
graph = _build_send_fan_out(async_checkpointer)
config = _thread("t")
await graph.ainvoke(_both("in"), config, durability=durability)
live = await graph.ainvoke(
Command(resume="yes", goto=[Send("q", _both("unused"))]),
config,
durability=durability,
)
state = await graph.aget_state(config)
assert state.values["log"] == state.values["plain"] == live["log"], (
f"'p' ran in the fan-out the resume replaced, but the reload reads "
f"{state.values['log']} against the live {live['log']}"
)
def test_replay_interrupted_in_its_first_step_still_seals_the_fork(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
graph = _build_send_fan_out(sync_checkpointer)
config = _thread("t")
graph.invoke(_both("in"), config, durability=durability)
graph.invoke(None, graph.get_state(config).config, durability=durability)
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"] == ["in", "p"], (
f"the replay reran p, so the fork must not also replay the first p, "
f"but it reads {state.values['log']}"
)
async def test_areplay_interrupted_in_its_first_step_still_seals_the_fork(
async_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
graph = _build_send_fan_out(async_checkpointer)
config = _thread("t")
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.values["log"] == state.values["plain"] == ["in", "p"], (
f"the replay reran p, so the fork must not also replay the first p, "
f"but it reads {state.values['log']}"
)
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"]
)
@@ -94,6 +94,28 @@ async def test_forced_snapshot_single_run() -> None:
assert "seed-a" in state.values["a"]
async def test_supersteps_bound_skips_a_channel_never_written() -> None:
with patch(
"langgraph.pregel._checkpoint.DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT",
3,
):
saver = InMemorySaver()
graph = _build_two_channel_graph(saver, n_loops=4)
config = {"configurable": {"thread_id": "never-written"}}
graph.invoke({"a": ["seed-a"]}, config)
minted = [
t.config["configurable"]["checkpoint_id"]
for t in saver.list(config)
if "b" in t.checkpoint["channel_versions"]
]
assert not minted, (
f"b was never written, but {len(minted)} checkpoints minted it a version"
)
assert graph.get_state(config).values["b"] == []
async def test_forced_snapshot_accumulates_across_runs() -> None:
"""Supersteps counter for an unwritten channel persists across separate
invoke() calls. After enough runs, the channel is force-snapshotted."""
@@ -139,13 +161,17 @@ async def test_predicate_fires_on_supersteps_overflow() -> None:
channels = {"x": ch_instance}
counters: dict[str, tuple[int, int]] = {"x": (0, 5000)}
result = delta_channels_to_snapshot(channels, counters)
result = delta_channels_to_snapshot(channels, counters, {"x": 1})
assert "x" in result
counters_below: dict[str, tuple[int, int]] = {"x": (0, 4999)}
result2 = delta_channels_to_snapshot(channels, counters_below)
result2 = delta_channels_to_snapshot(channels, counters_below, {"x": 1})
assert "x" not in result2
assert not delta_channels_to_snapshot(channels, counters, {}), (
"a channel with no version was never written, so it has nothing to snapshot"
)
async def test_counter_reset_after_supersteps_snapshot() -> None:
"""After the supersteps bound triggers a snapshot, the counters for
@@ -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 == ()
@@ -1,534 +0,0 @@
"""State reads while some tasks of a superstep are finished and others are paused.
When parallel tasks each call `interrupt()` and only some of them are resumed,
the superstep stays open. Its recorded writes then contain the old interrupt of
each finished task next to that task's output. These tests check that state
reads, which are rebuilt from the checkpointer, report only the interrupts that
still need an answer.
"""
import operator
import sys
import uuid
from collections import Counter
from typing import Annotated, Any
import pytest
from langgraph.checkpoint.base import BaseCheckpointSaver
from typing_extensions import TypedDict
from langgraph._internal._constants import (
ERROR,
INTERRUPT,
NO_WRITES,
NULL_TASK_ID,
RESUME,
RETURN,
)
from langgraph.func import entrypoint, task
from langgraph.graph import END, START, StateGraph
from langgraph.pregel._task_status import read_task_statuses
from langgraph.types import Command, Durability, Interrupt, Send, interrupt
pytestmark = pytest.mark.anyio
NEEDS_CONTEXTVARS = pytest.mark.skipif(
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
class State(TypedDict, total=False):
log: Annotated[list[str], operator.add]
count: int
def _config() -> dict[str, Any]:
return {"configurable": {"thread_id": str(uuid.uuid4())}}
def _build_parallel(
checkpointer: BaseCheckpointSaver,
calls: Counter[str],
*,
a_questions: int = 1,
a_returns: Any = "log",
):
"""Build a graph where nodes `a` and `b` start in parallel and both ask questions.
`a` asks `a_questions` questions in a row. `a_returns` controls what `a`
returns after its last answer. The default `"log"` returns the answers in
`log`. Any other value is returned as-is.
"""
def a(state: State) -> Any:
calls["a"] += 1
answers = [interrupt(f"A{i + 1}") for i in range(a_questions)]
if a_returns == "log":
return {"log": [f"a:{answer}" for answer in answers]}
return a_returns
def b(state: State) -> State:
calls["b"] += 1
return {"log": [f"b:{interrupt('B')}"]}
builder = StateGraph(State)
builder.add_node("a", a)
builder.add_node("b", b)
builder.add_edge(START, "a")
builder.add_edge(START, "b")
builder.add_edge("a", END)
builder.add_edge("b", END)
return builder.compile(checkpointer=checkpointer)
def _interrupt_by_value(snapshot: Any, value: str) -> Interrupt:
return next(i for i in snapshot.interrupts if i.value == value)
def _task(snapshot: Any, name: str) -> Any:
return next(t for t in snapshot.tasks if t.name == name)
def _interrupt_values(interrupts: Any) -> list[str]:
return sorted(i.value for i in interrupts)
# --- Task A answered and finished, task B still paused ---
def test_finished_task_does_not_report_answered_interrupt(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
calls: Counter[str] = Counter()
graph = _build_parallel(sync_checkpointer, calls)
config = _config()
graph.invoke({"log": []}, config, durability=durability)
snapshot = graph.get_state(config)
assert _interrupt_values(snapshot.interrupts) == ["A1", "B"]
graph.invoke(
Command(resume={_interrupt_by_value(snapshot, "A1").id: "yes"}),
config,
durability=durability,
)
snapshot = graph.get_state(config)
assert _interrupt_values(snapshot.interrupts) == ["B"]
assert snapshot.next == ("b",)
assert _task(snapshot, "a").interrupts == ()
assert _task(snapshot, "a").result == {"log": ["a:yes"]}
assert _interrupt_values(_task(snapshot, "b").interrupts) == ["B"]
assert _task(snapshot, "b").result is None
# Reading the same checkpoint by id gives the record of the step: every task
# in it, and every question asked, including the one A already answered.
record = graph.get_state(snapshot.config)
assert sorted(record.next) == ["a", "b"]
assert _interrupt_values(record.interrupts) == ["A1", "B"]
assert _interrupt_values(_task(record, "a").interrupts) == ["A1"]
assert _task(record, "a").result == {"log": ["a:yes"]}
# B can still be answered, and the graph finishes normally.
result = graph.invoke(
Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}),
config,
durability=durability,
)
assert sorted(result["log"]) == ["a:yes", "b:ok"]
assert calls == {"a": 2, "b": 3}
snapshot = graph.get_state(config)
assert snapshot.next == ()
assert snapshot.interrupts == ()
# History still shows where each question was asked.
asked = [
_interrupt_values(s.interrupts)
for s in graph.get_state_history(config)
if s.interrupts
]
if durability != "exit":
assert asked == [["A1", "B"]]
@NEEDS_CONTEXTVARS
async def test_finished_task_does_not_report_answered_interrupt_async(
async_checkpointer: BaseCheckpointSaver,
) -> None:
calls: Counter[str] = Counter()
graph = _build_parallel(async_checkpointer, calls)
config = _config()
await graph.ainvoke({"log": []}, config)
snapshot = await graph.aget_state(config)
await graph.ainvoke(
Command(resume={_interrupt_by_value(snapshot, "A1").id: "yes"}), config
)
snapshot = await graph.aget_state(config)
assert _interrupt_values(snapshot.interrupts) == ["B"]
assert snapshot.next == ("b",)
assert _task(snapshot, "a").interrupts == ()
assert _task(snapshot, "a").result == {"log": ["a:yes"]}
assert _interrupt_values(_task(snapshot, "b").interrupts) == ["B"]
record = await graph.aget_state(snapshot.config)
assert _interrupt_values(record.interrupts) == ["A1", "B"]
assert _interrupt_values(_task(record, "a").interrupts) == ["A1"]
result = await graph.ainvoke(
Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}), config
)
assert sorted(result["log"]) == ["a:yes", "b:ok"]
assert calls == {"a": 2, "b": 3}
# --- Task A answered its first question and asked a second one ---
def test_task_paused_at_second_question_stays_pending(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
calls: Counter[str] = Counter()
graph = _build_parallel(sync_checkpointer, calls, a_questions=2)
config = _config()
graph.invoke({"log": []}, config)
snapshot = graph.get_state(config)
graph.invoke(
Command(resume={_interrupt_by_value(snapshot, "A1").id: "one"}), config
)
snapshot = graph.get_state(config)
assert _interrupt_values(snapshot.interrupts) == ["A2", "B"]
# A is not finished: it has a saved answer, but no output.
assert sorted(snapshot.next) == ["a", "b"]
assert _interrupt_values(_task(snapshot, "a").interrupts) == ["A2"]
assert _task(snapshot, "a").result is None
assert _interrupt_values(_task(snapshot, "b").interrupts) == ["B"]
# Both remaining questions can be answered together.
result = graph.invoke(
Command(
resume={
_interrupt_by_value(snapshot, "A2").id: "two",
_interrupt_by_value(snapshot, "B").id: "ok",
}
),
config,
)
assert sorted(result["log"]) == ["a:one", "a:two", "b:ok"]
snapshot = graph.get_state(config)
assert snapshot.next == ()
assert snapshot.interrupts == ()
@NEEDS_CONTEXTVARS
async def test_task_paused_at_second_question_stays_pending_async(
async_checkpointer: BaseCheckpointSaver,
) -> None:
calls: Counter[str] = Counter()
graph = _build_parallel(async_checkpointer, calls, a_questions=2)
config = _config()
await graph.ainvoke({"log": []}, config)
snapshot = await graph.aget_state(config)
await graph.ainvoke(
Command(resume={_interrupt_by_value(snapshot, "A1").id: "one"}), config
)
snapshot = await graph.aget_state(config)
assert _interrupt_values(snapshot.interrupts) == ["A2", "B"]
assert sorted(snapshot.next) == ["a", "b"]
assert _interrupt_values(_task(snapshot, "a").interrupts) == ["A2"]
assert _task(snapshot, "a").result is None
def test_task_paused_at_second_question_then_other_task_finishes(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
calls: Counter[str] = Counter()
graph = _build_parallel(sync_checkpointer, calls, a_questions=2)
config = _config()
graph.invoke({"log": []}, config)
snapshot = graph.get_state(config)
graph.invoke(
Command(resume={_interrupt_by_value(snapshot, "A1").id: "one"}), config
)
snapshot = graph.get_state(config)
graph.invoke(Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}), config)
snapshot = graph.get_state(config)
assert _interrupt_values(snapshot.interrupts) == ["A2"]
assert snapshot.next == ("a",)
assert _task(snapshot, "b").interrupts == ()
assert _task(snapshot, "b").result == {"log": ["b:ok"]}
result = graph.invoke(
Command(resume={_interrupt_by_value(snapshot, "A2").id: "two"}), config
)
assert sorted(result["log"]) == ["a:one", "a:two", "b:ok"]
def test_resume_without_id_rejected_when_second_question_and_other_task_pending(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
calls: Counter[str] = Counter()
graph = _build_parallel(sync_checkpointer, calls, a_questions=2)
config = _config()
graph.invoke({"log": []}, config)
snapshot = graph.get_state(config)
graph.invoke(
Command(resume={_interrupt_by_value(snapshot, "A1").id: "one"}), config
)
# A2 and B are both waiting, so a resume value without an id is ambiguous.
with pytest.raises(RuntimeError, match="multiple pending interrupts"):
graph.invoke(Command(resume="ambiguous"), config)
def test_resume_without_id_rejected_when_subgraph_has_parallel_interrupts(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
# A subgraph node whose child graph pauses in two parallel nodes records
# both interrupts under one parent task. Both count as pending, so a resume
# value without an id is ambiguous. (Before, only the first was counted and
# the value went to whichever interrupt consumed it first.)
child_builder = StateGraph(State)
child_builder.add_node("a", lambda s: {"log": [f"a:{interrupt('A')}"]})
child_builder.add_node("b", lambda s: {"log": [f"b:{interrupt('B')}"]})
child_builder.add_edge(START, "a")
child_builder.add_edge(START, "b")
builder = StateGraph(State)
builder.add_node("child", child_builder.compile())
builder.add_edge(START, "child")
graph = builder.compile(checkpointer=sync_checkpointer)
config = _config()
graph.invoke({"log": []}, config)
snapshot = graph.get_state(config)
assert _interrupt_values(snapshot.interrupts) == ["A", "B"]
with pytest.raises(RuntimeError, match="multiple pending interrupts"):
graph.invoke(Command(resume="ambiguous"), config)
result = graph.invoke(
Command(
resume={
_interrupt_by_value(snapshot, "A").id: "x",
_interrupt_by_value(snapshot, "B").id: "y",
}
),
config,
)
assert sorted(result["log"]) == ["a:x", "b:y"]
# --- Task A finished with an empty or falsy result ---
@pytest.mark.parametrize(
"a_returns",
[None, {}, {"count": 0}, {"log": []}],
ids=["none", "empty_dict", "zero", "empty_list"],
)
def test_task_finished_with_falsy_result(
sync_checkpointer: BaseCheckpointSaver, a_returns: Any
) -> None:
calls: Counter[str] = Counter()
graph = _build_parallel(sync_checkpointer, calls, a_returns=a_returns)
config = _config()
graph.invoke({"log": []}, config)
snapshot = graph.get_state(config)
graph.invoke(
Command(resume={_interrupt_by_value(snapshot, "A1").id: "yes"}), config
)
snapshot = graph.get_state(config)
assert _interrupt_values(snapshot.interrupts) == ["B"]
assert snapshot.next == ("b",)
assert _task(snapshot, "a").interrupts == ()
graph.invoke(Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}), config)
# A already finished, so resuming B must not run A again.
assert calls == {"a": 2, "b": 3}
snapshot = graph.get_state(config)
assert snapshot.next == ()
assert snapshot.interrupts == ()
@pytest.mark.parametrize("a_returns", [None, {"count": 0}], ids=["none", "zero"])
@NEEDS_CONTEXTVARS
async def test_task_finished_with_falsy_result_async(
async_checkpointer: BaseCheckpointSaver, a_returns: Any
) -> None:
calls: Counter[str] = Counter()
graph = _build_parallel(async_checkpointer, calls, a_returns=a_returns)
config = _config()
await graph.ainvoke({"log": []}, config)
snapshot = await graph.aget_state(config)
await graph.ainvoke(
Command(resume={_interrupt_by_value(snapshot, "A1").id: "yes"}), config
)
snapshot = await graph.aget_state(config)
assert _interrupt_values(snapshot.interrupts) == ["B"]
assert snapshot.next == ("b",)
assert _task(snapshot, "a").interrupts == ()
await graph.ainvoke(
Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}), config
)
assert calls == {"a": 2, "b": 3}
# --- Subgraphs and the functional API ---
def test_parallel_subgraphs_report_only_pending_interrupts(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
class ChildState(TypedDict):
prompt: str
answers: Annotated[list[str], operator.add]
def ask(state: ChildState) -> dict[str, Any]:
return {"answers": [interrupt(state["prompt"])]}
child_builder = StateGraph(ChildState)
child_builder.add_node("ask", ask)
child_builder.add_edge(START, "ask")
child = child_builder.compile()
class ParentState(TypedDict):
answers: Annotated[list[str], operator.add]
builder = StateGraph(ParentState)
builder.add_node("child", child)
builder.add_conditional_edges(
START,
lambda _: [Send("child", {"prompt": p, "answers": []}) for p in ("a", "b")],
["child"],
)
graph = builder.compile(checkpointer=sync_checkpointer)
config = _config()
graph.invoke({"answers": []}, config)
snapshot = graph.get_state(config)
assert _interrupt_values(snapshot.interrupts) == ["a", "b"]
graph.invoke(Command(resume={_interrupt_by_value(snapshot, "a").id: "x"}), config)
snapshot = graph.get_state(config)
assert _interrupt_values(snapshot.interrupts) == ["b"]
assert snapshot.next == ("child",)
finished = next(t for t in snapshot.tasks if t.result is not None)
assert finished.interrupts == ()
assert finished.result == {"answers": ["x"]}
result = graph.invoke(
Command(resume={_interrupt_by_value(snapshot, "b").id: "y"}), config
)
assert sorted(result["answers"]) == ["x", "y"]
def test_functional_task_finished_with_none_is_not_rerun(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
calls: Counter[str] = Counter()
@task
def ask_a() -> None:
calls["a"] += 1
interrupt("A")
@task
def ask_b() -> str:
calls["b"] += 1
return interrupt("B")
@entrypoint(checkpointer=sync_checkpointer)
def workflow(_: Any) -> list[Any]:
a, b = ask_a(), ask_b()
return [a.result(), b.result()]
config = _config()
workflow.invoke(1, config)
snapshot = workflow.get_state(config)
workflow.invoke(
Command(resume={_interrupt_by_value(snapshot, "A").id: "x"}), config
)
snapshot = workflow.get_state(config)
assert _interrupt_values(snapshot.interrupts) == ["B"]
result = workflow.invoke(
Command(resume={_interrupt_by_value(snapshot, "B").id: "y"}), config
)
assert result == [None, "y"]
assert calls == {"a": 2, "b": 3}
# --- Reading task status from recorded writes ---
def test_read_task_statuses() -> None:
a1 = Interrupt(value="A1", id="a")
a2 = Interrupt(value="A2", id="a")
b = Interrupt(value="B", id="b")
error = ValueError("boom")
statuses = read_task_statuses(
[
# answered and finished: old interrupt stays recorded
("finished", INTERRUPT, (a1,)),
("finished", RESUME, ["yes"]),
("finished", "log", ["a:yes"]),
# answered once, then paused at a second question
("paused", INTERRUPT, (a2,)),
("paused", RESUME, ["one"]),
# finished with no output
("no_output", INTERRUPT, (b,)),
("no_output", RESUME, ["ok"]),
("no_output", NO_WRITES, None),
# functional task that returned None
("returned_none", RETURN, None),
# failed
("failed", ERROR, error),
# not a task
(NULL_TASK_ID, RESUME, "global"),
]
)
assert set(statuses) == {
"finished",
"paused",
"no_output",
"returned_none",
"failed",
}
assert statuses["finished"].finished
assert statuses["finished"].interrupts == (a1,)
assert statuses["finished"].pending_interrupts == ()
assert statuses["finished"].output == (("log", ["a:yes"]),)
assert not statuses["paused"].finished
assert statuses["paused"].interrupts == (a2,)
assert statuses["paused"].pending_interrupts == (a2,)
assert statuses["paused"].output == ()
assert statuses["no_output"].finished
assert statuses["no_output"].interrupts == (b,)
assert statuses["no_output"].pending_interrupts == ()
assert statuses["returned_none"].finished
assert statuses["returned_none"].output == ((RETURN, None),)
assert not statuses["failed"].finished
assert statuses["failed"].error is error
+55
View File
@@ -9410,6 +9410,61 @@ def test_fork_does_not_apply_pending_writes(
assert result == {"value": 121}
def _extend(state: list, writes: list[list]) -> list:
return [*state, *(v for write in writes for v in write)]
class _IntVersionSaver(InMemorySaver):
"""Integer versions tie exactly where `InMemorySaver`'s break at random."""
get_next_version = BaseCheckpointSaver.get_next_version
def _build_chain_after_a_delta_channel() -> Pregel:
return Pregel(
nodes={
"a": NodeBuilder().subscribe_only("inp").do(lambda _: ["a"]).write_to("d"),
"b": NodeBuilder().subscribe_only("d").do(lambda _: "b").write_to("x"),
"c": NodeBuilder().subscribe_only("x").do(lambda _: "c").write_to("out"),
},
channels={
"inp": LastValue(str),
"d": DeltaChannel(_extend, snapshot_frequency=1),
"x": LastValue(str),
"out": LastValue(str),
},
input_channels=["inp"],
output_channels=["out"],
checkpointer=_IntVersionSaver(),
)
def test_update_state_after_an_exit_snapshot_infers_the_last_writer() -> None:
graph = _build_chain_after_a_delta_channel()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"inp": "go"}, config, durability="exit")
graph.update_state(config, "u")
values = graph.get_state(config).values
assert values["out"] == "u", (
f"the update should apply as c, the last node to write, but state is {values}"
)
async def test_aupdate_state_after_an_exit_snapshot_infers_the_last_writer() -> None:
graph = _build_chain_after_a_delta_channel()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"inp": "go"}, config, durability="exit")
await graph.aupdate_state(config, "u")
values = (await graph.aget_state(config)).values
assert values["out"] == "u", (
f"the update should apply as c, the last node to write, but state is {values}"
)
async def test_delta_channel_end_to_end_inmemory() -> None:
"""Full graph run: DeltaChannel accumulates correctly across multiple turns."""