This repository has been archived on 2026-08-04. You can view files and clone it, but cannot push or open issues or pull requests.
orchestrator/agent-team/tests/test_schema.py

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]