"""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, init_db, migrate, reopen_question, supersede_question, ) 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_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