From 23714cfe7c9e07e22f2aea504b5c5ef52c691929 Mon Sep 17 00:00:00 2001 From: avi Date: Sat, 12 Sep 2026 17:44:12 -0500 Subject: [PATCH] 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. --- backend/shonar/services/processing.py | 44 +++++++++++++++++++++++---- backend/shonar/worker.py | 5 +++ backend/tests/test_ai_pipeline.py | 35 +++++++++++++++++++++ 3 files changed, 78 insertions(+), 6 deletions(-) diff --git a/backend/shonar/services/processing.py b/backend/shonar/services/processing.py index cf888ca..25a1a13 100644 --- a/backend/shonar/services/processing.py +++ b/backend/shonar/services/processing.py @@ -18,6 +18,7 @@ from __future__ import annotations import logging import uuid +from datetime import timedelta from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession @@ -47,6 +48,11 @@ logger = logging.getLogger("shonar.processing") 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 ----------------------------------------------------------------- @@ -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.""" from arq import create_pool from arq.connections import RedisSettings + from arq.constants import result_key_prefix + job_id = f"{job_type.value}:{recording_id}" try: pool = await create_pool(RedisSettings.from_dsn(get_settings().redis_url)) 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: await pool.aclose() 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: job.status = JobStatus.skipped continue - if job.status == JobStatus.running and job.attempt >= MAX_TRIES: - job.status = JobStatus.failed - job.error = "Worker died too many times." - job.finished_at = utcnow() - continue + if job.status == JobStatus.running: + if job.attempt >= MAX_TRIES: + job.status = JobStatus.failed + job.error = "Worker died too many times." + 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.error = None count += 1 diff --git a/backend/shonar/worker.py b/backend/shonar/worker.py index 68fcef6..31b2c4e 100644 --- a/backend/shonar/worker.py +++ b/backend/shonar/worker.py @@ -60,3 +60,8 @@ class WorkerSettings: # Retry budget for transient provider failures; the tasks themselves # mark jobs failed on the last try (see processing.MAX_TRIES). 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 diff --git a/backend/tests/test_ai_pipeline.py b/backend/tests/test_ai_pipeline.py index c4c1af7..52e4755 100644 --- a/backend/tests/test_ai_pipeline.py +++ b/backend/tests/test_ai_pipeline.py @@ -359,6 +359,41 @@ async def test_sweep_requeues_stale_jobs(client, monkeypatch): 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 ---------------------------------------------------------------------