"""Unit tests for agent_team.db.schema (§3.3.1, §6.7).""" from __future__ import annotations import sqlite3 import threading from pathlib import Path import pytest from agent_team.db.schema import ( BUDGET_LEDGER_DDL, PENDING_QUESTIONS_DDL, QUESTION_STATES, SCHEMA_VERSION, answer_question, connect, expire_question, find_open_question_kind_by_channel_ref, init_db, issue_already_ingested, migrate, record_issue_ingested, reopen_question, supersede_question, ) def _tables(conn: sqlite3.Connection) -> set[str]: return { row[0] for row in conn.execute("SELECT name FROM sqlite_master WHERE type='table'") } def _insert_open_question(conn: sqlite3.Connection, qid: str, turn: int = 0) -> None: conn.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport) " "VALUES (?, 'thread-1', ?, 'open', 'slack')", (qid, turn), ) def test_ddl_constants_are_nonempty_strings() -> None: assert isinstance(PENDING_QUESTIONS_DDL, str) and PENDING_QUESTIONS_DDL assert isinstance(BUDGET_LEDGER_DDL, str) and BUDGET_LEDGER_DDL assert "pending_questions" in PENDING_QUESTIONS_DDL assert "budget_ledger" in BUDGET_LEDGER_DDL def test_schema_version_is_int() -> None: assert isinstance(SCHEMA_VERSION, int) def test_question_states_match_ddl_check() -> None: assert QUESTION_STATES == ("open", "answered", "expired", "superseded") for state in QUESTION_STATES: assert f"'{state}'" in PENDING_QUESTIONS_DDL def test_find_open_question_kind_by_channel_ref_returns_qid_and_kind( tmp_path: Path, ) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: conn.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport, channel_ref, kind) " "VALUES ('q1', 't1', 0, 'open', 'slack', 'TS.1', 'plan_decision')" ) conn.commit() assert find_open_question_kind_by_channel_ref(conn, "TS.1") == ( "q1", "plan_decision", ) # Empty ref / no match -> None. assert find_open_question_kind_by_channel_ref(conn, "") is None assert find_open_question_kind_by_channel_ref(conn, "NOPE") is None finally: conn.close() def test_find_open_question_kind_by_channel_ref_constrained_to_open( tmp_path: Path, ) -> None: """Anti-replay: a non-open row's channel_ref resolves to None.""" db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: conn.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport, channel_ref, kind) " "VALUES ('q1', 't1', 0, 'answered', 'slack', 'TS.1', 'plan_decision')" ) conn.commit() assert find_open_question_kind_by_channel_ref(conn, "TS.1") is None finally: conn.close() def test_connect_sets_pragmas(tmp_path: Path) -> None: conn = connect(tmp_path / "db.sqlite") try: assert conn.execute("PRAGMA journal_mode").fetchone()[0].lower() == "wal" assert conn.execute("PRAGMA foreign_keys").fetchone()[0] == 1 assert conn.execute("PRAGMA busy_timeout").fetchone()[0] >= 1 finally: conn.close() def test_init_db_creates_tables(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: names = { r[0] for r in conn.execute( "SELECT name FROM sqlite_master WHERE type='table'" ).fetchall() } finally: conn.close() assert {"pending_questions", "budget_ledger", "schema_meta"} <= names def test_init_db_is_idempotent(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) init_db(db) # must not raise conn = connect(db) try: version = conn.execute( "SELECT schema_version FROM schema_meta WHERE id=1" ).fetchone()[0] finally: conn.close() assert version == SCHEMA_VERSION def test_pending_questions_status_check_constraint(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: with pytest.raises(sqlite3.IntegrityError): conn.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport) " "VALUES ('q', 't', 0, 'bogus', 'slack')" ) finally: conn.close() def test_migrate_stamps_version(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" conn = connect(db) try: migrate(conn) version = conn.execute( "SELECT schema_version FROM schema_meta WHERE id=1" ).fetchone()[0] # Tables exist after migrate. conn.execute("SELECT 1 FROM pending_questions LIMIT 1") conn.execute("SELECT 1 FROM budget_ledger LIMIT 1") finally: conn.close() assert version == SCHEMA_VERSION def test_answer_question_first_wins(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: _insert_open_question(conn, "q1") first = answer_question( conn, question_id="q1", answer_json='{"a":1}', answered_via="slack" ) second = answer_question( conn, question_id="q1", answer_json='{"a":2}', answered_via="github" ) assert first is True assert second is False # duplicate/late loses the compare-and-set row = conn.execute( "SELECT status, answer_json, answered_via, answered_at " "FROM pending_questions WHERE question_id='q1'" ).fetchone() finally: conn.close() assert row["status"] == "answered" assert row["answer_json"] == '{"a":1}' # first answer retained assert row["answered_via"] == "slack" assert row["answered_at"] def test_answer_after_expire_loses(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: _insert_open_question(conn, "q2") assert expire_question(conn, question_id="q2") is True assert ( answer_question( conn, question_id="q2", answer_json="{}", answered_via="slack" ) is False ) status = conn.execute( "SELECT status FROM pending_questions WHERE question_id='q2'" ).fetchone()["status"] finally: conn.close() assert status == "expired" def test_expire_only_open(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: _insert_open_question(conn, "q3") answer_question(conn, question_id="q3", answer_json="{}", answered_via="slack") # already answered -> cannot expire assert expire_question(conn, question_id="q3") is False finally: conn.close() def test_supersede_open_or_answered(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: _insert_open_question(conn, "q4") answer_question(conn, question_id="q4", answer_json="{}", answered_via="slack") assert supersede_question(conn, question_id="q4") is True # already superseded -> no-op assert supersede_question(conn, question_id="q4") is False finally: conn.close() def test_reopen_question_unparks_expired_only(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: _insert_open_question(conn, "exp") _insert_open_question(conn, "ans") assert expire_question(conn, question_id="exp") is True assert answer_question( conn, question_id="ans", answer_json="{}", answered_via="t" ) # Expired -> reopened. assert reopen_question(conn, question_id="exp") is True row = conn.execute( "SELECT status, deadline_at FROM pending_questions WHERE question_id='exp'" ).fetchone() assert row["status"] == "open" assert row["deadline_at"] is None # no deadline until one is set # Answered row is NOT reopenable (only expired rows are). assert reopen_question(conn, question_id="ans") is False assert ( conn.execute( "SELECT status FROM pending_questions WHERE question_id='ans'" ).fetchone()["status"] == "answered" ) finally: conn.close() def test_concurrent_answers_single_winner(tmp_path: Path) -> None: """Two threads racing to answer the same open question: exactly one wins.""" db = tmp_path / "db.sqlite" init_db(db) seed = connect(db) try: _insert_open_question(seed, "race") finally: seed.close() results: list[bool] = [] barrier = threading.Barrier(2) lock = threading.Lock() def worker(via: str) -> None: conn = connect(db) try: barrier.wait() won = answer_question( conn, question_id="race", answer_json='{"v":1}', answered_via=via ) with lock: results.append(won) finally: conn.close() threads = [threading.Thread(target=worker, args=(f"c{i}",)) for i in range(2)] for t in threads: t.start() for t in threads: t.join() assert sorted(results) == [False, True] def test_shared_connection_concurrent_distinct_questions(tmp_path: Path) -> None: """Regression: many threads share ONE connection, answering DISTINCT questions. The responder and resume worker hold a single connection and call the CAS helpers from different threads concurrently (``connect()`` sets ``check_same_thread=False``). A single connection cannot hold two explicit ``BEGIN IMMEDIATE`` transactions at once, so the previous implementation raised "cannot start a transaction within a transaction" for all but one thread. The CAS now runs each write on its own private connection, so every distinct question is answered with no error. """ db = tmp_path / "db.sqlite" init_db(db) shared = connect(db) n = 8 try: for i in range(n): _insert_open_question(shared, f"q{i}") barrier = threading.Barrier(n) lock = threading.Lock() wins: list[bool] = [] errors: list[BaseException] = [] def worker(qid: str) -> None: try: barrier.wait() won = answer_question( shared, question_id=qid, answer_json='{"v":1}', answered_via="t" ) with lock: wins.append(won) except BaseException as exc: # noqa: BLE001 - record for assertion with lock: errors.append(exc) threads = [threading.Thread(target=worker, args=(f"q{i}",)) for i in range(n)] for t in threads: t.start() for t in threads: t.join() finally: shared.close() assert errors == [], f"shared-connection CAS raised: {errors!r}" assert wins == [True] * n def test_shared_connection_concurrent_same_question_single_winner( tmp_path: Path, ) -> None: """Regression: shared connection, many threads racing the SAME question. Exactly one first-answer-wins, the rest no-op (rowcount 0), and no thread raises a transaction-nesting or lock error. """ db = tmp_path / "db.sqlite" init_db(db) shared = connect(db) n = 8 try: _insert_open_question(shared, "race") barrier = threading.Barrier(n) lock = threading.Lock() wins: list[bool] = [] errors: list[BaseException] = [] def worker(via: str) -> None: try: barrier.wait() won = answer_question( shared, question_id="race", answer_json='{"v":1}', answered_via=via ) with lock: wins.append(won) except BaseException as exc: # noqa: BLE001 - record for assertion with lock: errors.append(exc) threads = [threading.Thread(target=worker, args=(f"c{i}",)) for i in range(n)] for t in threads: t.start() for t in threads: t.join() finally: shared.close() assert errors == [], f"shared-connection CAS raised: {errors!r}" assert sum(wins) == 1 assert wins.count(False) == n - 1 def _insert_open_question_with_ref( conn: sqlite3.Connection, qid: str, channel_ref: str | None, *, status: str = "open", turn: int = 0, ) -> None: conn.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport, channel_ref) " "VALUES (?, 'thread-1', ?, ?, 'slack', ?)", (qid, turn, status, channel_ref), ) def test_open_channel_ref_partial_unique_rejects_second_open_row( tmp_path: Path, ) -> None: """Defense-in-depth: two OPEN rows can never share a non-null channel_ref. The partial unique index ``uq_pending_questions_open_channel_ref`` makes a duplicate (channel_ref, status='open') pair impossible, so a thread reply's thread_ts can never resolve to two open questions. """ db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: _insert_open_question_with_ref(conn, "q1", "ts-100") with pytest.raises(sqlite3.IntegrityError): _insert_open_question_with_ref(conn, "q2", "ts-100") finally: conn.close() def test_open_channel_ref_partial_unique_allows_multiple_nulls( tmp_path: Path, ) -> None: """NULL channel_refs are unaffected: many open rows may have a null ref. The WHERE clause excludes nulls (and SQLite treats multiple NULLs as distinct in a unique index anyway), so the lost-post recovery path's unposted-open rows are never blocked. """ db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: _insert_open_question_with_ref(conn, "q1", None) _insert_open_question_with_ref(conn, "q2", None) # must not raise count = conn.execute( "SELECT COUNT(*) FROM pending_questions WHERE channel_ref IS NULL" ).fetchone()[0] finally: conn.close() assert count == 2 def test_open_channel_ref_partial_unique_ignores_closed_rows( tmp_path: Path, ) -> None: """Closed rows are excluded: a non-open row may share an open row's ref. The index predicate is ``status='open'``, so once a question is answered / expired / superseded its channel_ref no longer participates — a fresh open question can reuse it (e.g. a re-asked turn on the same thread anchor). """ db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: # An answered row and a superseded row both carry 'ts-200'. _insert_open_question_with_ref(conn, "ans", "ts-200", status="answered") _insert_open_question_with_ref(conn, "sup", "ts-200", status="superseded") # A NEW open row may still take 'ts-200' (no open row holds it). _insert_open_question_with_ref(conn, "open1", "ts-200") # must not raise # But a SECOND open row with the same ref is rejected. with pytest.raises(sqlite3.IntegrityError): _insert_open_question_with_ref(conn, "open2", "ts-200") finally: conn.close() def test_migrate_adds_open_channel_ref_partial_unique(tmp_path: Path) -> None: """migrate() (not just init_db) installs the partial unique index. Existing DBs picked up via migrate() must get the defense-in-depth index too, so the uniqueness guarantee holds after an in-place schema step. """ db = tmp_path / "db.sqlite" conn = connect(db) try: # Migrate twice: the second call runs against an already-stamped v1 DB, # exercising the UNCONDITIONAL index install (the `current < 1` block is # skipped), which is exactly the existing-DB upgrade path. migrate(conn) migrate(conn) names = { r[0] for r in conn.execute( "SELECT name FROM sqlite_master WHERE type='index'" ).fetchall() } assert "uq_pending_questions_open_channel_ref" in names # And it actually enforces: a duplicate open ref is rejected. _insert_open_question_with_ref(conn, "q1", "ts-300") with pytest.raises(sqlite3.IntegrityError): _insert_open_question_with_ref(conn, "q2", "ts-300") finally: conn.close() # --------------------------------------------------------------------------- # # ingested_issues durable de-dup (schema v2) # --------------------------------------------------------------------------- # def test_schema_version_is_at_least_2() -> None: assert SCHEMA_VERSION >= 2 def test_init_db_creates_ingested_issues(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: assert "ingested_issues" in _tables(conn) finally: conn.close() def test_issue_ingest_helpers_record_and_detect(tmp_path: Path) -> None: db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: src = "github:o/r" assert not issue_already_ingested(conn, source=src, issue_id="1") # first record inserts a new row assert record_issue_ingested(conn, source=src, issue_id="1") is True assert issue_already_ingested(conn, source=src, issue_id="1") # idempotent: a repeat record is a no-op (False) but still "seen" assert record_issue_ingested(conn, source=src, issue_id="1") is False assert issue_already_ingested(conn, source=src, issue_id="1") # source namespacing: same id under a different repo is independent assert not issue_already_ingested(conn, source="github:o/other", issue_id="1") finally: conn.close() def test_migrate_adds_ingested_issues_to_a_legacy_v1_db(tmp_path: Path) -> None: """A DB stamped at v1 (no ingested_issues) gains the table + a v2 stamp.""" db = tmp_path / "legacy.sqlite" conn = connect(db) try: # Simulate a legacy v1 DB: pending_questions + a schema_meta stamped at 1, # WITHOUT the v2 ingested_issues table. conn.execute(PENDING_QUESTIONS_DDL) conn.execute( "CREATE TABLE IF NOT EXISTS schema_meta " "(id INTEGER PRIMARY KEY CHECK (id = 1), schema_version INTEGER NOT NULL)" ) conn.execute("INSERT INTO schema_meta (id, schema_version) VALUES (1, 1)") assert "ingested_issues" not in _tables(conn) migrate(conn) assert "ingested_issues" in _tables(conn) ver = conn.execute( "SELECT schema_version FROM schema_meta WHERE id = 1" ).fetchone()[0] assert ver == SCHEMA_VERSION finally: conn.close() # --------------------------------------------------------------------------- # # pending_questions.kind discriminator (schema v4) # --------------------------------------------------------------------------- # # Legacy (pre-kind) pending_questions DDL, used to construct a DB whose table # predates the additive migration. _LEGACY_PENDING_QUESTIONS_DDL = """ CREATE TABLE IF NOT EXISTS pending_questions ( question_id TEXT PRIMARY KEY, thread_id TEXT NOT NULL, turn INTEGER NOT NULL, status TEXT NOT NULL CHECK (status IN ('open', 'answered', 'expired', 'superseded')), transport TEXT NOT NULL, channel_ref TEXT, posted_at TEXT, deadline_at TEXT, answer_json TEXT, answered_at TEXT, answered_via TEXT ) """.strip() def _pq_columns(conn: sqlite3.Connection) -> list[str]: return [r["name"] for r in conn.execute("PRAGMA table_info(pending_questions)")] def test_schema_version_is_at_least_4() -> None: assert SCHEMA_VERSION >= 4 def test_init_db_pending_questions_has_kind_defaulting_clarify( tmp_path: Path, ) -> None: """A fresh init_db gives pending_questions a kind column defaulting clarify.""" db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: assert "kind" in _pq_columns(conn) _insert_open_question(conn, "q-default") kind = conn.execute( "SELECT kind FROM pending_questions WHERE question_id='q-default'" ).fetchone()["kind"] finally: conn.close() assert kind == "clarify" def test_init_db_kind_is_idempotent(tmp_path: Path) -> None: """Running init_db twice does not error and kind exists exactly once.""" db = tmp_path / "db.sqlite" init_db(db) init_db(db) # must not raise (no duplicate-column error) conn = connect(db) try: cols = _pq_columns(conn) finally: conn.close() assert cols.count("kind") == 1 def test_migrate_adds_kind_to_legacy_db_rows_read_clarify(tmp_path: Path) -> None: """A legacy pending_questions (no kind) gains the column; old rows read clarify.""" db = tmp_path / "legacy.sqlite" conn = connect(db) try: # Build the OLD table by hand and seed a row, with NO kind column. conn.execute(_LEGACY_PENDING_QUESTIONS_DDL) conn.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport) " "VALUES ('legacy', 't', 0, 'open', 'slack')" ) assert "kind" not in _pq_columns(conn) init_db(db) assert "kind" in _pq_columns(conn) # The pre-existing row reads back as 'clarify' (NOT null). kind = conn.execute( "SELECT kind FROM pending_questions WHERE question_id='legacy'" ).fetchone()["kind"] finally: conn.close() assert kind == "clarify" def test_migrate_helper_adds_kind_to_legacy_db(tmp_path: Path) -> None: """migrate() (not just init_db) installs the v4 kind column on a legacy DB.""" db = tmp_path / "legacy2.sqlite" conn = connect(db) try: conn.execute(_LEGACY_PENDING_QUESTIONS_DDL) conn.execute( "CREATE TABLE IF NOT EXISTS schema_meta " "(id INTEGER PRIMARY KEY CHECK (id = 1), schema_version INTEGER NOT NULL)" ) conn.execute("INSERT INTO schema_meta (id, schema_version) VALUES (1, 3)") assert "kind" not in _pq_columns(conn) migrate(conn) assert "kind" in _pq_columns(conn) ver = conn.execute( "SELECT schema_version FROM schema_meta WHERE id = 1" ).fetchone()[0] finally: conn.close() assert ver == SCHEMA_VERSION def test_init_db_advances_existing_version_stamp(tmp_path: Path) -> None: """init_db (the daemon's only schema entry point) bumps a stale version stamp. Regression: init_db's own schema_meta write was ON CONFLICT DO NOTHING, so an already-stamped DB (e.g. an old v3 ledger) kept its stale version forever — the daemon calls init_db, never migrate(), so the stamp never advanced even though the column was ensured. init_db now drives migrate(), which upserts. """ db = tmp_path / "stale.sqlite" conn = connect(db) try: conn.execute(_LEGACY_PENDING_QUESTIONS_DDL) conn.execute( "CREATE TABLE IF NOT EXISTS schema_meta " "(id INTEGER PRIMARY KEY CHECK (id = 1), schema_version INTEGER NOT NULL)" ) conn.execute("INSERT INTO schema_meta (id, schema_version) VALUES (1, 3)") finally: conn.close() init_db(db) conn = connect(db) try: ver = conn.execute( "SELECT schema_version FROM schema_meta WHERE id = 1" ).fetchone()[0] assert "kind" in _pq_columns(conn) finally: conn.close() assert ver == SCHEMA_VERSION def test_kind_plan_decision_round_trips(tmp_path: Path) -> None: """A row written with kind='plan_decision' round-trips; default is 'clarify'.""" db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: conn.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport, kind) " "VALUES ('pd', 't', 0, 'open', 'slack', 'plan_decision')" ) _insert_open_question(conn, "cl") # no kind -> default pd_kind = conn.execute( "SELECT kind FROM pending_questions WHERE question_id='pd'" ).fetchone()["kind"] cl_kind = conn.execute( "SELECT kind FROM pending_questions WHERE question_id='cl'" ).fetchone()["kind"] finally: conn.close() assert pd_kind == "plan_decision" assert cl_kind == "clarify" def test_kind_check_rejects_unknown_value(tmp_path: Path) -> None: """The CHECK constraint on a fresh DB rejects an out-of-range kind.""" db = tmp_path / "db.sqlite" init_db(db) conn = connect(db) try: with pytest.raises(sqlite3.IntegrityError): conn.execute( "INSERT INTO pending_questions " "(question_id, thread_id, turn, status, transport, kind) " "VALUES ('bad', 't', 0, 'open', 'slack', 'bogus')" ) finally: conn.close()