"""Unit tests for agent_team.resume_worker (§3.3.1 single-flight, turn-guarded). These tests prove the design's three resume guarantees: * turn-guarded: resume applies only while the graph is interrupted on the answer's turn; a graph that already advanced is superseded and skipped; * no double-apply: a redelivered job for an already-advanced thread no-ops; * single-flight: resumes for one ``thread_id`` are serialized while different threads run concurrently; plus the restart-recovery sweep, and an end-to-end check against a real compiled LangGraph app when ``langgraph`` is importable. """ from __future__ import annotations import json import sqlite3 import threading import time from dataclasses import dataclass from pathlib import Path from typing import Any import pytest from agent_team import resume_worker from agent_team.db.schema import answer_question, init_db, connect from agent_team.resume_worker import ( GraphLike, ResumeOutcome, ResumeResult, ResumeWorker, build_resume_command, snapshot_interrupt_turns, ) # --------------------------------------------------------------------------- # # Fakes # --------------------------------------------------------------------------- # @dataclass class _FakeInterrupt: """Mimics a langgraph Interrupt: carries a ``.value`` payload.""" value: Any @dataclass class _FakeSnapshot: """Mimics a langgraph StateSnapshot's relevant surface.""" next: tuple[str, ...] = () interrupts: tuple[_FakeInterrupt, ...] = () class _FakeGraph: """A GraphLike test double over an explicit interrupt-turn. ``interrupted_turn`` is the turn the graph is currently suspended on, or ``None`` if it has advanced past every interrupt. ``invoke`` records every resume payload so double-apply is directly observable, and advances the graph (clears the interrupt) the way a real resume would. """ def __init__(self, interrupted_turn: int | None) -> None: self.interrupted_turn = interrupted_turn self.invocations: list[Any] = [] self.get_state_calls: list[dict[str, Any]] = [] self._invoke_hook: Any = None def get_state(self, config: dict[str, Any]) -> _FakeSnapshot: self.get_state_calls.append(config) if self.interrupted_turn is None: return _FakeSnapshot(next=(), interrupts=()) payload = {"turn": self.interrupted_turn, "question_id": "q"} return _FakeSnapshot( next=("clarify",), interrupts=(_FakeInterrupt(value=payload),), ) def invoke(self, command: Any, config: dict[str, Any]) -> Any: if self._invoke_hook is not None: self._invoke_hook() self.invocations.append(command) # A real resume clears the interrupt and advances the graph. self.interrupted_turn = None return {"resumed": True, "command": command} def _insert_question( conn: sqlite3.Connection, *, question_id: str, thread_id: str, turn: int, status: str = "open", ) -> None: conn.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport) " "VALUES (?, ?, ?, ?, 'slack')", (question_id, thread_id, turn, status), ) def _status(conn: sqlite3.Connection, question_id: str) -> str: row = conn.execute( "SELECT status FROM pending_questions WHERE question_id = ?", (question_id,), ).fetchone() return row["status"] @pytest.fixture() def conn(tmp_path: Path) -> sqlite3.Connection: db = tmp_path / "agent_team.sqlite" init_db(db) connection = connect(db) yield connection connection.close() # --------------------------------------------------------------------------- # # snapshot_interrupt_turns # --------------------------------------------------------------------------- # def test_snapshot_turns_from_mapping_payload() -> None: snap = _FakeSnapshot( next=("n",), interrupts=(_FakeInterrupt(value={"turn": 7}),), ) assert snapshot_interrupt_turns(snap) == {7} def test_snapshot_turns_from_object_payload() -> None: @dataclass class _ObjPayload: turn: int snap = _FakeSnapshot(interrupts=(_FakeInterrupt(value=_ObjPayload(turn=2)),)) assert snapshot_interrupt_turns(snap) == {2} def test_snapshot_turns_empty_when_not_interrupted() -> None: assert snapshot_interrupt_turns(_FakeSnapshot()) == set() def test_snapshot_turns_ignores_unreadable_turn() -> None: snap = _FakeSnapshot(interrupts=(_FakeInterrupt(value={"no_turn": 1}),)) assert snapshot_interrupt_turns(snap) == set() def test_snapshot_turns_ignores_bool_turn() -> None: # bool is an int subclass; a True/False must not be read as a turn number. snap = _FakeSnapshot(interrupts=(_FakeInterrupt(value={"turn": True}),)) assert snapshot_interrupt_turns(snap) == set() def test_snapshot_turns_collects_multiple() -> None: snap = _FakeSnapshot( interrupts=( _FakeInterrupt(value={"turn": 1}), _FakeInterrupt(value={"turn": 4}), ) ) assert snapshot_interrupt_turns(snap) == {1, 4} # --------------------------------------------------------------------------- # # resume — turn guard # --------------------------------------------------------------------------- # def test_resume_applies_when_interrupted_on_turn(conn: sqlite3.Connection) -> None: _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") graph = _FakeGraph(interrupted_turn=3) worker = ResumeWorker(graph, conn) result = worker.resume( thread_id="t1", question_id="q1", turn=3, answer="the answer" ) assert result.outcome is ResumeOutcome.RESUMED assert result.resumed is True assert result.graph_result == {"resumed": True, "command": graph.invocations[0]} assert len(graph.invocations) == 1 # Question is untouched by the worker on a successful resume (the responder # already flipped it to answered). assert _status(conn, "q1") == "answered" def test_resume_superseded_when_graph_advanced(conn: sqlite3.Connection) -> None: _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") graph = _FakeGraph(interrupted_turn=None) # already advanced past turn 3 worker = ResumeWorker(graph, conn) result = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="x") assert result.outcome is ResumeOutcome.SUPERSEDED assert result.resumed is False assert graph.invocations == [] # never invoked -> never applied assert _status(conn, "q1") == "superseded" def test_resume_superseded_when_interrupted_on_different_turn( conn: sqlite3.Connection, ) -> None: # Graph moved on to a *later* interrupt (turn 4); a turn-3 job is stale. _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") graph = _FakeGraph(interrupted_turn=4) worker = ResumeWorker(graph, conn) result = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="x") assert result.outcome is ResumeOutcome.SUPERSEDED assert graph.invocations == [] assert _status(conn, "q1") == "superseded" def test_resume_stale_when_nothing_to_supersede(conn: sqlite3.Connection) -> None: # Question already expired; graph advanced. Nothing to supersede. _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="expired") graph = _FakeGraph(interrupted_turn=None) worker = ResumeWorker(graph, conn) result = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="x") assert result.outcome is ResumeOutcome.STALE assert graph.invocations == [] assert _status(conn, "q1") == "expired" # --------------------------------------------------------------------------- # # no double-apply # --------------------------------------------------------------------------- # def test_redelivered_job_does_not_double_apply(conn: sqlite3.Connection) -> None: _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") graph = _FakeGraph(interrupted_turn=3) worker = ResumeWorker(graph, conn) first = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="a") # A redelivered/duplicate resume job for the same turn arrives. second = worker.resume(thread_id="t1", question_id="q1", turn=3, answer="a") assert first.outcome is ResumeOutcome.RESUMED assert second.outcome is ResumeOutcome.SUPERSEDED # The graph was invoked exactly once across both jobs. assert len(graph.invocations) == 1 assert _status(conn, "q1") == "superseded" # --------------------------------------------------------------------------- # # single-flight serialization # --------------------------------------------------------------------------- # def test_same_thread_resumes_are_serialized(conn: sqlite3.Connection) -> None: _insert_question(conn, question_id="q1", thread_id="t1", turn=3, status="answered") graph = _FakeGraph(interrupted_turn=3) in_invoke = threading.Event() overlap_detected: list[bool] = [] concurrency = {"current": 0, "max": 0} lock = threading.Lock() def _hook() -> None: with lock: concurrency["current"] += 1 concurrency["max"] = max(concurrency["max"], concurrency["current"]) in_invoke.set() time.sleep(0.05) with lock: concurrency["current"] -= 1 graph._invoke_hook = _hook worker = ResumeWorker(graph, conn) results: list[ResumeResult] = [] results_lock = threading.Lock() def _run(answer: str) -> None: r = worker.resume(thread_id="t1", question_id="q1", turn=3, answer=answer) with results_lock: results.append(r) threads = [threading.Thread(target=_run, args=(f"a{i}",)) for i in range(5)] for t in threads: t.start() for t in threads: t.join() # Lock serialized them: invoke never overlapped. assert concurrency["max"] == 1 assert overlap_detected == [] # Exactly one resumed; the rest were superseded (turn guard) -> no double. resumed = [r for r in results if r.outcome is ResumeOutcome.RESUMED] assert len(resumed) == 1 assert len(graph.invocations) == 1 def test_distinct_threads_use_distinct_locks(conn: sqlite3.Connection) -> None: graph_a = _FakeGraph(interrupted_turn=1) graph_b = _FakeGraph(interrupted_turn=1) # One worker can only hold one graph; emulate isolation by giving each # thread its own worker over its own graph, sharing the ledger. _insert_question(conn, question_id="qa", thread_id="ta", turn=1, status="answered") _insert_question(conn, question_id="qb", thread_id="tb", turn=1, status="answered") worker_a = ResumeWorker(graph_a, conn) worker_b = ResumeWorker(graph_b, conn) # Distinct thread_ids must mint distinct locks within a single worker. single = ResumeWorker(_FakeGraph(interrupted_turn=1), conn) assert single._lock_for("ta") is not single._lock_for("tb") assert single._lock_for("ta") is single._lock_for("ta") ra = worker_a.resume(thread_id="ta", question_id="qa", turn=1, answer="x") rb = worker_b.resume(thread_id="tb", question_id="qb", turn=1, answer="y") assert ra.outcome is ResumeOutcome.RESUMED assert rb.outcome is ResumeOutcome.RESUMED # --------------------------------------------------------------------------- # # restart recovery sweep # --------------------------------------------------------------------------- # def test_recover_resumes_answered_rows_still_interrupted( conn: sqlite3.Connection, ) -> None: _insert_question(conn, question_id="q1", thread_id="t1", turn=2, status="open") answer_question( conn, question_id="q1", answer_json=json.dumps("ans"), answered_via="slack" ) graph = _FakeGraph(interrupted_turn=2) worker = ResumeWorker(graph, conn) results = worker.recover_pending_resumes() assert len(results) == 1 assert results[0].outcome is ResumeOutcome.RESUMED assert results[0].thread_id == "t1" # The decoded answer reached the graph as a Command(resume=...). assert len(graph.invocations) == 1 def test_recover_is_idempotent_when_graph_already_advanced( conn: sqlite3.Connection, ) -> None: _insert_question(conn, question_id="q1", thread_id="t1", turn=2, status="open") answer_question( conn, question_id="q1", answer_json=json.dumps("ans"), answered_via="slack" ) # Graph already advanced (the resume applied before the crash). graph = _FakeGraph(interrupted_turn=None) worker = ResumeWorker(graph, conn) results = worker.recover_pending_resumes() assert len(results) == 1 assert results[0].outcome is ResumeOutcome.SUPERSEDED assert graph.invocations == [] # no double-apply across a restart assert _status(conn, "q1") == "superseded" def test_recover_skips_non_answered_rows(conn: sqlite3.Connection) -> None: _insert_question(conn, question_id="open1", thread_id="t1", turn=0, status="open") _insert_question(conn, question_id="exp1", thread_id="t2", turn=0, status="expired") graph = _FakeGraph(interrupted_turn=0) worker = ResumeWorker(graph, conn) results = worker.recover_pending_resumes() assert results == [] assert graph.invocations == [] def test_recover_processes_answered_oldest_first(conn: sqlite3.Connection) -> None: # Two answered rows on distinct threads; recovery must visit older first. _insert_question( conn, question_id="q_old", thread_id="t_old", turn=0, status="open" ) answer_question( conn, question_id="q_old", answer_json=json.dumps("old"), answered_via="slack", answered_at="2026-01-01T00:00:00+00:00", ) _insert_question( conn, question_id="q_new", thread_id="t_new", turn=0, status="open" ) answer_question( conn, question_id="q_new", answer_json=json.dumps("new"), answered_via="slack", answered_at="2026-06-01T00:00:00+00:00", ) graph = _FakeGraph(interrupted_turn=0) worker = ResumeWorker(graph, conn) results = worker.recover_pending_resumes() assert [r.thread_id for r in results] == ["t_old", "t_new"] # --------------------------------------------------------------------------- # # answer decoding # --------------------------------------------------------------------------- # def test_decode_answer_json_roundtrip() -> None: assert resume_worker._decode_answer(json.dumps({"k": 1})) == {"k": 1} def test_decode_answer_none() -> None: assert resume_worker._decode_answer(None) is None def test_decode_answer_non_json_passthrough() -> None: assert resume_worker._decode_answer("not-json{{") == "not-json{{" # --------------------------------------------------------------------------- # # resume_ci (CI machine-gate: turn-guarded on the awaiting-CI marker) # --------------------------------------------------------------------------- # class _CiGraph: """A GraphLike double suspended at VERIFY awaiting a given CI ``run_id``. ``awaiting_run`` is the run the thread is currently suspended on, or ``None`` if it has advanced past the CI gate (resumed / parked / done). ``invoke`` records every resume and advances the graph (clears the interrupt) as a real resume would, so a double-resume is directly observable. """ def __init__(self, awaiting_run: str | None) -> None: self.awaiting_run = awaiting_run self.invocations: list[Any] = [] def get_state(self, config: dict[str, Any]) -> _FakeSnapshot: if self.awaiting_run is None: return _FakeSnapshot(next=(), interrupts=()) payload = {"awaiting_ci": True, "run_id": self.awaiting_run} return _FakeSnapshot( next=("verify",), interrupts=(_FakeInterrupt(value=payload),), ) def invoke(self, command: Any, config: dict[str, Any]) -> Any: self.invocations.append(command) self.awaiting_run = None return {"resumed": True} def test_resume_ci_applies_when_suspended_on_run(conn: sqlite3.Connection) -> None: graph = _CiGraph(awaiting_run="999") worker = ResumeWorker(graph, conn) result = worker.resume_ci(thread_id="t1", run_id="999", answer={"conclusion": "ok"}) assert result.outcome is ResumeOutcome.RESUMED assert result.resumed is True assert len(graph.invocations) == 1 def test_resume_ci_skips_when_thread_already_advanced( conn: sqlite3.Connection, ) -> None: # Already resumed/parked/done: no awaiting-CI interrupt -> guard skips invoke. graph = _CiGraph(awaiting_run=None) worker = ResumeWorker(graph, conn) result = worker.resume_ci(thread_id="t1", run_id="999", answer={}) assert result.outcome is ResumeOutcome.STALE assert graph.invocations == [] def test_resume_ci_skips_when_awaiting_a_different_run( conn: sqlite3.Connection, ) -> None: # Suspended awaiting a DIFFERENT run (e.g. a re-dispatch): must not resume. graph = _CiGraph(awaiting_run="other") worker = ResumeWorker(graph, conn) result = worker.resume_ci(thread_id="t1", run_id="999", answer={}) assert result.outcome is ResumeOutcome.STALE assert graph.invocations == [] def test_resume_ci_double_resume_is_idempotent(conn: sqlite3.Connection) -> None: # Two terminal observations of the same run on overlapping sweeps: the first # applies; the second finds the thread advanced (guard) and skips. State is # never double-applied. graph = _CiGraph(awaiting_run="999") worker = ResumeWorker(graph, conn) first = worker.resume_ci(thread_id="t1", run_id="999", answer={}) second = worker.resume_ci(thread_id="t1", run_id="999", answer={}) assert first.outcome is ResumeOutcome.RESUMED assert second.outcome is ResumeOutcome.STALE assert len(graph.invocations) == 1 # --------------------------------------------------------------------------- # # command builder # --------------------------------------------------------------------------- # def test_build_resume_command_wraps_answer() -> None: pytest.importorskip("langgraph") cmd = build_resume_command("hello") assert getattr(cmd, "resume", None) == "hello" # --------------------------------------------------------------------------- # # module contract # --------------------------------------------------------------------------- # def test_module_exports_public_contract() -> None: for name in ( "GraphLike", "ResumeOutcome", "ResumeResult", "ResumeWorker", "build_resume_command", "snapshot_interrupt_turns", ): assert name in resume_worker.__all__ assert hasattr(resume_worker, name) def test_graphlike_is_runtime_checkable() -> None: assert isinstance(_FakeGraph(interrupted_turn=None), GraphLike) # --------------------------------------------------------------------------- # # end-to-end against a real compiled LangGraph app # --------------------------------------------------------------------------- # def test_end_to_end_against_real_langgraph(conn: sqlite3.Connection) -> None: pytest.importorskip("langgraph") from langgraph.graph import StateGraph, START, END from langgraph.checkpoint.memory import MemorySaver from langgraph.types import interrupt from typing import TypedDict class S(TypedDict, total=False): turn: int answer: Any def clarify(state: S) -> dict[str, Any]: ans = interrupt({"turn": state.get("turn", 0), "question_id": "q1"}) return {"answer": ans} g = StateGraph(S) g.add_node("clarify", clarify) g.add_edge(START, "clarify") g.add_edge("clarify", END) app = g.compile(checkpointer=MemorySaver()) cfg = {"configurable": {"thread_id": "real-1"}} app.invoke({"turn": 5}, cfg) # suspends on interrupt at turn 5 _insert_question( conn, question_id="q1", thread_id="real-1", turn=5, status="answered" ) worker = ResumeWorker(app, conn) first = worker.resume( thread_id="real-1", question_id="q1", turn=5, answer="confirmed" ) assert first.outcome is ResumeOutcome.RESUMED assert first.graph_result.get("answer") == "confirmed" # A redelivered job after the real graph advanced must not double-apply. second = worker.resume( thread_id="real-1", question_id="q1", turn=5, answer="confirmed" ) assert second.outcome is ResumeOutcome.SUPERSEDED assert _status(conn, "q1") == "superseded"