505 lines
16 KiB
Python
505 lines
16 KiB
Python
|
|
"""Unit tests for agent_team.deadline_timer (design §3.3.1).
|
||
|
|
|
||
|
|
Covers the deadline / no-answer timer loop: overdue selection, the
|
||
|
|
deterministic answer-vs-timeout race on the ``open`` -> ``expired`` flip, the
|
||
|
|
PARK and DEFAULT_ANSWER policies, restart idempotency, side-effect isolation,
|
||
|
|
and the concurrent responder-vs-timer race.
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import sqlite3
|
||
|
|
import threading
|
||
|
|
from datetime import datetime, timedelta, timezone
|
||
|
|
from pathlib import Path
|
||
|
|
|
||
|
|
import pytest
|
||
|
|
|
||
|
|
from agent_team.db.schema import answer_question, connect, init_db
|
||
|
|
from agent_team.deadline_timer import (
|
||
|
|
DeadlinePolicy,
|
||
|
|
ExpiryAction,
|
||
|
|
OverdueQuestion,
|
||
|
|
TimerLoopReport,
|
||
|
|
overdue_open_questions,
|
||
|
|
run_deadline_timer,
|
||
|
|
)
|
||
|
|
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
# Helpers
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
|
||
|
|
_PAST = "2000-01-01T00:00:00+00:00"
|
||
|
|
_FUTURE = "2999-01-01T00:00:00+00:00"
|
||
|
|
|
||
|
|
|
||
|
|
def _iso(dt: datetime) -> str:
|
||
|
|
return dt.isoformat()
|
||
|
|
|
||
|
|
|
||
|
|
def _insert_question(
|
||
|
|
conn: sqlite3.Connection,
|
||
|
|
qid: str,
|
||
|
|
*,
|
||
|
|
thread_id: str = "thread-1",
|
||
|
|
turn: int = 0,
|
||
|
|
status: str = "open",
|
||
|
|
transport: str = "slack",
|
||
|
|
deadline_at: str | None = _PAST,
|
||
|
|
) -> None:
|
||
|
|
conn.execute(
|
||
|
|
"INSERT INTO pending_questions "
|
||
|
|
"(question_id, thread_id, turn, status, transport, posted_at, deadline_at) "
|
||
|
|
"VALUES (?, ?, ?, ?, ?, ?, ?)",
|
||
|
|
(qid, thread_id, turn, status, transport, _PAST, deadline_at),
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.fixture()
|
||
|
|
def conn(tmp_path: Path) -> sqlite3.Connection:
|
||
|
|
db = tmp_path / "db.sqlite"
|
||
|
|
init_db(db)
|
||
|
|
c = connect(db)
|
||
|
|
yield c
|
||
|
|
c.close()
|
||
|
|
|
||
|
|
|
||
|
|
def _status(conn: sqlite3.Connection, qid: str) -> str:
|
||
|
|
return conn.execute(
|
||
|
|
"SELECT status FROM pending_questions WHERE question_id=?", (qid,)
|
||
|
|
).fetchone()["status"]
|
||
|
|
|
||
|
|
|
||
|
|
class _Collector:
|
||
|
|
"""Records the questions handed to an injected side-effect callback."""
|
||
|
|
|
||
|
|
def __init__(self) -> None:
|
||
|
|
self.calls: list[OverdueQuestion] = []
|
||
|
|
|
||
|
|
def __call__(self, question: OverdueQuestion) -> None:
|
||
|
|
self.calls.append(question)
|
||
|
|
|
||
|
|
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
# overdue_open_questions
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
|
||
|
|
|
||
|
|
def test_overdue_selects_only_past_open_with_deadline(conn: sqlite3.Connection) -> None:
|
||
|
|
_insert_question(conn, "past", deadline_at=_PAST)
|
||
|
|
_insert_question(conn, "future", deadline_at=_FUTURE)
|
||
|
|
_insert_question(conn, "no-deadline", deadline_at=None)
|
||
|
|
_insert_question(conn, "answered", status="answered", deadline_at=_PAST)
|
||
|
|
_insert_question(conn, "expired", status="expired", deadline_at=_PAST)
|
||
|
|
|
||
|
|
overdue = overdue_open_questions(conn)
|
||
|
|
assert [q.question_id for q in overdue] == ["past"]
|
||
|
|
|
||
|
|
|
||
|
|
def test_overdue_uses_now_cutoff(conn: sqlite3.Connection) -> None:
|
||
|
|
now = datetime(2026, 6, 17, tzinfo=timezone.utc)
|
||
|
|
just_past = _iso(now - timedelta(seconds=1))
|
||
|
|
just_future = _iso(now + timedelta(seconds=1))
|
||
|
|
_insert_question(conn, "before", deadline_at=just_past)
|
||
|
|
_insert_question(conn, "after", deadline_at=just_future)
|
||
|
|
|
||
|
|
overdue = overdue_open_questions(conn, now=_iso(now))
|
||
|
|
assert [q.question_id for q in overdue] == ["before"]
|
||
|
|
|
||
|
|
|
||
|
|
def test_overdue_includes_deadline_equal_to_now(conn: sqlite3.Connection) -> None:
|
||
|
|
now = "2026-06-17T00:00:00+00:00"
|
||
|
|
_insert_question(conn, "exact", deadline_at=now)
|
||
|
|
overdue = overdue_open_questions(conn, now=now)
|
||
|
|
assert [q.question_id for q in overdue] == ["exact"]
|
||
|
|
|
||
|
|
|
||
|
|
def test_overdue_ordered_oldest_deadline_first(conn: sqlite3.Connection) -> None:
|
||
|
|
_insert_question(conn, "newer", deadline_at="2010-01-01T00:00:00+00:00")
|
||
|
|
_insert_question(conn, "older", deadline_at="2001-01-01T00:00:00+00:00")
|
||
|
|
overdue = overdue_open_questions(conn)
|
||
|
|
assert [q.question_id for q in overdue] == ["older", "newer"]
|
||
|
|
|
||
|
|
|
||
|
|
def test_overdue_row_fields_mapped(conn: sqlite3.Connection) -> None:
|
||
|
|
_insert_question(
|
||
|
|
conn, "q", thread_id="t-42", turn=3, transport="github", deadline_at=_PAST
|
||
|
|
)
|
||
|
|
(q,) = overdue_open_questions(conn)
|
||
|
|
assert q.question_id == "q"
|
||
|
|
assert q.thread_id == "t-42"
|
||
|
|
assert q.turn == 3
|
||
|
|
assert q.transport == "github"
|
||
|
|
assert q.deadline_at == _PAST
|
||
|
|
|
||
|
|
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
# run_deadline_timer — PARK policy (default)
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
|
||
|
|
|
||
|
|
def test_park_policy_expires_and_alarms(conn: sqlite3.Connection) -> None:
|
||
|
|
_insert_question(conn, "q1")
|
||
|
|
on_park = _Collector()
|
||
|
|
|
||
|
|
report = run_deadline_timer(conn, on_park=on_park)
|
||
|
|
|
||
|
|
assert _status(conn, "q1") == "expired"
|
||
|
|
assert [q.question_id for q in on_park.calls] == ["q1"]
|
||
|
|
assert report.examined == 1
|
||
|
|
assert report.parked == 1
|
||
|
|
assert report.expired == 1
|
||
|
|
assert report.lost_race == 0
|
||
|
|
assert report.outcomes[0].action is ExpiryAction.PARKED
|
||
|
|
assert report.outcomes[0].policy is DeadlinePolicy.PARK
|
||
|
|
|
||
|
|
|
||
|
|
def test_default_policy_is_park(conn: sqlite3.Connection) -> None:
|
||
|
|
_insert_question(conn, "q1")
|
||
|
|
on_park = _Collector()
|
||
|
|
# No policy_resolver supplied -> PARK for every overdue question.
|
||
|
|
run_deadline_timer(conn, on_park=on_park)
|
||
|
|
assert len(on_park.calls) == 1
|
||
|
|
|
||
|
|
|
||
|
|
def test_multiple_overdue_all_parked(conn: sqlite3.Connection) -> None:
|
||
|
|
for i in range(5):
|
||
|
|
_insert_question(conn, f"q{i}")
|
||
|
|
on_park = _Collector()
|
||
|
|
|
||
|
|
report = run_deadline_timer(conn, on_park=on_park)
|
||
|
|
|
||
|
|
assert report.examined == 5
|
||
|
|
assert report.parked == 5
|
||
|
|
assert {c.question_id for c in on_park.calls} == {f"q{i}" for i in range(5)}
|
||
|
|
for i in range(5):
|
||
|
|
assert _status(conn, f"q{i}") == "expired"
|
||
|
|
|
||
|
|
|
||
|
|
def test_future_and_null_deadlines_untouched(conn: sqlite3.Connection) -> None:
|
||
|
|
_insert_question(conn, "future", deadline_at=_FUTURE)
|
||
|
|
_insert_question(conn, "none", deadline_at=None)
|
||
|
|
on_park = _Collector()
|
||
|
|
|
||
|
|
report = run_deadline_timer(conn, on_park=on_park)
|
||
|
|
|
||
|
|
assert report.examined == 0
|
||
|
|
assert on_park.calls == []
|
||
|
|
assert _status(conn, "future") == "open"
|
||
|
|
assert _status(conn, "none") == "open"
|
||
|
|
|
||
|
|
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
# run_deadline_timer — DEFAULT_ANSWER policy
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
|
||
|
|
|
||
|
|
def test_default_answer_policy_resumes(conn: sqlite3.Connection) -> None:
|
||
|
|
_insert_question(conn, "q1")
|
||
|
|
on_park = _Collector()
|
||
|
|
resume = _Collector()
|
||
|
|
|
||
|
|
report = run_deadline_timer(
|
||
|
|
conn,
|
||
|
|
on_park=on_park,
|
||
|
|
resume_with_default=resume,
|
||
|
|
policy_resolver=lambda _q: DeadlinePolicy.DEFAULT_ANSWER,
|
||
|
|
)
|
||
|
|
|
||
|
|
assert _status(conn, "q1") == "expired"
|
||
|
|
assert on_park.calls == []
|
||
|
|
assert [q.question_id for q in resume.calls] == ["q1"]
|
||
|
|
assert report.defaulted == 1
|
||
|
|
assert report.parked == 0
|
||
|
|
assert report.outcomes[0].action is ExpiryAction.DEFAULTED
|
||
|
|
|
||
|
|
|
||
|
|
def test_default_answer_without_callback_raises(conn: sqlite3.Connection) -> None:
|
||
|
|
_insert_question(conn, "q1")
|
||
|
|
on_park = _Collector()
|
||
|
|
|
||
|
|
with pytest.raises(ValueError, match="DEFAULT_ANSWER"):
|
||
|
|
run_deadline_timer(
|
||
|
|
conn,
|
||
|
|
on_park=on_park,
|
||
|
|
policy_resolver=lambda _q: DeadlinePolicy.DEFAULT_ANSWER,
|
||
|
|
)
|
||
|
|
# The row was already flipped to expired before the config error surfaced.
|
||
|
|
assert _status(conn, "q1") == "expired"
|
||
|
|
|
||
|
|
|
||
|
|
def test_mixed_policies_routed_per_question(conn: sqlite3.Connection) -> None:
|
||
|
|
_insert_question(conn, "park-me")
|
||
|
|
_insert_question(conn, "default-me")
|
||
|
|
on_park = _Collector()
|
||
|
|
resume = _Collector()
|
||
|
|
|
||
|
|
def resolver(q: OverdueQuestion) -> DeadlinePolicy:
|
||
|
|
return (
|
||
|
|
DeadlinePolicy.DEFAULT_ANSWER
|
||
|
|
if q.question_id == "default-me"
|
||
|
|
else DeadlinePolicy.PARK
|
||
|
|
)
|
||
|
|
|
||
|
|
report = run_deadline_timer(
|
||
|
|
conn,
|
||
|
|
on_park=on_park,
|
||
|
|
resume_with_default=resume,
|
||
|
|
policy_resolver=resolver,
|
||
|
|
)
|
||
|
|
|
||
|
|
assert [c.question_id for c in on_park.calls] == ["park-me"]
|
||
|
|
assert [c.question_id for c in resume.calls] == ["default-me"]
|
||
|
|
assert report.parked == 1
|
||
|
|
assert report.defaulted == 1
|
||
|
|
|
||
|
|
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
# Deterministic answer-vs-timeout race (§3.3.1)
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
|
||
|
|
|
||
|
|
def test_answered_first_loses_race_no_policy(conn: sqlite3.Connection) -> None:
|
||
|
|
"""A question answered before the timer runs must NOT be parked/defaulted."""
|
||
|
|
_insert_question(conn, "q1")
|
||
|
|
# Responder wins first.
|
||
|
|
assert answer_question(
|
||
|
|
conn, question_id="q1", answer_json="{}", answered_via="slack"
|
||
|
|
)
|
||
|
|
on_park = _Collector()
|
||
|
|
resume = _Collector()
|
||
|
|
|
||
|
|
report = run_deadline_timer(conn, on_park=on_park, resume_with_default=resume)
|
||
|
|
|
||
|
|
# The row is no longer ``open`` so it is not even returned by the overdue
|
||
|
|
# query — examined is zero, no policy applied.
|
||
|
|
assert report.examined == 0
|
||
|
|
assert on_park.calls == []
|
||
|
|
assert resume.calls == []
|
||
|
|
assert _status(conn, "q1") == "answered"
|
||
|
|
|
||
|
|
|
||
|
|
def test_lost_race_when_answered_between_select_and_flip(
|
||
|
|
conn: sqlite3.Connection, monkeypatch: pytest.MonkeyPatch
|
||
|
|
) -> None:
|
||
|
|
"""If a responder answers a row after it was selected as overdue but before
|
||
|
|
the timer flips it, the timer's compare-and-set returns False -> LOST_RACE,
|
||
|
|
and NO no-answer policy is applied (the question was actually answered)."""
|
||
|
|
import agent_team.deadline_timer as dt
|
||
|
|
|
||
|
|
_insert_question(conn, "q1")
|
||
|
|
on_park = _Collector()
|
||
|
|
|
||
|
|
real_expire = dt.expire_question
|
||
|
|
|
||
|
|
def racing_expire(c: sqlite3.Connection, *, question_id: str) -> bool:
|
||
|
|
# Simulate the responder winning the compare-and-set in the window
|
||
|
|
# between overdue selection and this flip. Uses the SAME connection so
|
||
|
|
# the ordering is fully deterministic (no thread scheduling needed).
|
||
|
|
answer_question(
|
||
|
|
c, question_id=question_id, answer_json="{}", answered_via="slack"
|
||
|
|
)
|
||
|
|
return real_expire(c, question_id=question_id)
|
||
|
|
|
||
|
|
monkeypatch.setattr(dt, "expire_question", racing_expire)
|
||
|
|
|
||
|
|
report = run_deadline_timer(conn, on_park=on_park)
|
||
|
|
|
||
|
|
# The flip lost: no park, row is ``answered`` not ``expired``.
|
||
|
|
assert on_park.calls == []
|
||
|
|
assert report.examined == 1
|
||
|
|
assert report.lost_race == 1
|
||
|
|
assert report.parked == 0
|
||
|
|
assert report.expired == 0
|
||
|
|
assert report.outcomes[0].action is ExpiryAction.LOST_RACE
|
||
|
|
assert _status(conn, "q1") == "answered"
|
||
|
|
|
||
|
|
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
# Restart idempotency
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
|
||
|
|
|
||
|
|
def test_rerun_is_idempotent(conn: sqlite3.Connection) -> None:
|
||
|
|
"""A second pass after a 'crash' processes no already-expired rows."""
|
||
|
|
_insert_question(conn, "q1")
|
||
|
|
on_park = _Collector()
|
||
|
|
|
||
|
|
first = run_deadline_timer(conn, on_park=on_park)
|
||
|
|
assert first.parked == 1
|
||
|
|
|
||
|
|
second = run_deadline_timer(conn, on_park=on_park)
|
||
|
|
# Already expired -> no longer ``open`` -> not selected -> no double park.
|
||
|
|
assert second.examined == 0
|
||
|
|
assert len(on_park.calls) == 1
|
||
|
|
assert _status(conn, "q1") == "expired"
|
||
|
|
|
||
|
|
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
# Side-effect isolation
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
|
||
|
|
|
||
|
|
def test_side_effect_error_isolated_row_still_expired(
|
||
|
|
conn: sqlite3.Connection,
|
||
|
|
) -> None:
|
||
|
|
_insert_question(conn, "boom")
|
||
|
|
_insert_question(conn, "ok")
|
||
|
|
|
||
|
|
def flaky_park(q: OverdueQuestion) -> None:
|
||
|
|
if q.question_id == "boom":
|
||
|
|
raise RuntimeError("alarm transport down")
|
||
|
|
|
||
|
|
report = run_deadline_timer(conn, on_park=flaky_park)
|
||
|
|
|
||
|
|
# Both rows are durably expired (flip commits before the side effect).
|
||
|
|
assert _status(conn, "boom") == "expired"
|
||
|
|
assert _status(conn, "ok") == "expired"
|
||
|
|
assert report.errored == 1
|
||
|
|
assert report.parked == 1
|
||
|
|
errored = next(o for o in report.outcomes if o.action is ExpiryAction.ERRORED)
|
||
|
|
assert errored.question_id == "boom"
|
||
|
|
assert "RuntimeError" in (errored.error or "")
|
||
|
|
assert errored.policy is DeadlinePolicy.PARK
|
||
|
|
|
||
|
|
|
||
|
|
def test_error_in_one_row_does_not_abort_batch(conn: sqlite3.Connection) -> None:
|
||
|
|
for i in range(4):
|
||
|
|
_insert_question(conn, f"q{i}")
|
||
|
|
|
||
|
|
def park(q: OverdueQuestion) -> None:
|
||
|
|
if q.question_id == "q1":
|
||
|
|
raise ValueError("nope")
|
||
|
|
|
||
|
|
report = run_deadline_timer(conn, on_park=park)
|
||
|
|
|
||
|
|
assert report.examined == 4
|
||
|
|
assert report.errored == 1
|
||
|
|
assert report.parked == 3
|
||
|
|
for i in range(4):
|
||
|
|
assert _status(conn, f"q{i}") == "expired"
|
||
|
|
|
||
|
|
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
# TimerLoopReport counters
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
|
||
|
|
|
||
|
|
def test_empty_report_counters() -> None:
|
||
|
|
report = TimerLoopReport()
|
||
|
|
assert report.examined == 0
|
||
|
|
assert report.parked == 0
|
||
|
|
assert report.defaulted == 0
|
||
|
|
assert report.lost_race == 0
|
||
|
|
assert report.errored == 0
|
||
|
|
assert report.expired == 0
|
||
|
|
|
||
|
|
|
||
|
|
def test_expired_equals_examined_minus_lost_race(
|
||
|
|
conn: sqlite3.Connection, monkeypatch: pytest.MonkeyPatch
|
||
|
|
) -> None:
|
||
|
|
import agent_team.deadline_timer as dt
|
||
|
|
|
||
|
|
_insert_question(conn, "parkable")
|
||
|
|
_insert_question(conn, "raced")
|
||
|
|
on_park = _Collector()
|
||
|
|
|
||
|
|
real_expire = dt.expire_question
|
||
|
|
|
||
|
|
def racing_expire(c: sqlite3.Connection, *, question_id: str) -> bool:
|
||
|
|
if question_id == "raced":
|
||
|
|
answer_question(
|
||
|
|
c, question_id="raced", answer_json="{}", answered_via="slack"
|
||
|
|
)
|
||
|
|
return real_expire(c, question_id=question_id)
|
||
|
|
|
||
|
|
monkeypatch.setattr(dt, "expire_question", racing_expire)
|
||
|
|
|
||
|
|
report = run_deadline_timer(conn, on_park=on_park)
|
||
|
|
|
||
|
|
assert report.examined == 2
|
||
|
|
assert report.lost_race == 1
|
||
|
|
assert report.expired == report.examined - report.lost_race == 1
|
||
|
|
|
||
|
|
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
# Concurrency: responder thread vs timer thread on the same question
|
||
|
|
# ---------------------------------------------------------------------------
|
||
|
|
|
||
|
|
|
||
|
|
def test_concurrent_timer_and_responder_single_winner(tmp_path: Path) -> None:
|
||
|
|
"""A timer pass and a responder race the same open question: the
|
||
|
|
BEGIN-IMMEDIATE compare-and-set guarantees exactly one of 'expire'/'answer'
|
||
|
|
wins, and the timer parks IFF it actually flipped the row to ``expired``.
|
||
|
|
|
||
|
|
Repeated across many rows so the threads genuinely interleave (the barrier
|
||
|
|
aligns each pair at the start), catching any non-determinism in the race.
|
||
|
|
"""
|
||
|
|
db = tmp_path / "race.sqlite"
|
||
|
|
init_db(db)
|
||
|
|
|
||
|
|
n = 40
|
||
|
|
seed = connect(db)
|
||
|
|
try:
|
||
|
|
for i in range(n):
|
||
|
|
_insert_question(seed, f"race{i}", deadline_at=_PAST)
|
||
|
|
finally:
|
||
|
|
seed.close()
|
||
|
|
|
||
|
|
park_calls: list[str] = []
|
||
|
|
park_lock = threading.Lock()
|
||
|
|
|
||
|
|
def run_one(qid: str) -> None:
|
||
|
|
barrier = threading.Barrier(2)
|
||
|
|
|
||
|
|
def timer_worker() -> None:
|
||
|
|
c = connect(db)
|
||
|
|
try:
|
||
|
|
|
||
|
|
def on_park(q: OverdueQuestion) -> None:
|
||
|
|
with park_lock:
|
||
|
|
park_calls.append(q.question_id)
|
||
|
|
|
||
|
|
barrier.wait()
|
||
|
|
run_deadline_timer(c, on_park=on_park)
|
||
|
|
finally:
|
||
|
|
c.close()
|
||
|
|
|
||
|
|
def responder_worker() -> None:
|
||
|
|
c = connect(db)
|
||
|
|
try:
|
||
|
|
barrier.wait()
|
||
|
|
answer_question(
|
||
|
|
c, question_id=qid, answer_json="{}", answered_via="slack"
|
||
|
|
)
|
||
|
|
finally:
|
||
|
|
c.close()
|
||
|
|
|
||
|
|
t1 = threading.Thread(target=timer_worker)
|
||
|
|
t2 = threading.Thread(target=responder_worker)
|
||
|
|
t1.start()
|
||
|
|
t2.start()
|
||
|
|
t1.join()
|
||
|
|
t2.join()
|
||
|
|
|
||
|
|
for i in range(n):
|
||
|
|
run_one(f"race{i}")
|
||
|
|
|
||
|
|
check = connect(db)
|
||
|
|
try:
|
||
|
|
rows = {
|
||
|
|
r["question_id"]: r["status"]
|
||
|
|
for r in check.execute(
|
||
|
|
"SELECT question_id, status FROM pending_questions"
|
||
|
|
).fetchall()
|
||
|
|
}
|
||
|
|
finally:
|
||
|
|
check.close()
|
||
|
|
|
||
|
|
# Every row ends in exactly one terminal state.
|
||
|
|
for i in range(n):
|
||
|
|
assert rows[f"race{i}"] in {"expired", "answered"}
|
||
|
|
# Every parked question must be one the timer actually expired.
|
||
|
|
for qid in park_calls:
|
||
|
|
assert rows[qid] == "expired"
|