mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-29 13:05:15 +02:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d2035cbf32 | ||
|
|
e21d19b936 |
@@ -119,6 +119,7 @@ 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,
|
||||
@@ -736,17 +737,14 @@ class PregelLoop:
|
||||
def _reapply_writes_to_succeeded_nodes(
|
||||
self, tasks: Mapping[str, PregelExecutableTask]
|
||||
) -> None:
|
||||
"""Restore successful channel writes from checkpoint to in-memory tasks.
|
||||
"""Restore the output of finished tasks from checkpoint to in-memory tasks.
|
||||
|
||||
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.
|
||||
Unfinished (failed or interrupted) tasks keep empty writes, so the
|
||||
runner re-executes them or routes them to error handlers.
|
||||
"""
|
||||
for tid, k, v in self.checkpoint_pending_writes:
|
||||
if k in (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME):
|
||||
continue
|
||||
for tid, status in read_task_statuses(self.checkpoint_pending_writes).items():
|
||||
if task := tasks.get(tid):
|
||||
task.writes.append((k, v))
|
||||
task.writes.extend(status.output)
|
||||
|
||||
def _resume_error_handlers_if_applicable(self) -> None:
|
||||
"""On resume, schedule error handlers for tasks that failed in a prior run.
|
||||
@@ -816,35 +814,13 @@ class PregelLoop:
|
||||
self.tasks[handler_task.id] = handler_task
|
||||
|
||||
def _pending_interrupts(self) -> set[str]:
|
||||
"""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
|
||||
"""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
|
||||
}
|
||||
|
||||
# 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:
|
||||
|
||||
@@ -45,6 +45,7 @@ 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,
|
||||
@@ -606,8 +607,9 @@ class PregelRunner:
|
||||
task.config is None or TAG_HIDDEN not in task.config.get("tags", [])
|
||||
):
|
||||
self.node_finished(task.name)
|
||||
if not task.writes:
|
||||
# add no writes marker
|
||||
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`)
|
||||
task.writes.append((NO_WRITES, None))
|
||||
# save task writes to checkpointer
|
||||
self.put_writes()(task.id, task.writes) # type: ignore[misc]
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
"""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,6 +26,7 @@ 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,
|
||||
@@ -37,6 +38,8 @@ 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."""
|
||||
@@ -211,35 +214,21 @@ 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."""
|
||||
pending_writes = pending_writes or []
|
||||
"""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 [])
|
||||
out: list[PregelTask] = []
|
||||
for task in tasks:
|
||||
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)
|
||||
]
|
||||
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]
|
||||
|
||||
if rtn is not MISSING:
|
||||
task_result = rtn
|
||||
@@ -261,19 +250,15 @@ 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,
|
||||
task_error,
|
||||
task_interrupts,
|
||||
status.error,
|
||||
status.pending_interrupts if live else status.interrupts,
|
||||
states.get(task.id) if states else None,
|
||||
task_result if has_writes else None,
|
||||
task_result if status.finished else None,
|
||||
)
|
||||
)
|
||||
return tuple(out)
|
||||
|
||||
@@ -79,7 +79,6 @@ from langgraph._internal._constants import (
|
||||
CONFIG_KEY_STREAM_MESSAGES_V2,
|
||||
CONFIG_KEY_TASK_ID,
|
||||
CONFIG_KEY_THREAD_ID,
|
||||
ERROR,
|
||||
INPUT,
|
||||
INTERRUPT,
|
||||
NS_END,
|
||||
@@ -149,6 +148,7 @@ 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,8 +1147,16 @@ class Pregel(
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
recurse: BaseCheckpointSaver | None = None,
|
||||
apply_pending_writes: bool = False,
|
||||
live: 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={},
|
||||
@@ -1236,13 +1244,10 @@ class Pregel(
|
||||
None,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
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 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 tasks := [t for t in next_tasks.values() if t.writes]:
|
||||
apply_writes(
|
||||
saved.checkpoint, channels, tasks, None, self.trigger_to_nodes
|
||||
@@ -1252,6 +1257,7 @@ class Pregel(
|
||||
saved.pending_writes,
|
||||
task_states,
|
||||
self.stream_channels_asis,
|
||||
live=live,
|
||||
)
|
||||
# assemble the state snapshot
|
||||
return StateSnapshot(
|
||||
@@ -1270,8 +1276,16 @@ class Pregel(
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
recurse: BaseCheckpointSaver | None = None,
|
||||
apply_pending_writes: bool = False,
|
||||
live: 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={},
|
||||
@@ -1359,13 +1373,10 @@ class Pregel(
|
||||
None,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
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 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 tasks := [t for t in next_tasks.values() if t.writes]:
|
||||
apply_writes(
|
||||
saved.checkpoint, channels, tasks, None, self.trigger_to_nodes
|
||||
@@ -1376,6 +1387,7 @@ class Pregel(
|
||||
saved.pending_writes,
|
||||
task_states,
|
||||
self.stream_channels_asis,
|
||||
live=live,
|
||||
)
|
||||
# assemble the state snapshot
|
||||
return StateSnapshot(
|
||||
@@ -1430,7 +1442,7 @@ class Pregel(
|
||||
config,
|
||||
saved,
|
||||
recurse=checkpointer if subgraphs else None,
|
||||
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
live=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
)
|
||||
|
||||
async def aget_state(
|
||||
@@ -1474,7 +1486,7 @@ class Pregel(
|
||||
config,
|
||||
saved,
|
||||
recurse=checkpointer if subgraphs else None,
|
||||
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
live=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
)
|
||||
|
||||
def get_state_history(
|
||||
@@ -1710,13 +1722,12 @@ class Pregel(
|
||||
checkpointer.get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
# 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))
|
||||
# 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)
|
||||
# clear all current tasks
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
@@ -2174,13 +2185,12 @@ class Pregel(
|
||||
checkpointer.get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
# 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))
|
||||
# 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)
|
||||
# clear all current tasks
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
|
||||
@@ -726,7 +726,13 @@ 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 that are pending resolution."""
|
||||
"""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.
|
||||
"""
|
||||
|
||||
|
||||
class Send:
|
||||
|
||||
@@ -0,0 +1,534 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user