From ef429769eb3a543013a4ef5d0aa922aaf9586f55 Mon Sep 17 00:00:00 2001 From: Ashen <310210685+ashen0x@users.noreply.github.com> Date: Sun, 27 Sep 2026 20:05:38 +0530 Subject: [PATCH] test: cover outcome writer commit races, heartbeat and task restart --- tests/unit/test_terminal_outcomes.py | 140 +++++++++++++++++++++++++++ 1 file changed, 140 insertions(+) diff --git a/tests/unit/test_terminal_outcomes.py b/tests/unit/test_terminal_outcomes.py index 9fc878ee..b825a304 100644 --- a/tests/unit/test_terminal_outcomes.py +++ b/tests/unit/test_terminal_outcomes.py @@ -9,6 +9,7 @@ from pathlib import Path from unittest.mock import MagicMock import pytest +from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlmodel import SQLModel, col, select from sqlmodel.ext.asyncio.session import AsyncSession @@ -795,3 +796,142 @@ async def test_restart_alongside_live_writer_preserves_continuous_coverage( finally: assert await writers[1].stop(timeout=1) assert await writers[2].stop(timeout=1) + + +def _commit_fails_once(sessions: SessionFactory) -> SessionFactory: + """A competing writer wins the first commit, as it can between processes.""" + failed = False + + @asynccontextmanager + async def factory() -> AsyncGenerator[AsyncSession, None]: + async with sessions() as session: + commit = session.commit + + async def commit_once() -> None: + nonlocal failed + if not failed: + failed = True + raise IntegrityError("COMMIT", {}, Exception("competing writer")) + await commit() + + session.commit = commit_once # type: ignore[method-assign] + yield session + + return factory + + +@pytest.mark.parametrize("step", ["ensure", "rotate", "pending", "unattended"]) +async def test_epoch_steps_retry_after_losing_a_commit_race( + ledger: tuple[AsyncEngine, SessionFactory], step: str +) -> None: + _, sessions = ledger + day = date(2026, 9, 1) + clock = MutableClock(_timestamp(day)) + if step != "ensure": + await TerminalOutcomeWriter( + session_factory=sessions, clock=clock + )._ensure_epoch(clock.value) + clock.value = _timestamp(day + timedelta(days=3)) + if step == "pending": + async with sessions() as session: + session.add( + TerminalOutcomeWriterRun( + run_id="lost-run", + status="lost", + started_at_ms=_timestamp(day), + heartbeat_at_ms=_timestamp(day), + closed_at_ms=_timestamp(day), + loss_day=day + timedelta(days=2), + ) + ) + await session.commit() + writer = TerminalOutcomeWriter( + session_factory=_commit_fails_once(sessions), clock=clock + ) + + if step == "ensure": + await writer._ensure_epoch(clock.value) + elif step == "rotate": + await writer._rotate_epoch(day + timedelta(days=2)) + elif step == "pending": + await writer._recover_pending_runs() + else: + await writer._recover_unattended_coverage(clock.value) + + async with sessions() as session: + epochs = ( + await session.exec( + select(TerminalOutcomeEpoch).order_by(col(TerminalOutcomeEpoch.epoch)) + ) + ).all() + run = await session.get(TerminalOutcomeWriterRun, "lost-run") + coverage = [ + (epoch.coverage_start_day, epoch.coverage_end_day, epoch.current_slot) + for epoch in epochs + ] + reopened = (day + timedelta(days=4), None, 1) + expected = { + "ensure": [reopened], + "rotate": [(day + timedelta(days=1), day + timedelta(days=1), None), reopened], + "pending": [(day + timedelta(days=1), day + timedelta(days=1), None), reopened], + "unattended": [(day + timedelta(days=1), day, None), reopened], + } + assert coverage == expected[step] + assert writer.epoch == len(epochs) - 1 + if step == "pending": + assert run is not None and run.status == "recovered" + + +async def test_heartbeat_renews_the_lease_and_checkpoint( + ledger: tuple[AsyncEngine, SessionFactory], +) -> None: + _, sessions = ledger + clock = MutableClock(_timestamp(date(2026, 9, 1))) + writer = TerminalOutcomeWriter(session_factory=sessions, clock=clock) + assert await writer.start() + clock.value += 60_000 + + await writer._heartbeat() + + async with sessions() as session: + run = (await session.exec(select(TerminalOutcomeWriterRun))).one() + assert run.heartbeat_at_ms == clock.value + assert run.flushed_through_ms == clock.value + assert await writer.stop(timeout=1) + + +async def test_writer_task_that_dies_declares_a_loss_and_restarts( + ledger: tuple[AsyncEngine, SessionFactory], + monkeypatch: pytest.MonkeyPatch, +) -> None: + _, sessions = ledger + writer = TerminalOutcomeWriter( + session_factory=sessions, heartbeat_seconds=0.01, lease_timeout_seconds=1 + ) + + class WriterDied(BaseException): + pass + + heartbeat = writer._heartbeat + died = False + + async def die_once() -> None: + nonlocal died + if not died: + died = True + raise WriterDied + await heartbeat() + + monkeypatch.setattr(writer, "_heartbeat", die_once) + assert await writer.start() + first = writer._task + assert first is not None + for _ in range(300): + if first.done() and writer.epoch == 1 and not writer.loss_pending: + break + await asyncio.sleep(0.01) + + # The loss rotates coverage, and a fresh task keeps collecting. + assert writer.epoch == 1 and not writer.loss_pending + assert writer.running and writer._task is not first + await writer.stop(timeout=1)