mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-02 06:25:09 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
79781ccfbe | ||
|
|
c0405d246a | ||
|
|
e4d77cd222 | ||
|
|
858e55f232 | ||
|
|
e59ecbc23d | ||
|
|
719a4d71bc | ||
|
|
ddaf708cd0 | ||
|
|
9d5b0f1991 | ||
|
|
9d16b52955 | ||
|
|
2de0c47c1f | ||
|
|
377083220e | ||
|
|
9c5914861b | ||
|
|
24cf33f348 |
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user