mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-06 04:38:22 +00:00
test: cover outcome writer commit races, heartbeat and task restart
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user