Fix worker copy storm: dedupe arq enqueues, sweep grace window, max_jobs=1

sweep_stale re-enqueued every queued/running job every 5 minutes with
undifferentiated enqueue_job ids, so a long transcription accumulated a
duplicate copy per sweep; arq's default max_jobs=10 then ran them all
concurrently in one worker (each a multi-GB local whisper pass) and the
worker ballooned to 11 GB + 500% CPU, swapping the box.

- transport_enqueue: deterministic _job_id (type:recording) so arq
  refuses duplicate copies while one is queued/running; stale result key
  is dropped first so deliberate re-runs of finished work still enqueue.
- sweep_stale: running jobs started within STALE_RUNNING_AFTER (2h) are
  live work, not crash relics, and are left alone.
- WorkerSettings.max_jobs = 1: one whisper pass per worker process;
  scale via more worker processes.
- regression test: sweep leaves a live running job, requeues an orphan.
This commit is contained in:
avi 2026-09-12 17:44:12 -05:00
commit 23714cfe7c
3 changed files with 78 additions and 6 deletions

View file

@ -18,6 +18,7 @@ from __future__ import annotations
import logging import logging
import uuid import uuid
from datetime import timedelta
from sqlalchemy import func, select from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -47,6 +48,11 @@ logger = logging.getLogger("shonar.processing")
MAX_TRIES = 3 MAX_TRIES = 3
# A `running` job younger than this is treated as live work, not a crash
# relic — long transcriptions legitimately outrun the 5-minute sweep, and
# re-enqueuing them mid-flight caused a copy storm (one whisper per copy).
STALE_RUNNING_AFTER = timedelta(hours=2)
# --- entry ----------------------------------------------------------------- # --- entry -----------------------------------------------------------------
@ -140,11 +146,28 @@ async def transport_enqueue(job_type: JobType, recording_id: uuid.UUID) -> None:
queue and the worker sweep picks up anything the transport missed.""" queue and the worker sweep picks up anything the transport missed."""
from arq import create_pool from arq import create_pool
from arq.connections import RedisSettings from arq.connections import RedisSettings
from arq.constants import result_key_prefix
job_id = f"{job_type.value}:{recording_id}"
try: try:
pool = await create_pool(RedisSettings.from_dsn(get_settings().redis_url)) pool = await create_pool(RedisSettings.from_dsn(get_settings().redis_url))
try: try:
await pool.enqueue_job("run_" + job_type.value, str(recording_id)) # Deterministic _job_id: arq refuses a duplicate while a copy of
# (job_type, recording) is queued or running, so sweep pokes and
# retried enqueues can never pile up concurrent copies of the
# same work.
#
# arq dedupes on job key OR result key, and failed runs leave a
# result behind for keep_result days — that would silently swallow
# deliberate re-runs of previously-finished work. Drop only the
# stale RESULT key first (never the job key: that is exactly what
# shields a live copy from duplicate enqueue).
await pool.delete(result_key_prefix + job_id)
await pool.enqueue_job(
"run_" + job_type.value,
str(recording_id),
_job_id=job_id,
)
finally: finally:
await pool.aclose() await pool.aclose()
except Exception as e: # noqa: BLE001 — transport must never break uploads except Exception as e: # noqa: BLE001 — transport must never break uploads
@ -437,11 +460,20 @@ async def sweep_stale(limit: int = 100) -> int:
if rec is None or rec.deleted_at is not None: if rec is None or rec.deleted_at is not None:
job.status = JobStatus.skipped job.status = JobStatus.skipped
continue continue
if job.status == JobStatus.running and job.attempt >= MAX_TRIES: if job.status == JobStatus.running:
job.status = JobStatus.failed if job.attempt >= MAX_TRIES:
job.error = "Worker died too many times." job.status = JobStatus.failed
job.finished_at = utcnow() job.error = "Worker died too many times."
continue job.finished_at = utcnow()
continue
# Live work, not a crash relic: a running job that started
# recently belongs to a worker still chewing it (long audio
# outruns the sweep cadence). Only orphaned runs — no
# started_at, or older than the grace window — get requeued.
if job.started_at is not None and (
utcnow() - job.started_at
) < STALE_RUNNING_AFTER:
continue
job.status = JobStatus.queued job.status = JobStatus.queued
job.error = None job.error = None
count += 1 count += 1

View file

@ -60,3 +60,8 @@ class WorkerSettings:
# Retry budget for transient provider failures; the tasks themselves # Retry budget for transient provider failures; the tasks themselves
# mark jobs failed on the last try (see processing.MAX_TRIES). # mark jobs failed on the last try (see processing.MAX_TRIES).
max_tries = 3 max_tries = 3
# One job per worker process at a time. Local faster-whisper is
# multi-GB per concurrent pass; default max_jobs=10 let one worker run
# ~8 transcriptions at once and OOM-swap the box. Scale by running more
# worker processes, never by raising this.
max_jobs = 1

View file

@ -359,6 +359,41 @@ async def test_sweep_requeues_stale_jobs(client, monkeypatch):
assert jobs[JobType.summarize].status == JobStatus.queued assert jobs[JobType.summarize].status == JobStatus.queued
async def test_sweep_leaves_live_running_jobs_alone(client, monkeypatch):
"""Regression: sweep re-enqueued every in-flight transcription every
5 minutes; with local whisper each copy ran concurrently and one worker
OOM-swap-swelled the box. A running job started recently is live work."""
from datetime import UTC, datetime, timedelta
use_fakes(monkeypatch, FakeTranscriber(), FakeLlm())
token = await user_tokens(client)
rec = await upload_recording(client, token, client_id="m7-sweep-live")
rid = uuid.UUID(rec["id"])
async with db_session._session_factory() as s:
from shonar.db.models import Recording
row = await s.get(Recording, rid)
await processing.reset_or_create(s, row, JobType.transcribe)
live = (await s.scalars(
select(ProcessingJob).where(
ProcessingJob.recording_id == rid,
ProcessingJob.job_type == JobType.transcribe)
)).one()
live.status = JobStatus.running
live.started_at = datetime.now(UTC) # worker is chewing it right now
stale = ProcessingJob(recording_id=rid, job_type=JobType.summarize,
status=JobStatus.running, attempt=1)
stale.started_at = datetime.now(UTC) - timedelta(
hours=processing.STALE_RUNNING_AFTER.total_seconds() / 3600 + 1)
s.add(stale)
await s.commit()
count = await processing.sweep_stale()
assert count == 1 # only the orphaned old run
jobs = await jobs_for(rec["id"])
assert jobs[JobType.transcribe].status == JobStatus.running
assert jobs[JobType.summarize].status == JobStatus.queued
# --- ownership --------------------------------------------------------------------- # --- ownership ---------------------------------------------------------------------