From e21d19b936d575243a88507bee51491d2970deb9 Mon Sep 17 00:00:00 2001 From: Igor Soarez Date: Mon, 28 Sep 2026 14:48:00 +0300 Subject: [PATCH] fix(langgraph): stop reporting answered interrupts in get_state Fixes #6956 Fixes #8579 When parallel tasks interrupt, resuming one leaves its old `__interrupt__` write in the open superstep. `get_state()` treated that write as pending even after the task finished, so the answered question reappeared. Use `read_task_statuses()` for execution, state reads, and resume checks. An output write marks a task finished; `RESUME` alone does not, because the task may have interrupted again. Record `NO_WRITES` when a resumed task completes without output, so it isn't run again. Latest-state reads now omit interrupts from finished tasks. History and explicit checkpoint reads still show the last recorded interrupt per task. A task paused at its second interrupt appears in `next` with `result=None`. `Command(resume=value)` without an id now raises when more than one interrupt is pending, including in a subgraph. Older checkpoints where a resumed task finished without output may still appear unfinished. On SQLite and Postgres, insert-ignore can leave an old `RESUME` value for a finished task; that task does not run again. Tests: `tests/test_interrupt_state.py` (sync and async). `make format`, `make lint`; memory/SQLite: 2075 passed, 4 skipped; Postgres 16: 3659 passed, 6 skipped. --- libs/langgraph/langgraph/pregel/_loop.py | 46 +- libs/langgraph/langgraph/pregel/_runner.py | 6 +- .../langgraph/pregel/_task_status.py | 127 +++++ libs/langgraph/langgraph/pregel/debug.py | 51 +- libs/langgraph/langgraph/pregel/main.py | 76 +-- libs/langgraph/langgraph/types.py | 8 +- libs/langgraph/tests/test_interrupt_state.py | 525 ++++++++++++++++++ 7 files changed, 735 insertions(+), 104 deletions(-) create mode 100644 libs/langgraph/langgraph/pregel/_task_status.py create mode 100644 libs/langgraph/tests/test_interrupt_state.py diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 371037768..796fb5efc 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -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: diff --git a/libs/langgraph/langgraph/pregel/_runner.py b/libs/langgraph/langgraph/pregel/_runner.py index 4d53b9f9d..5e283716f 100644 --- a/libs/langgraph/langgraph/pregel/_runner.py +++ b/libs/langgraph/langgraph/pregel/_runner.py @@ -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] diff --git a/libs/langgraph/langgraph/pregel/_task_status.py b/libs/langgraph/langgraph/pregel/_task_status.py new file mode 100644 index 000000000..1110860a6 --- /dev/null +++ b/libs/langgraph/langgraph/pregel/_task_status.py @@ -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() + } diff --git a/libs/langgraph/langgraph/pregel/debug.py b/libs/langgraph/langgraph/pregel/debug.py index b3ef5a874..46f92eb5b 100644 --- a/libs/langgraph/langgraph/pregel/debug.py +++ b/libs/langgraph/langgraph/pregel/debug.py @@ -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) diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index f4ca48024..dc9f497be 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -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, diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 5e2a1a897..54510c8fc 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -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: diff --git a/libs/langgraph/tests/test_interrupt_state.py b/libs/langgraph/tests/test_interrupt_state.py new file mode 100644 index 000000000..28d2391df --- /dev/null +++ b/libs/langgraph/tests/test_interrupt_state.py @@ -0,0 +1,525 @@ +"""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 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 + + +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"]] + + +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 == () + + +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"]) +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