"""Unit tests for agent_team.recovery — the restart-recovery sweep (§3.3.1, §6.7). Exercises the three convergence steps (deadline / redeliver / resume-or-supersede), their idempotency, the first-answer-wins races, post-restore reconciliation, and per-row error isolation. Uses the committed foundation contracts verbatim (``agent_team.db.schema`` for the ledger, ``agent_team.transport.base`` for the transport ABC) — nothing here redefines a foundation interface. """ from __future__ import annotations import json from datetime import datetime, timedelta, timezone from pathlib import Path from typing import Any import pytest from agent_team.db.schema import ( answer_question, connect, init_db, ) from agent_team.recovery import ( DeadlineOutcome, PendingQuestion, RecoveryReport, apply_deadline_policy, load_pending_questions, redeliver_open_questions, reenqueue_answered_resumes, run_restart_recovery, ) from agent_team.transport.base import QuestionSet, Transport UTC = timezone.utc # --------------------------------------------------------------------------- # # Fixtures & helpers # # --------------------------------------------------------------------------- # @pytest.fixture() def db_path(tmp_path: Path) -> Path: """A freshly initialized agent-team DB file.""" path = tmp_path / "agent_team.db" init_db(path) return path @pytest.fixture() def conn(db_path: Path): """An open connection to the initialized DB (closed at teardown).""" connection = connect(db_path) yield connection connection.close() def _insert( connection, *, question_id: str, thread_id: str = "t1", turn: int = 0, status: str = "open", transport: str = "slack", channel_ref: str | None = None, posted_at: str | None = None, deadline_at: str | None = None, answer_json: str | None = None, answered_via: str | None = None, ) -> None: """Insert a raw ``pending_questions`` row for a test scenario.""" connection.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport, channel_ref, " " posted_at, deadline_at, answer_json, answered_at, answered_via) " "VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", ( question_id, thread_id, turn, status, transport, channel_ref, posted_at, deadline_at, answer_json, None, answered_via, ), ) def _status(connection, question_id: str) -> str: row = connection.execute( "SELECT status FROM pending_questions WHERE question_id = ?", (question_id,), ).fetchone() return row["status"] def _channel_ref(connection, question_id: str) -> str | None: row = connection.execute( "SELECT channel_ref FROM pending_questions WHERE question_id = ?", (question_id,), ).fetchone() return row["channel_ref"] class RecordingTransport(Transport): """A Transport that records posts and returns a deterministic channel_ref.""" def __init__(self, ref: str = "ts-123", *, fail: bool = False) -> None: self.ref = ref self.fail = fail self.posts: list[dict[str, Any]] = [] def post_question( self, *, thread_id: str, question_id: str, turn: int, question_set: QuestionSet, deadline: str, ) -> str: if self.fail: raise RuntimeError("transport down") self.posts.append( { "thread_id": thread_id, "question_id": question_id, "turn": turn, "deadline": deadline, } ) return self.ref def parse_answer(self, raw: Any) -> tuple[str, Any, str]: # pragma: no cover raise NotImplementedError def _park_policy(question: PendingQuestion) -> DeadlineOutcome: return DeadlineOutcome(question_id=question.question_id, action="parked") # --------------------------------------------------------------------------- # # load_pending_questions / PendingQuestion # # --------------------------------------------------------------------------- # def test_load_pending_questions_filters_by_status(conn) -> None: _insert(conn, question_id="q-open", status="open") _insert(conn, question_id="q-ans", status="answered") opens = load_pending_questions(conn, status="open") assert [q.question_id for q in opens] == ["q-open"] def test_load_pending_questions_no_filter_returns_all(conn) -> None: _insert(conn, question_id="q1", status="open") _insert(conn, question_id="q2", status="answered") assert len(load_pending_questions(conn)) == 2 def test_load_pending_questions_rejects_unknown_status(conn) -> None: with pytest.raises(ValueError): load_pending_questions(conn, status="bogus") def test_pending_question_from_row_maps_columns(conn) -> None: _insert( conn, question_id="q1", thread_id="thread-x", turn=3, status="open", transport="github", channel_ref="ref-1", deadline_at="2026-01-01T00:00:00+00:00", ) (q,) = load_pending_questions(conn, status="open") assert q == PendingQuestion( question_id="q1", thread_id="thread-x", turn=3, status="open", transport="github", channel_ref="ref-1", posted_at=None, deadline_at="2026-01-01T00:00:00+00:00", answer_json=None, answered_at=None, answered_via=None, ) # --------------------------------------------------------------------------- # # Step 1 — redeliver lost posts # # --------------------------------------------------------------------------- # def test_redeliver_posts_open_row_without_ref(conn) -> None: _insert(conn, question_id="q1", status="open", channel_ref=None) transport = RecordingTransport(ref="slack-ts-9") report = RecoveryReport() redeliver_open_questions( conn, resolve_transport=lambda _t: transport, report=report ) assert report.redelivered == ["q1"] assert len(transport.posts) == 1 assert _channel_ref(conn, "q1") == "slack-ts-9" def test_redeliver_skips_row_that_already_has_ref(conn) -> None: _insert(conn, question_id="q1", status="open", channel_ref="already-here") transport = RecordingTransport() report = RecoveryReport() redeliver_open_questions( conn, resolve_transport=lambda _t: transport, report=report ) assert report.redelivered == [] assert transport.posts == [] def test_redeliver_defers_when_transport_unreachable(conn) -> None: _insert(conn, question_id="q1", status="open", channel_ref=None) report = RecoveryReport() redeliver_open_questions(conn, resolve_transport=lambda _t: None, report=report) assert report.redelivery_deferred == ["q1"] assert report.redelivered == [] # Row stays open with no ref so a later sweep retries. assert _status(conn, "q1") == "open" assert _channel_ref(conn, "q1") is None def test_redeliver_isolates_transport_exception(conn) -> None: _insert(conn, question_id="q1", status="open", channel_ref=None) transport = RecordingTransport(fail=True) report = RecoveryReport() redeliver_open_questions( conn, resolve_transport=lambda _t: transport, report=report ) assert report.redelivery_deferred == ["q1"] assert report.errors and report.errors[0][0] == "q1" def test_redeliver_treats_empty_ref_as_deferred(conn) -> None: _insert(conn, question_id="q1", status="open", channel_ref=None) transport = RecordingTransport(ref="") report = RecoveryReport() redeliver_open_questions( conn, resolve_transport=lambda _t: transport, report=report ) assert report.redelivery_deferred == ["q1"] assert _channel_ref(conn, "q1") is None # --------------------------------------------------------------------------- # # Step 2 — re-enqueue resumes / supersede # # --------------------------------------------------------------------------- # def test_reenqueue_resumes_when_graph_still_interrupted(conn) -> None: _insert(conn, question_id="q1", thread_id="t1", turn=2, status="answered") enqueued: list[tuple[str, str, int]] = [] report = RecoveryReport() reenqueue_answered_resumes( conn, is_interrupted_on_turn=lambda _tid, _turn: True, enqueue_resume=lambda tid, qid, turn: enqueued.append((tid, qid, turn)) or True, report=report, ) assert report.resumes_enqueued == ["q1"] assert enqueued == [("t1", "q1", 2)] # Row remains answered — the resume worker owns the terminal transition. assert _status(conn, "q1") == "answered" def test_reenqueue_supersedes_when_graph_advanced(conn) -> None: _insert(conn, question_id="q1", thread_id="t1", turn=2, status="answered") enqueued: list[Any] = [] report = RecoveryReport() reenqueue_answered_resumes( conn, is_interrupted_on_turn=lambda _tid, _turn: False, enqueue_resume=lambda *a: enqueued.append(a) or True, report=report, ) assert report.superseded == ["q1"] assert report.resumes_enqueued == [] assert enqueued == [] assert _status(conn, "q1") == "superseded" def test_reenqueue_probe_receives_thread_and_turn(conn) -> None: _insert(conn, question_id="q1", thread_id="thread-9", turn=7, status="answered") seen: list[tuple[str, int]] = [] reenqueue_answered_resumes( conn, is_interrupted_on_turn=lambda tid, turn: seen.append((tid, turn)) or True, enqueue_resume=lambda *a: True, report=RecoveryReport(), ) assert seen == [("thread-9", 7)] def test_reenqueue_isolates_probe_exception(conn) -> None: _insert(conn, question_id="q1", status="answered") def boom(_tid: str, _turn: int) -> bool: raise RuntimeError("probe failed") report = RecoveryReport() reenqueue_answered_resumes( conn, is_interrupted_on_turn=boom, enqueue_resume=lambda *a: True, report=report, ) assert report.errors and "probe" in report.errors[0][1] assert report.resumes_enqueued == [] def test_reenqueue_isolates_enqueue_exception(conn) -> None: _insert(conn, question_id="q1", status="answered") def boom(*_a: Any) -> bool: raise RuntimeError("queue down") report = RecoveryReport() reenqueue_answered_resumes( conn, is_interrupted_on_turn=lambda *a: True, enqueue_resume=boom, report=report, ) assert report.errors and "resume" in report.errors[0][1] # --------------------------------------------------------------------------- # # Step 2 — post-restore reconciliation (§6.7) # # --------------------------------------------------------------------------- # class _Reconciler: def __init__(self, *, safe: bool = True, raise_exc: bool = False) -> None: self.safe = safe self.raise_exc = raise_exc self.calls: list[str] = [] def reconcile(self, question: PendingQuestion) -> bool: self.calls.append(question.question_id) if self.raise_exc: raise RuntimeError("reconcile blew up") return self.safe def test_reconciler_allows_resume_when_external_state_consistent(conn) -> None: _insert(conn, question_id="q1", status="answered") reconciler = _Reconciler(safe=True) enqueued: list[Any] = [] report = RecoveryReport() reenqueue_answered_resumes( conn, is_interrupted_on_turn=lambda *a: True, enqueue_resume=lambda *a: enqueued.append(a) or True, report=report, reconciler=reconciler, ) assert reconciler.calls == ["q1"] assert report.resumes_enqueued == ["q1"] def test_reconciler_holds_resume_when_external_state_diverged(conn) -> None: _insert(conn, question_id="q1", status="answered") reconciler = _Reconciler(safe=False) enqueued: list[Any] = [] report = RecoveryReport() reenqueue_answered_resumes( conn, is_interrupted_on_turn=lambda *a: True, enqueue_resume=lambda *a: enqueued.append(a) or True, report=report, reconciler=reconciler, ) assert report.reconcile_held == ["q1"] assert report.resumes_enqueued == [] assert enqueued == [] def test_reconciler_exception_holds_and_records_error(conn) -> None: _insert(conn, question_id="q1", status="answered") reconciler = _Reconciler(raise_exc=True) report = RecoveryReport() reenqueue_answered_resumes( conn, is_interrupted_on_turn=lambda *a: True, enqueue_resume=lambda *a: True, report=report, reconciler=reconciler, ) assert report.reconcile_held == ["q1"] assert report.errors and "reconcile" in report.errors[0][1] def test_reconciler_not_consulted_when_graph_advanced(conn) -> None: _insert(conn, question_id="q1", status="answered") reconciler = _Reconciler(safe=True) reenqueue_answered_resumes( conn, is_interrupted_on_turn=lambda *a: False, enqueue_resume=lambda *a: True, report=RecoveryReport(), reconciler=reconciler, ) # Superseded path never reaches reconciliation. assert reconciler.calls == [] # --------------------------------------------------------------------------- # # Step 3 — deadline policy # # --------------------------------------------------------------------------- # def test_deadline_expires_overdue_open_question(conn) -> None: past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() _insert(conn, question_id="q1", status="open", deadline_at=past) report = RecoveryReport() apply_deadline_policy(conn, policy=_park_policy, report=report) assert report.expired == ["q1"] assert _status(conn, "q1") == "expired" assert report.deadline_outcomes[0].action == "parked" def test_deadline_leaves_future_question_open(conn) -> None: future = (datetime.now(UTC) + timedelta(hours=1)).isoformat() _insert(conn, question_id="q1", status="open", deadline_at=future) report = RecoveryReport() apply_deadline_policy(conn, policy=_park_policy, report=report) assert report.expired == [] assert _status(conn, "q1") == "open" def test_deadline_ignores_row_without_deadline(conn) -> None: _insert(conn, question_id="q1", status="open", deadline_at=None) report = RecoveryReport() apply_deadline_policy(conn, policy=_park_policy, report=report) assert report.expired == [] assert _status(conn, "q1") == "open" def test_deadline_respects_injected_now(conn) -> None: deadline = "2026-06-01T00:00:00+00:00" _insert(conn, question_id="q1", status="open", deadline_at=deadline) before = datetime(2026, 5, 1, tzinfo=UTC) report = RecoveryReport() apply_deadline_policy(conn, policy=_park_policy, report=report, now=before) assert report.expired == [] # not yet overdue at injected now assert _status(conn, "q1") == "open" def test_deadline_treats_naive_timestamp_as_utc(conn) -> None: past_naive = ( (datetime.now(UTC) - timedelta(hours=2)).replace(tzinfo=None).isoformat() ) _insert(conn, question_id="q1", status="open", deadline_at=past_naive) report = RecoveryReport() apply_deadline_policy(conn, policy=_park_policy, report=report) assert report.expired == ["q1"] def test_deadline_ignores_unparseable_timestamp(conn) -> None: _insert(conn, question_id="q1", status="open", deadline_at="not-a-date") report = RecoveryReport() apply_deadline_policy(conn, policy=_park_policy, report=report) assert report.expired == [] assert _status(conn, "q1") == "open" def test_deadline_policy_not_invoked_when_already_answered(conn) -> None: # An answer that won the race before the sweep: compare-and-set finds no # open row, so no expiry and no policy call. past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() _insert(conn, question_id="q1", status="answered", deadline_at=past) invoked: list[str] = [] report = RecoveryReport() apply_deadline_policy( conn, policy=lambda q: ( invoked.append(q.question_id) or DeadlineOutcome(q.question_id, "parked") ), report=report, ) assert report.expired == [] assert invoked == [] def test_deadline_isolates_policy_exception(conn) -> None: past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() _insert(conn, question_id="q1", status="open", deadline_at=past) def boom(_q: PendingQuestion) -> DeadlineOutcome: raise RuntimeError("policy failed") report = RecoveryReport() apply_deadline_policy(conn, policy=boom, report=report) # Row still durably expired even though the policy callback failed. assert report.expired == ["q1"] assert _status(conn, "q1") == "expired" assert report.errors and "policy" in report.errors[0][1] # --------------------------------------------------------------------------- # # Full sweep orchestration # # --------------------------------------------------------------------------- # def test_run_restart_recovery_drives_all_three_steps(db_path: Path) -> None: setup = connect(db_path) past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() # 1: overdue open -> expired _insert(setup, question_id="q-late", status="open", deadline_at=past) # 2: open w/o ref -> redelivered _insert(setup, question_id="q-lost", status="open", channel_ref=None) # 3: answered, graph still waits -> resume enqueued _insert(setup, question_id="q-ans", thread_id="ta", turn=1, status="answered") setup.close() transport = RecordingTransport(ref="ts-x") enqueued: list[Any] = [] report = run_restart_recovery( db_path, resolve_transport=lambda _t: transport, is_interrupted_on_turn=lambda *a: True, enqueue_resume=lambda *a: enqueued.append(a) or True, deadline_policy=_park_policy, ) assert report.expired == ["q-late"] assert report.redelivered == ["q-lost"] assert report.resumes_enqueued == ["q-ans"] assert not report.clean def test_run_restart_recovery_expires_before_redelivering(db_path: Path) -> None: # An overdue open row must be expired by step 1, never redelivered by # step 2 — proving deadline-first ordering. setup = connect(db_path) past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() _insert( setup, question_id="q1", status="open", channel_ref=None, deadline_at=past, ) setup.close() transport = RecordingTransport() report = run_restart_recovery( db_path, resolve_transport=lambda _t: transport, is_interrupted_on_turn=lambda *a: True, enqueue_resume=lambda *a: True, deadline_policy=_park_policy, ) assert report.expired == ["q1"] assert report.redelivered == [] assert transport.posts == [] # never posted an already-expired question def test_run_restart_recovery_clean_when_nothing_pending(db_path: Path) -> None: report = run_restart_recovery( db_path, resolve_transport=lambda _t: RecordingTransport(), is_interrupted_on_turn=lambda *a: True, enqueue_resume=lambda *a: True, deadline_policy=_park_policy, ) assert report.clean def test_run_restart_recovery_with_injected_conn(conn) -> None: _insert(conn, question_id="q1", status="open", channel_ref=None) report = run_restart_recovery( ":memory:", # ignored because conn is injected resolve_transport=lambda _t: RecordingTransport(ref="r"), is_interrupted_on_turn=lambda *a: True, enqueue_resume=lambda *a: True, deadline_policy=_park_policy, conn=conn, ) assert report.redelivered == ["q1"] # Injected connection is left open for the caller. assert _status(conn, "q1") == "open" def test_run_restart_recovery_post_restore_holds_diverged_task(db_path: Path) -> None: setup = connect(db_path) _insert(setup, question_id="q1", status="answered") setup.close() reconciler = _Reconciler(safe=False) enqueued: list[Any] = [] report = run_restart_recovery( db_path, resolve_transport=lambda _t: RecordingTransport(), is_interrupted_on_turn=lambda *a: True, enqueue_resume=lambda *a: enqueued.append(a) or True, deadline_policy=_park_policy, reconciler=reconciler, ) assert report.reconcile_held == ["q1"] assert enqueued == [] # --------------------------------------------------------------------------- # # Idempotency — re-running the sweep converges, doesn't duplicate # # --------------------------------------------------------------------------- # def test_sweep_is_idempotent_for_redelivery(db_path: Path) -> None: setup = connect(db_path) _insert(setup, question_id="q1", status="open", channel_ref=None) setup.close() kwargs: dict[str, Any] = dict( resolve_transport=lambda _t: RecordingTransport(ref="r1"), is_interrupted_on_turn=lambda *a: True, enqueue_resume=lambda *a: True, deadline_policy=_park_policy, ) first = run_restart_recovery(db_path, **kwargs) second = run_restart_recovery(db_path, **kwargs) assert first.redelivered == ["q1"] # Second pass: row now has a ref, so nothing to redeliver -> clean. assert second.redelivered == [] assert second.clean def test_first_answer_wins_against_concurrent_expiry(conn) -> None: # A real first-answer-wins race: answer lands, then the deadline sweep # runs. The compare-and-set protects the answered row from expiry. past = (datetime.now(UTC) - timedelta(hours=1)).isoformat() _insert(conn, question_id="q1", status="open", deadline_at=past) won = answer_question( conn, question_id="q1", answer_json=json.dumps({"ok": True}), answered_via="slack", ) assert won is True report = RecoveryReport() apply_deadline_policy(conn, policy=_park_policy, report=report) assert report.expired == [] assert _status(conn, "q1") == "answered" # --------------------------------------------------------------------------- # # RecoveryReport.clean # # --------------------------------------------------------------------------- # def test_report_clean_true_for_empty_report() -> None: assert RecoveryReport().clean is True @pytest.mark.parametrize( "field_name", [ "redelivered", "redelivery_deferred", "resumes_enqueued", "superseded", "expired", "reconcile_held", ], ) def test_report_not_clean_when_any_action_list_populated(field_name: str) -> None: report = RecoveryReport() getattr(report, field_name).append("q1") assert report.clean is False def test_report_not_clean_when_errors_present() -> None: report = RecoveryReport() report.errors.append(("q1", "boom")) assert report.clean is False