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
Adam Moussa 9caf483117 fix(agent-team): harden listener respawn/close, broaden handle_event guard, channel_ref partial-unique
Three robustness/hardening fixes surfaced by /sh-security-review on the
agent-team listener/coordinator surface. None alter AUTHZ-01 allowlist
behavior or the first-answer-wins compare-and-set semantics.

1. Respawn close() leak (CWE-772). The watchdog _supervise_slack_listener
   respawned the inbound Slack listener without tearing down the dead one,
   leaking a Socket Mode WebSocket / SDK thread set per flap. Now the dead
   listener is closed before respawn (new _close_dead_listener, idempotent),
   AND _run_listener has a finally that always closes the listener so a
   crashed serve() releases its socket. SlackListener.close() is idempotent,
   so the belt-and-braces close stays a safe no-op.

2. Broadened exception guard in handle_event (CWE-248). The submit block
   only caught ValueError; the accept path (submit_answer ->
   _question_turn/_question_thread) can raise KeyError on a concurrently
   mutated row, and the CAS can raise sqlite3.Error. An uncaught exception
   would escape into the Bolt dispatch. Added a separate `except Exception`
   that logs at WARNING (not silent, not debug) and returns None. The
   existing ValueError-as-debug behavior is unchanged; authorization still
   runs first, so the trust boundary is not widened.

3. channel_ref partial-unique index (defense-in-depth). Added
   uq_pending_questions_open_channel_ref — a PARTIAL UNIQUE index on
   (channel_ref) WHERE channel_ref IS NOT NULL AND status='open' — so two
   OPEN rows can never share a non-null channel_ref (a thread_ts can never
   map to two open questions). Installed in init_db AND unconditionally in
   migrate (idempotent IF NOT EXISTS) so existing v1 DBs gain it. NULLs and
   closed rows are excluded; mirrored verbatim into schema.sql.

Tests: +8 (was 960, now 968). New: schema partial-unique reject/null/closed/
migrate cases; handle_event KeyError + sqlite3.Error swallow cases;
coordinator close-before-respawn + run_listener-closes-on-crash. Fixed the
operator-cli test fixture to use a per-question channel_ref (it previously
inserted multiple open rows sharing one ref, which the new index correctly
rejects).
2026-06-22 15:58:01 -04:00

465 lines
15 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,
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
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()