"""Single-flight, turn-guarded resume worker (design §3.3.1). When a human answer lands and wins the first-answer-wins compare-and-set (:func:`agent_team.db.schema.answer_question`), a resume job is enqueued to drive the suspended LangGraph thread forward. This module owns that resume mechanic, and the design pins three guarantees on it: * **Single-flight per ``thread_id``.** Two resume jobs for the same task never run concurrently. A task is its own thread; different threads resume concurrently, but one thread is serialized so a redelivered/duplicated job cannot race itself. * **Turn-guarded.** Before resuming, the worker reads the *live* checkpoint and confirms the graph is still interrupted on the answer's ``turn``. If the graph already advanced (a stale or redelivered job, or a resume that already applied), the worker marks the question ``superseded`` and skips. This is the mechanism by which **a resume can never double-apply** (§3.3.1). * **Restart-recoverable.** A startup sweep re-enqueues a resume for every ``answered`` ledger row whose graph is still interrupted on that turn; the turn guard makes re-enqueue idempotent, so converging after a reboot cannot double-apply either. The worker depends only on a small structural :class:`GraphLike` protocol (``get_state`` + ``invoke``), satisfied by a compiled LangGraph app, so the durable resume logic stays decoupled from any specific checkpointer and remains unit-testable without a live graph. The ledger reads/writes go through the committed :mod:`agent_team.db.schema` helpers (imported verbatim); this module adds no SQL of its own. """ from __future__ import annotations import sqlite3 import threading from dataclasses import dataclass from enum import Enum from typing import Any, Protocol, runtime_checkable from agent_team.db.schema import supersede_question __all__ = [ "GraphLike", "ResumeOutcome", "ResumeResult", "ResumeWorker", "build_resume_command", "snapshot_interrupt_turns", ] # --------------------------------------------------------------------------- # # Graph seam # --------------------------------------------------------------------------- # @runtime_checkable class GraphLike(Protocol): """Structural protocol for the compiled LangGraph app the worker drives. A real ``langgraph`` compiled graph satisfies this: ``get_state`` returns a ``StateSnapshot`` (with ``.next`` and ``.interrupts``) and ``invoke`` accepts a ``Command(resume=...)`` plus the thread config. Depending on the structural protocol rather than the concrete class keeps the resume logic decoupled from the checkpointer and trivially testable (§3.3.1). """ def get_state(self, config: dict[str, Any]) -> Any: """Return the live :class:`StateSnapshot` for ``config``'s thread.""" ... def invoke(self, input: Any, config: dict[str, Any]) -> Any: """Resume/run the graph for ``config``'s thread with ``input``.""" ... def build_resume_command(answer: Any) -> Any: """Build the LangGraph ``Command(resume=answer)`` resume input. ``Command`` is imported lazily so this module imports even where ``langgraph`` is absent (the durable ledger logic does not need it). When ``langgraph`` is installed, the real ``Command`` is used so the worker drives an actual compiled graph; otherwise a clear :class:`RuntimeError` is raised at call time. """ try: from langgraph.types import Command except ImportError as exc: # pragma: no cover - environment-dependent raise RuntimeError( "langgraph is required to resume a graph; install langgraph or " "inject a graph whose invoke() accepts a plain resume payload" ) from exc return Command(resume=answer) def snapshot_interrupt_turns(snapshot: Any) -> set[int]: """Extract the set of ``turn`` values the snapshot is interrupted on. The §3.3.1 interrupt payload is ``{thread_id, question_id, turn, ...}``. This reads each pending ``Interrupt.value`` and collects its ``turn``. A snapshot that is not interrupted (``snapshot.interrupts`` empty) yields an empty set, which the turn guard treats as "graph already advanced". Tolerant of either a mapping payload (``value['turn']``) or an object payload (``value.turn``); anything without a readable integer ``turn`` is ignored rather than crashing the worker. """ turns: set[int] = set() interrupts = getattr(snapshot, "interrupts", None) or () for item in interrupts: value = getattr(item, "value", item) turn: Any = None if isinstance(value, dict): turn = value.get("turn") else: turn = getattr(value, "turn", None) if isinstance(turn, bool): # bool is an int subclass; not a real turn continue if isinstance(turn, int): turns.add(turn) return turns def _snapshot_is_interrupted(snapshot: Any) -> bool: """True if the snapshot is suspended on an interrupt (``next`` non-empty). LangGraph reports a pending interrupt via a non-empty ``next`` tuple and a populated ``interrupts`` tuple. We treat either signal as "still interrupted"; the turn check then narrows it to *this* turn. """ if getattr(snapshot, "interrupts", None): return True nxt = getattr(snapshot, "next", None) return bool(nxt) # --------------------------------------------------------------------------- # # Result types # --------------------------------------------------------------------------- # class ResumeOutcome(Enum): """Outcome of a single :meth:`ResumeWorker.resume` attempt (§3.3.1).""" #: The graph was interrupted on this turn; ``Command(resume=...)`` applied. RESUMED = "resumed" #: The graph already advanced past this turn; question marked superseded, #: resume skipped. This is the no-double-apply guard firing. SUPERSEDED = "superseded" #: The graph already advanced but the question was no longer open/answered, #: so there was nothing to supersede; resume skipped. STALE = "stale" @dataclass class ResumeResult: """Structured result of a resume attempt. ``outcome`` is the :class:`ResumeOutcome`; ``thread_id`` / ``question_id`` / ``turn`` echo the job; ``graph_result`` carries the graph's return value when (and only when) the resume actually applied. """ outcome: ResumeOutcome thread_id: str question_id: str turn: int graph_result: Any = None @property def resumed(self) -> bool: """True iff the resume applied (``Command(resume=...)`` was invoked).""" return self.outcome is ResumeOutcome.RESUMED # --------------------------------------------------------------------------- # # Worker # --------------------------------------------------------------------------- # def _thread_config(thread_id: str) -> dict[str, Any]: """The LangGraph config addressing a single durable thread.""" return {"configurable": {"thread_id": thread_id}} class ResumeWorker: """Serializes and turn-guards graph resumes (§3.3.1 "single-flight"). One worker drives many threads; it holds a per-``thread_id`` lock registry so resumes for the *same* task are serialized (single-flight) while different tasks resume concurrently. ``graph`` is any :class:`GraphLike` (a compiled LangGraph app in production); ``conn`` is the agent-team SQLite connection (see :func:`agent_team.db.schema.connect`) used to read pending rows and to mark a stale question ``superseded`` via the committed compare-and-set helper. The worker performs no SQL of its own: lifecycle writes go through :func:`agent_team.db.schema.supersede_question`. It does not itself flip a question to ``answered`` — that is the responder's first-answer-wins compare-and-set, which gates whether a resume job is enqueued at all. """ def __init__(self, graph: GraphLike, conn: sqlite3.Connection) -> None: self._graph = graph self._conn = conn # Registry of per-thread locks. Guarded by _registry_lock so two # threads minting the lock for the same thread_id get the *same* lock. self._locks: dict[str, threading.Lock] = {} self._registry_lock = threading.Lock() def _lock_for(self, thread_id: str) -> threading.Lock: """Return the (memoized) serialization lock for ``thread_id``.""" with self._registry_lock: lock = self._locks.get(thread_id) if lock is None: lock = threading.Lock() self._locks[thread_id] = lock return lock def resume( self, *, thread_id: str, question_id: str, turn: int, answer: Any, ) -> ResumeResult: """Resume ``thread_id`` with ``answer``, single-flight + turn-guarded. Acquires the per-thread lock so this thread's resumes are serialized, then reads the live checkpoint. If the graph is still interrupted on ``turn`` it invokes ``Command(resume=answer)`` and returns :attr:`ResumeOutcome.RESUMED`. If the graph has already advanced past ``turn`` (stale/redelivered job, or a resume that already applied), it marks ``question_id`` ``superseded`` and skips — returning :attr:`ResumeOutcome.SUPERSEDED` (the row was open/answered) or :attr:`ResumeOutcome.STALE` (nothing left to supersede). A resume can therefore never double-apply (§3.3.1). """ lock = self._lock_for(thread_id) with lock: return self._resume_locked( thread_id=thread_id, question_id=question_id, turn=turn, answer=answer, ) def _resume_locked( self, *, thread_id: str, question_id: str, turn: int, answer: Any, ) -> ResumeResult: """Resume body that runs while holding this thread's lock.""" config = _thread_config(thread_id) snapshot = self._graph.get_state(config) interrupted_here = _snapshot_is_interrupted( snapshot ) and turn in snapshot_interrupt_turns(snapshot) if not interrupted_here: # Graph already advanced past this turn: the turn guard. Mark the # question superseded so it can never enqueue another resume, and # skip. supersede_question is the committed atomic compare-and-set; # rowcount 1 => we superseded it, 0 => it was already terminal. superseded = supersede_question(self._conn, question_id=question_id) outcome = ResumeOutcome.SUPERSEDED if superseded else ResumeOutcome.STALE return ResumeResult( outcome=outcome, thread_id=thread_id, question_id=question_id, turn=turn, ) graph_result = self._graph.invoke(build_resume_command(answer), config) return ResumeResult( outcome=ResumeOutcome.RESUMED, thread_id=thread_id, question_id=question_id, turn=turn, graph_result=graph_result, ) def recover_pending_resumes(self) -> list[ResumeResult]: """Restart sweep: re-enqueue resumes for durable ``answered`` rows. On reboot, in-memory resume jobs are gone but the ledger is durable (§3.3.1 "restart recovery"). This re-drives every ``answered`` question whose graph is still interrupted on its turn. The per-thread turn guard makes this idempotent: a thread that already advanced is superseded and skipped, so converging after a restart cannot double-apply. Returns one :class:`ResumeResult` per processed row (whatever its outcome) so a caller can log/ALARM. Rows are processed oldest-first by ``answered_at`` to preserve answer ordering across a recovery. """ rows = self._conn.execute( "SELECT question_id, thread_id, turn, answer_json " "FROM pending_questions " "WHERE status = 'answered' " "ORDER BY answered_at IS NULL, answered_at ASC" ).fetchall() results: list[ResumeResult] = [] for row in rows: answer = _decode_answer(row["answer_json"]) results.append( self.resume( thread_id=row["thread_id"], question_id=row["question_id"], turn=int(row["turn"]), answer=answer, ) ) return results def _decode_answer(answer_json: str | None) -> Any: """Decode a ledger ``answer_json`` payload back to a Python value. The responder stores answers as a JSON string in ``answer_json``. A non-JSON or ``NULL`` value is returned as-is (``None`` for ``NULL``), so a malformed row does not crash the recovery sweep — the turn guard still governs whether anything is applied. """ if answer_json is None: return None import json try: return json.loads(answer_json) except (ValueError, TypeError): return answer_json