From 56c0bf280f3f58693f989741199869241a7d968d Mon Sep 17 00:00:00 2001 From: Anto Subash Date: Fri, 4 Sep 2026 23:34:14 +0200 Subject: [PATCH] fix(background_tasks): make a concurrent "Retry failed" sweep safe MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `RetryCoordinator.retry_failed` read the eligible rows, published every one of them to the broker, and only then inserted the child rows that make them ineligible. Nothing held a lock across those three steps, so two operators pressing the button at once — or one double-click — both read the same batch and both published it. The guard meant to prevent that, `bulk_retry_conditions` excluding parents that already have a child, cannot fire until the child exists, and the child did not exist until after the publish. The result was every task in the sweep running twice, with two retry chains hanging off the same parent. Three changes, in the order they matter: * The claim takes `FOR UPDATE SKIP LOCKED`. A second sweep skips what the first is holding instead of reading it again, and the rows stay held until the transaction that wrote the children commits — which is exactly the moment the child guard takes over. `skip_locked` rather than plain `FOR UPDATE` so the second operator gets an answer instead of blocking behind a 500-row batch. * The child rows are written before anything is published. Losing a race is now an orphaned `pending` row that was never sent — visible, retryable, and cheap — rather than a task that already ran twice. This means reserving the broker id rather than reading it back from `send_task`, so the row still names the task that carries it. * `TaskRetried` is published from the request's post-commit hook (#268) rather than inline. Subscribers were being told about a transaction that could still roll back, and for a 500-row sweep were told nothing at all until the batch finished and then everything at once. `RetryCoordinator` now takes a `RequestSession` to say that dependency out loud. Also in the same screen: * `GET /api/background-tasks/executions` now takes the task-name search as `q`, the spelling the page and `retry-failed` already use, so a URL copied between them keeps its filter instead of silently listing everything. `task_name` stays accepted and deprecated. * Every strip tile states the period it covers ("All time" / "Last 24h"). The row mixes units by design — a success count is only interesting as a rate — but unlabelled it made "Succeeded 24h" reading lower than the table's total for the same filter look like a bug. On tests: `with_for_update` compiles to nothing on SQLite, so a concurrency test here would pass with the fix removed. The lock is asserted on the SQL `claim_query` compiles to against the Postgres dialect; the write-before- publish and announce-after-commit orderings are asserted on what the engine and the broker actually saw, which holds on any backend. Closes #298 --- .../background_tasks/background_tasks/deps.py | 4 +- .../background_tasks/endpoints/api_admin.py | 13 +- .../background_tasks/locales/en.json | 6 +- .../pages/components/StatusStrip.tsx | 26 ++- .../background_tasks/retry_service.py | 144 +++++++++---- .../background_tasks/service.py | 4 +- .../background_tasks/tests/test_admin_api.py | 29 +++ .../background_tasks/tests/test_bg_service.py | 31 +-- .../tests/test_retry_sweep_ordering.py | 200 ++++++++++++++++++ packages/i18n/src/generated-resources.ts | 4 +- packages/i18n/src/keys.generated.ts | 4 +- 11 files changed, 397 insertions(+), 68 deletions(-) create mode 100644 modules/background_tasks/tests/test_retry_sweep_ordering.py diff --git a/modules/background_tasks/background_tasks/deps.py b/modules/background_tasks/background_tasks/deps.py index a39cc1bc..6765920a 100644 --- a/modules/background_tasks/background_tasks/deps.py +++ b/modules/background_tasks/background_tasks/deps.py @@ -4,15 +4,15 @@ from fastapi import Depends, Request from simple_module_core.events import EventBus +from simple_module_db import RequestSession from simple_module_db.deps import get_db -from sqlalchemy.ext.asyncio import AsyncSession from background_tasks.service import BackgroundTaskService async def get_background_task_service( request: Request, - db: AsyncSession = Depends(get_db), + db: RequestSession = Depends(get_db), ) -> BackgroundTaskService: services = request.app.state.background_tasks bus: EventBus = request.app.state.sm.event_bus diff --git a/modules/background_tasks/background_tasks/endpoints/api_admin.py b/modules/background_tasks/background_tasks/endpoints/api_admin.py index b6559e50..2997e202 100644 --- a/modules/background_tasks/background_tasks/endpoints/api_admin.py +++ b/modules/background_tasks/background_tasks/endpoints/api_admin.py @@ -34,14 +34,23 @@ @router.get("/executions", response_model=TaskExecutionListResponse) async def list_executions( status: TaskStatus | None = Query(default=None), - task_name: str | None = Query(default=None), + q: str | None = Query(default=None), + task_name: str | None = Query(default=None, deprecated=True), queue: str | None = Query(default=None), page: int = Query(default=1, ge=1), per_page: int = Query(default=20, ge=1, le=200), service: BackgroundTaskService = Depends(get_background_task_service), ) -> TaskExecutionListResponse: + """List executions, filtered exactly as the page filters them. + + The task-name search is ``q`` — the same name the page and the sweep use, + so an operator can move a URL between the three and keep their filter. This + endpoint originally spelled it ``task_name``, which meant a link copied off + the screen silently listed *everything*; that spelling is still accepted so + existing API callers keep working, and ``q`` wins when both are given. + """ return await service.list( - status=status, task_name=task_name, queue=queue, page=page, per_page=per_page + status=status, task_name=q or task_name, queue=queue, page=page, per_page=per_page ) diff --git a/modules/background_tasks/background_tasks/locales/en.json b/modules/background_tasks/background_tasks/locales/en.json index 8f416a98..2a95bc52 100644 --- a/modules/background_tasks/background_tasks/locales/en.json +++ b/modules/background_tasks/background_tasks/locales/en.json @@ -12,9 +12,11 @@ "strip": { "queued": "Queued", "running": "Running", - "succeeded_24h": "Succeeded 24h", + "succeeded": "Succeeded", "failed": "Failed", - "stuck": "Stuck" + "stuck": "Stuck", + "window_all": "All time", + "window_24h": "Last 24h" }, "filters": { "status_label": "Filter by status", diff --git a/modules/background_tasks/background_tasks/pages/components/StatusStrip.tsx b/modules/background_tasks/background_tasks/pages/components/StatusStrip.tsx index c267bf96..e3d711bb 100644 --- a/modules/background_tasks/background_tasks/pages/components/StatusStrip.tsx +++ b/modules/background_tasks/background_tasks/pages/components/StatusStrip.tsx @@ -21,6 +21,8 @@ interface Tile { /** Status this tile narrows the table to. */ filter: TaskStatus; labelKey: StripLabelKey; + /** The period the figure covers, spelled out under it. */ + windowKey: StripLabelKey; count: (counts: StatusCounts) => number; tone: Tone; } @@ -29,21 +31,30 @@ interface Tile { * The five numbers that describe a queue, left to right in the order work * moves through it: waiting, running, done, broken, wedged. * - * "Succeeded 24h" is the only windowed figure, and deliberately so — the other - * four are states something is in *right now*, while success is only - * interesting as a rate. An all-time success total is a number that never goes - * down and therefore never says anything. + * Succeeded is the only windowed figure, and deliberately so — the other four + * are states something is in *right now*, while success is only interesting as + * a rate. An all-time success total is a number that never goes down and + * therefore never says anything. + * + * That makes the row mixed units, which is why every tile carries its period + * rather than only the odd one out: side by side and unlabelled, a 24-hour + * success count reading lower than the table's total for the same filter looks + * like a bug, and nothing on screen said otherwise. */ +const ALL_TIME = keys.background_tasks.strip.window_all; + const TILES: Tile[] = [ { filter: TASK_STATUS.PENDING, labelKey: keys.background_tasks.strip.queued, + windowKey: ALL_TIME, count: (c) => c[TASK_STATUS.PENDING] ?? 0, tone: 'default', }, { filter: TASK_STATUS.RUNNING, labelKey: keys.background_tasks.strip.running, + windowKey: ALL_TIME, count: (c) => c[TASK_STATUS.RUNNING] ?? 0, tone: 'default', }, @@ -52,19 +63,22 @@ const TILES: Tile[] = [ // windowed: the tile is how an operator reaches the successes at all, and // the segmented control below has no option for them. filter: TASK_STATUS.SUCCESS, - labelKey: keys.background_tasks.strip.succeeded_24h, + labelKey: keys.background_tasks.strip.succeeded, + windowKey: keys.background_tasks.strip.window_24h, count: (c) => c.success_24h ?? 0, tone: 'default', }, { filter: TASK_STATUS.FAILED, labelKey: keys.background_tasks.strip.failed, + windowKey: ALL_TIME, count: (c) => c[TASK_STATUS.FAILED] ?? 0, tone: 'destructive', }, { filter: TASK_STATUS.STUCK, labelKey: keys.background_tasks.strip.stuck, + windowKey: ALL_TIME, count: (c) => c[TASK_STATUS.STUCK] ?? 0, tone: 'warning', }, @@ -101,6 +115,8 @@ export function StatusStrip({ counts, active, onSelect }: Props) { diff --git a/modules/background_tasks/background_tasks/retry_service.py b/modules/background_tasks/background_tasks/retry_service.py index 3f4f42ed..0f029d2a 100644 --- a/modules/background_tasks/background_tasks/retry_service.py +++ b/modules/background_tasks/background_tasks/retry_service.py @@ -9,31 +9,67 @@ from __future__ import annotations import asyncio -from typing import TYPE_CHECKING +import uuid +from collections.abc import Sequence +from typing import TYPE_CHECKING, Any from simple_module_core.events import EventBus -from sqlalchemy import func, select -from sqlalchemy.ext.asyncio import AsyncSession +from simple_module_db import RequestSession +from sqlalchemy import Select, func, select from background_tasks.constants import RETRY_ALL_BATCH, TaskStatus from background_tasks.contracts.events import TaskRetried from background_tasks.contracts.schemas import RetryFailedResult -from background_tasks.filters import bulk_retry_conditions +from background_tasks.filters import Conditions, bulk_retry_conditions from background_tasks.models import TaskExecution if TYPE_CHECKING: from celery import Celery +def claim_query(conditions: Conditions, batch: int) -> Select[tuple[TaskExecution]]: + """The rows one sweep takes, held for the length of its transaction. + + ``FOR UPDATE SKIP LOCKED`` is what makes two simultaneous sweeps disjoint. + The guard that is *supposed* to keep a row out of a second sweep — + :func:`~background_tasks.filters.bulk_retry_conditions` excluding parents + that already have a child — only starts working once the child row exists, + and a concurrent sweep reads the same batch long before that. Locking the + parents makes the second reader skip them instead of publishing the same + tasks a second time, and the rows stay held until the transaction that + wrote the children commits, which is the moment the child guard takes over. + + ``skip_locked`` rather than plain ``FOR UPDATE``: a second operator should + get "nothing left to sweep" immediately, not block behind a 500-row batch. + + SQLite has no row locks and silently emits no clause; it also takes a + database-wide write lock, so the two sweeps this protects against cannot + overlap there in the first place. + """ + return ( + select(TaskExecution) + .where(*conditions) + # Oldest first: the queue an operator is unwedging should come back + # out in the order it went in. + .order_by(TaskExecution.queued_at.asc().nulls_last()) + .limit(batch) + .with_for_update(skip_locked=True) + ) + + class RetryCoordinator: """Sends tasks again and records each fresh attempt. The original row is immutable: a retry inserts a *new* ``TaskExecution`` carrying ``retried_from_id``, so the detail page can show the chain and nothing rewrites history. + + Takes a :class:`~simple_module_db.RequestSession` rather than a bare + ``AsyncSession`` because it announces retries through the session's + post-commit hook — see :meth:`_announce_after_commit`. """ - def __init__(self, db: AsyncSession, celery: Celery, event_bus: EventBus) -> None: + def __init__(self, db: RequestSession, celery: Celery, event_bus: EventBus) -> None: self.db = db self.celery = celery self.event_bus = event_bus @@ -50,7 +86,7 @@ async def retry_one(self, row: TaskExecution) -> TaskExecution: self.db.add(new_row) await self.db.flush() await self.db.refresh(new_row) - await self._announce(row, new_row) + self._announce_after_commit([(row, new_row)]) return new_row async def retry_failed( @@ -73,6 +109,13 @@ async def retry_failed( operator knows to press again. The new rows are ``pending`` and their parents now have a child, so neither this pass nor the next picks up its own work. + + Two operators pressing the button at once used to run every task twice: + both read the same batch, and the child rows that would have excluded + it were only written afterwards. :func:`claim_query` locks the batch, + and the children are written *before* the tasks are published — an + orphaned ``pending`` row that was never sent is a far cheaper failure + than a task that ran twice. """ # Read at call time, not bound as a default argument, so the cap has # exactly one source of truth that tests and operators can move. @@ -91,49 +134,52 @@ async def retry_failed( if not eligible: return RetryFailedResult(queued=0, remaining=0) - rows = list( - ( - await self.db.execute( - select(TaskExecution) - .where(*conditions) - # Oldest first: the queue an operator is unwedging should - # come back out in the order it went in. - .order_by(TaskExecution.queued_at.asc().nulls_last()) - .limit(batch) - ) - ).scalars() - ) - - # One thread for the whole batch. ``send_task`` is a blocking publish, - # and a few hundred of them inline would stall every other request for - # the length of the sweep. - celery_ids = await asyncio.to_thread(lambda: [self._publish(row) for row in rows]) - - new_rows = [ - self._new_attempt(row, celery_id) - for row, celery_id in zip(rows, celery_ids, strict=True) - ] + rows = list((await self.db.execute(claim_query(conditions, batch))).scalars()) + if not rows: + # Everything eligible is held by a sweep that got here first. Say + # so rather than reporting an empty queue: the work is moving, just + # not by this press. + return RetryFailedResult(queued=0, remaining=eligible) + + # The broker id is reserved here rather than read back from ``send_task`` + # so the row recording the attempt can be written first and still name + # the task that will carry it. Held as plain strings, not read back off + # the ORM objects, because the publish below runs on another thread. + reserved = [(row, str(uuid.uuid4())) for row in rows] + new_rows = [self._new_attempt(row, task_id) for row, task_id in reserved] self.db.add_all(new_rows) # One flush for the batch rather than one per row: the round trips are # the sweep's real cost, and the rows have no reason to land singly. await self.db.flush() - for original, new_row in zip(rows, new_rows, strict=True): - await self._announce(original, new_row) + # One thread for the whole batch. ``send_task`` is a blocking publish, + # and a few hundred of them inline would stall every other request for + # the length of the sweep. + await asyncio.to_thread( + lambda: [self._publish(row, task_id=task_id) for row, task_id in reserved] + ) + + self._announce_after_commit(list(zip(rows, new_rows, strict=True))) return RetryFailedResult(queued=len(rows), remaining=max(0, eligible - len(rows))) - def _publish(self, row: TaskExecution) -> str: + def _publish(self, row: TaskExecution, task_id: str | None = None) -> str: """Send one task to the broker and return its Celery id. + *task_id* pre-assigns that id, for callers that wrote the row naming it + before sending. Omitted, the broker assigns one and the caller stamps + the row with what comes back. + Blocking — kombu publishes synchronously. Call it off the event loop when sending more than one. """ - return self.celery.send_task( - row.task_name, - args=list(row.args or []), - kwargs=dict(row.kwargs or {}), - queue=row.queue, - ).id + options: dict[str, Any] = { + "args": list(row.args or []), + "kwargs": dict(row.kwargs or {}), + "queue": row.queue, + } + if task_id is not None: + options["task_id"] = task_id + return self.celery.send_task(row.task_name, **options).id def _new_attempt(self, row: TaskExecution, celery_task_id: str) -> TaskExecution: """The row recording a fresh attempt at *row*. Not yet flushed.""" @@ -147,12 +193,28 @@ def _new_attempt(self, row: TaskExecution, celery_task_id: str) -> TaskExecution retried_from_id=row.id, ) - async def _announce(self, original: TaskExecution, new_row: TaskExecution) -> None: - """Same event whichever path enqueued it — a retry is a retry.""" - await self.event_bus.publish( + def _announce_after_commit( + self, retried: Sequence[tuple[TaskExecution, TaskExecution]] + ) -> None: + """Announce each retry once its row is durable. Same event either path. + + Publishing inline told subscribers about a transaction that could still + roll back — and for a 500-row sweep it told them nothing at all until + the batch was done, then everything at once. The events are *built* + here, while the rows are attached and their ids are known, and sent + from the request's post-commit hook. + """ + events = [ TaskRetried( original_id=original.id, new_id=new_row.id, task_name=original.task_name, ) - ) + for original, new_row in retried + ] + + async def publish() -> None: + for event in events: + await self.event_bus.publish(event) + + self.db.on_commit(publish) diff --git a/modules/background_tasks/background_tasks/service.py b/modules/background_tasks/background_tasks/service.py index f5a86931..9f60fcc7 100644 --- a/modules/background_tasks/background_tasks/service.py +++ b/modules/background_tasks/background_tasks/service.py @@ -9,8 +9,8 @@ from fastapi import HTTPException, status from simple_module_core.events import EventBus +from simple_module_db import RequestSession from sqlalchemy import Row, func, select -from sqlalchemy.ext.asyncio import AsyncSession from background_tasks.constants import RETRYABLE_STATUSES, TaskStatus from background_tasks.contracts.schemas import ( @@ -43,7 +43,7 @@ class BackgroundTaskService: def __init__( self, - db: AsyncSession, + db: RequestSession, celery: Celery, event_bus: EventBus, ) -> None: diff --git a/modules/background_tasks/tests/test_admin_api.py b/modules/background_tasks/tests/test_admin_api.py index 7d357029..c5baa648 100644 --- a/modules/background_tasks/tests/test_admin_api.py +++ b/modules/background_tasks/tests/test_admin_api.py @@ -74,6 +74,35 @@ async def test_search_matches_like_metacharacters_literally( assert resp.status_code == 200 assert [i["task_name"] for i in resp.json()["items"]] == ["cleanup_old"] + async def test_the_search_is_spelled_q_like_the_page_and_the_sweep( + self, app, authenticated_client: httpx.AsyncClient + ): + """One filter name across all three surfaces. + + The page and ``retry-failed`` both take the search as ``q``; this + endpoint took it as ``task_name``, so a URL copied off the screen + listed everything and quietly looked like a filter that matched. + """ + await _seed_failed(app, task_name="files.thumbnail") + await _seed_failed(app, task_name="users.invite") + + resp = await authenticated_client.get(f"{ADMIN_BASE}/executions", params={"q": "files."}) + assert resp.status_code == 200 + assert [i["task_name"] for i in resp.json()["items"]] == ["files.thumbnail"] + + async def test_the_old_task_name_spelling_still_filters( + self, app, authenticated_client: httpx.AsyncClient + ): + """Deprecated, not removed — an existing API caller keeps working.""" + await _seed_failed(app, task_name="files.thumbnail") + await _seed_failed(app, task_name="users.invite") + + resp = await authenticated_client.get( + f"{ADMIN_BASE}/executions", params={"task_name": "users."} + ) + assert resp.status_code == 200 + assert [i["task_name"] for i in resp.json()["items"]] == ["users.invite"] + async def test_page_past_the_end_clamps_instead_of_reporting_zero( self, app, authenticated_client: httpx.AsyncClient ): diff --git a/modules/background_tasks/tests/test_bg_service.py b/modules/background_tasks/tests/test_bg_service.py index 3a33a404..3ec6b69d 100644 --- a/modules/background_tasks/tests/test_bg_service.py +++ b/modules/background_tasks/tests/test_bg_service.py @@ -14,7 +14,7 @@ from background_tasks.service import BackgroundTaskService from fastapi import HTTPException from simple_module_core.events import EventBus -from sqlalchemy.ext.asyncio import AsyncSession +from simple_module_db import RequestSession, finalize_session def _make_row( @@ -51,14 +51,14 @@ def mock_celery() -> MagicMock: @pytest.fixture def service( - db_session: AsyncSession, event_bus: EventBus, mock_celery: MagicMock + db_session: RequestSession, event_bus: EventBus, mock_celery: MagicMock ) -> BackgroundTaskService: return BackgroundTaskService(db=db_session, celery=mock_celery, event_bus=event_bus) class TestList: async def test_returns_paginated_rows_newest_first( - self, db_session: AsyncSession, service: BackgroundTaskService + self, db_session: RequestSession, service: BackgroundTaskService ): now = datetime.now(UTC) older = _make_row(task_name="demo.a", queued_at=now - timedelta(minutes=5)) @@ -71,7 +71,7 @@ async def test_returns_paginated_rows_newest_first( assert [i.task_name for i in resp.items] == ["demo.b", "demo.a"] async def test_filters_by_status( - self, db_session: AsyncSession, service: BackgroundTaskService + self, db_session: RequestSession, service: BackgroundTaskService ): db_session.add_all( [ @@ -85,7 +85,7 @@ async def test_filters_by_status( assert [i.task_name for i in resp.items] == ["demo.bad"] async def test_filters_by_task_name_substring( - self, db_session: AsyncSession, service: BackgroundTaskService + self, db_session: RequestSession, service: BackgroundTaskService ): db_session.add_all( [ @@ -103,7 +103,7 @@ class TestStatusCounts: """Feeds the failed/stuck ops strip above the executions table.""" async def test_counts_are_grouped_by_status( - self, db_session: AsyncSession, service: BackgroundTaskService + self, db_session: RequestSession, service: BackgroundTaskService ): for status in (TaskStatus.FAILED, TaskStatus.FAILED, TaskStatus.SUCCESS): db_session.add(_make_row(status=status)) @@ -117,7 +117,7 @@ async def test_empty_table_yields_no_counts(self, service: BackgroundTaskService assert await service.status_counts() == {} async def test_search_narrows_the_counts( - self, db_session: AsyncSession, service: BackgroundTaskService + self, db_session: RequestSession, service: BackgroundTaskService ): """The strip must describe the same rows the table is paging through.""" db_session.add(_make_row(task_name="orders.send_receipt", status=TaskStatus.FAILED)) @@ -128,7 +128,7 @@ async def test_search_narrows_the_counts( assert counts == {TaskStatus.FAILED.value: 1} async def test_counts_keys_are_plain_strings( - self, db_session: AsyncSession, service: BackgroundTaskService + self, db_session: RequestSession, service: BackgroundTaskService ): """Postgres hands back the enum, SQLite a str — the page needs one shape.""" db_session.add(_make_row(status=TaskStatus.STUCK)) @@ -144,7 +144,7 @@ async def test_returns_none_for_missing_id(self, service: BackgroundTaskService) assert await service.get(uuid.uuid4()) is None async def test_returns_detail_for_existing_row( - self, db_session: AsyncSession, service: BackgroundTaskService + self, db_session: RequestSession, service: BackgroundTaskService ): row = _make_row(args=[1, "two"], kwargs={"three": 3}) db_session.add(row) @@ -160,7 +160,7 @@ async def test_returns_detail_for_existing_row( class TestRetry: async def test_retry_failed_task_enqueues_and_creates_row( self, - db_session: AsyncSession, + db_session: RequestSession, service: BackgroundTaskService, mock_celery: MagicMock, event_bus: EventBus, @@ -194,12 +194,19 @@ async def _on_retried(event: TaskRetried) -> None: assert detail.retried_from_id == original.id assert detail.status == TaskStatus.PENDING assert detail.celery_task_id == "new-celery-id-123" + # Nothing is announced while the retry is still a pending write — a + # subscriber acting on a transaction that rolls back acts on a retry + # that never happened. + assert received == [] + + await finalize_session(db_session) + assert len(received) == 1 assert received[0].original_id == original.id assert received[0].new_id == detail.id async def test_retry_stuck_task_is_allowed( - self, db_session: AsyncSession, service: BackgroundTaskService + self, db_session: RequestSession, service: BackgroundTaskService ): row = _make_row(status=TaskStatus.STUCK) db_session.add(row) @@ -210,7 +217,7 @@ async def test_retry_stuck_task_is_allowed( assert detail.retried_from_id == row.id async def test_retry_rejects_non_retryable_status( - self, db_session: AsyncSession, service: BackgroundTaskService + self, db_session: RequestSession, service: BackgroundTaskService ): row = _make_row(status=TaskStatus.SUCCESS) db_session.add(row) diff --git a/modules/background_tasks/tests/test_retry_sweep_ordering.py b/modules/background_tasks/tests/test_retry_sweep_ordering.py new file mode 100644 index 00000000..95276b52 --- /dev/null +++ b/modules/background_tasks/tests/test_retry_sweep_ordering.py @@ -0,0 +1,200 @@ +"""Bulk retry: what must not happen twice, and in what order it happens. + +Scope lives in ``test_retry_failed_bulk.py`` and the per-press guards in +``test_retry_failed_guards.py``. This file covers the three properties that +keep two simultaneous presses from running every task twice: + +* the claim locks the batch it takes, so a second sweep cannot read it; +* the rows recording the attempts are written *before* anything is published; +* the retry is announced only once those rows are durable. + +**What each test actually proves.** The suite runs on SQLite, where +``FOR UPDATE SKIP LOCKED`` is a no-op — a two-task race here would pass with +the locking removed and prove nothing. So the lock is asserted on the *SQL the +claim compiles to* against the Postgres dialect, which is where the clause has +to be right; the other two are ordering invariants observable on any backend, +and they are what makes a lost race cheap rather than catastrophic. +""" + +from __future__ import annotations + +from collections.abc import Iterator +from typing import Any +from unittest.mock import MagicMock + +import httpx +import pytest +from background_tasks.constants import RETRY_ALL_BATCH, TABLE_TASK_EXECUTION, TaskStatus +from background_tasks.contracts.events import TaskRetried +from background_tasks.filters import bulk_retry_conditions +from background_tasks.retry_service import claim_query +from sqlalchemy import event +from sqlalchemy.dialects import postgresql + +ADMIN_BASE = "/api/background_tasks/admin" +RETRY_FAILED = f"{ADMIN_BASE}/executions/retry-failed" + + +def _compiled_claim() -> str: + conditions = bulk_retry_conditions() + assert conditions is not None + return str(claim_query(conditions, RETRY_ALL_BATCH).compile(dialect=postgresql.dialect())) + + +class TestTheClaimLocksItsBatch: + """Asserted on the emitted SQL, not on behaviour. + + SQLite has no row locks and SQLAlchemy silently drops the clause there, so + running two sweeps against the test database would pass either way. The + compiled Postgres statement is the only thing in reach of this suite that + can tell the fix from its absence. + """ + + def test_the_claim_takes_a_row_lock(self) -> None: + assert "FOR UPDATE" in _compiled_claim() + + def test_a_second_sweep_skips_a_held_row_rather_than_waiting_for_it(self) -> None: + """``SKIP LOCKED``: the second operator gets an answer, not a queue.""" + assert "SKIP LOCKED" in _compiled_claim() + + def test_the_claim_is_still_the_oldest_rows_first(self) -> None: + """Locking must not have cost the ordering the sweep is defined by.""" + sql = _compiled_claim() + assert "ORDER BY" in sql + assert sql.index("ORDER BY") < sql.index("LIMIT") < sql.index("FOR UPDATE") + + +@pytest.fixture +def sweep_trace(app) -> Iterator[list[str]]: + """Every insert, commit and broker publish this sweep makes, in order. + + Recorded at the engine and at the Celery stub rather than inside the + coordinator, so the test asserts on what the database and the broker + actually saw and not on the shape of the code that talked to them. + """ + trace: list[str] = [] + engine = app.state.sm.db.engine.sync_engine + + def _record_statement(conn, cursor, statement, parameters, context, executemany) -> None: + if statement.lstrip().lower().startswith(f"insert into {TABLE_TASK_EXECUTION}"): + trace.append("insert") + + def _record_commit(conn) -> None: + trace.append("commit") + + def _send_task(name: str, **options: Any) -> MagicMock: + trace.append("publish") + result = MagicMock() + result.id = options.get("task_id", "mocked-celery-id") + return result + + celery = MagicMock(name="Celery") + celery.send_task.side_effect = _send_task + app.state.background_tasks.celery = celery + + event.listen(engine, "before_cursor_execute", _record_statement) + event.listen(engine, "commit", _record_commit) + try: + yield trace + finally: + event.remove(engine, "before_cursor_execute", _record_statement) + event.remove(engine, "commit", _record_commit) + + +class TestRowsBeforeBroker: + """The row recording an attempt exists before the attempt is sent. + + Provable on any backend. It is what makes losing the race survivable: an + orphaned ``pending`` row nobody sent is a row an operator can see and retry, + while a task published against no row has already run twice by the time + anyone notices. + """ + + async def test_the_retry_row_is_written_before_the_task_is_published( + self, sweep_trace, seed_execution, authenticated_client: httpx.AsyncClient + ) -> None: + await seed_execution(status=TaskStatus.FAILED) + sweep_trace.clear() # Drop the seed's own insert and commit. + + resp = await authenticated_client.post(RETRY_FAILED) + + assert resp.json() == {"queued": 1, "remaining": 0} + assert sweep_trace.index("insert") < sweep_trace.index("publish") + + async def test_the_whole_batch_lands_before_the_first_publish( + self, sweep_trace, seed_execution, authenticated_client: httpx.AsyncClient + ) -> None: + """One flush for the batch, and it is complete before anything is sent.""" + for i in range(3): + await seed_execution(task_name=f"task.{i}", status=TaskStatus.FAILED) + sweep_trace.clear() + + resp = await authenticated_client.post(RETRY_FAILED) + + assert resp.json() == {"queued": 3, "remaining": 0} + first_publish = sweep_trace.index("publish") + assert sweep_trace.count("publish") == 3 + assert sweep_trace[:first_publish].count("insert") >= 1 + assert "insert" not in sweep_trace[first_publish:] + + async def test_the_row_names_the_task_that_was_actually_sent( + self, + app, + sweep_trace, + seed_execution, + execution_rows, + authenticated_client: httpx.AsyncClient, + ) -> None: + """Writing first means reserving the broker id, not inventing one later. + + If the reserved id and the published id could differ, an orphaned row + would be unmatchable against the broker and writing first would buy + nothing. + """ + original = await seed_execution(status=TaskStatus.FAILED) + + await authenticated_client.post(RETRY_FAILED) + + new_row = next(r for r in await execution_rows() if r.id != original.id) + sent = app.state.background_tasks.celery.send_task.call_args + assert new_row.celery_task_id == sent.kwargs["task_id"] + + +class TestAnnouncedAfterCommit: + """Subscribers hear about a retry that happened, not one that might. + + Publishing inline told them about an uncommitted transaction — and for a + 500-row sweep, told them nothing until the batch was done and then + everything at once, for work that could still roll back. + """ + + async def test_the_event_fires_after_the_rows_are_committed( + self, app, sweep_trace, seed_execution, authenticated_client: httpx.AsyncClient + ) -> None: + async def _on_retried(event_: TaskRetried) -> None: + sweep_trace.append("announced") + + app.state.sm.event_bus.subscribe(TaskRetried, _on_retried) + await seed_execution(status=TaskStatus.FAILED) + sweep_trace.clear() + + await authenticated_client.post(RETRY_FAILED) + + assert "announced" in sweep_trace + assert sweep_trace.index("commit") < sweep_trace.index("announced") + + async def test_one_event_per_retried_row( + self, app, sweep_trace, seed_execution, authenticated_client: httpx.AsyncClient + ) -> None: + received: list[TaskRetried] = [] + + async def _on_retried(event_: TaskRetried) -> None: + received.append(event_) + + app.state.sm.event_bus.subscribe(TaskRetried, _on_retried) + for i in range(3): + await seed_execution(task_name=f"task.{i}", status=TaskStatus.FAILED) + + await authenticated_client.post(RETRY_FAILED) + + assert sorted(e.task_name for e in received) == ["task.0", "task.1", "task.2"] diff --git a/packages/i18n/src/generated-resources.ts b/packages/i18n/src/generated-resources.ts index 87499582..7d374ee3 100644 --- a/packages/i18n/src/generated-resources.ts +++ b/packages/i18n/src/generated-resources.ts @@ -120,7 +120,9 @@ export default { 'background_tasks.strip.queued': '', 'background_tasks.strip.running': '', 'background_tasks.strip.stuck': '', - 'background_tasks.strip.succeeded_24h': '', + 'background_tasks.strip.succeeded': '', + 'background_tasks.strip.window_24h': '', + 'background_tasks.strip.window_all': '', 'background_tasks.table.actions': '', 'background_tasks.table.duration': '', 'background_tasks.table.queue': '', diff --git a/packages/i18n/src/keys.generated.ts b/packages/i18n/src/keys.generated.ts index 5771ad52..c94f3400 100644 --- a/packages/i18n/src/keys.generated.ts +++ b/packages/i18n/src/keys.generated.ts @@ -160,7 +160,9 @@ export const keys = { queued: 'background_tasks.strip.queued', running: 'background_tasks.strip.running', stuck: 'background_tasks.strip.stuck', - succeeded_24h: 'background_tasks.strip.succeeded_24h', + succeeded: 'background_tasks.strip.succeeded', + window_24h: 'background_tasks.strip.window_24h', + window_all: 'background_tasks.strip.window_all', }, table: { actions: 'background_tasks.table.actions',