test: cover outcome writer commit races, heartbeat and task restart

This commit is contained in:
Ashen
2026-10-01 11:21:53 +05:30
parent 67724aec6d
commit ef429769eb
+140
View File
@@ -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)