228 lines
6.6 KiB
Python
228 lines
6.6 KiB
Python
"""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,
|
|
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_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]
|