S.H.O.N.A.R._Desktop_Companion/backend/shonar/services/processing.py
avi 514d75fe92 Raw provider crash fails the job instead of escaping into the sweep requeue loop
A non-AIError from transcribe/summarize (av.InvalidDataError on corrupt
audio, seen live with 'My recording 63') escaped the consumer, left the
row 'running', and got requeued at every engine restart forever. Both
runners now catch it, fail the row with the exception type in the
message, and keep a usable transcript (summarize completes with a
'Summary failed' note). Regression tests pin both paths (proven red
without the fix).
2026-09-19 15:24:50 -05:00

742 lines
28 KiB
Python

"""AI processing pipeline (M7): uploaded -> transcribed -> summarized.
State lives in the database (ProcessingJob rows); arq/redis is transport
only. That ordering is deliberate: if redis is down, uploads still succeed
(the jobs sit queued) and the worker sweep picks them up. The only race is
a task running before the API transaction commits (recording invisible) —
missing rows are transient failures, so arq retry absorbs it.
Entry points:
- ``enqueue_for_recording``: called from upload finalize AND re-callable
later (manual transcripts in M8 re-enter here). Idempotent.
- ``run_transcribe`` / ``run_summarize``: arq task bodies. ``ctx`` is an
arq context in production and a plain dict in tests; only
``ctx.get("job_try", 1)`` is read.
"""
from __future__ import annotations
import asyncio
import logging
import uuid
from datetime import timedelta
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from shonar.core.config import get_settings
from shonar.db.models import (
Asset,
AssetKind,
JobStatus,
JobType,
ProcessingJob,
ProcessingStatus,
Recording,
Summary,
Transcript,
utcnow,
)
from shonar.services.ai import (
AIError,
ProviderConfigError,
ProviderTransientError,
get_llm_fallback_provider,
get_llm_provider,
get_transcription_provider,
)
from shonar.storage import get_storage
logger = logging.getLogger("shonar.processing")
# --- tone/request-integrity trace (pure logging; remove when the
# tone-mismatch investigation closes) ---------------------------------------
# POST /reprocess mints an rid per request and stashes it here keyed by
# (recording, job); the worker and store_summary read it back so all the
# lines for one click share the same rid= and can be grepped together.
# Single-process inline queue makes this safe; it carries no behavior.
_trace_rid: dict = {}
def trace_request(recording_id, job_type, tone):
"""Called from POST /reprocess. Returns the rid for the request log."""
rid = uuid.uuid4().hex[:8]
_trace_rid[(str(recording_id), str(job_type))] = rid
logger.info("TRACE reprocess rid=%s rec=%s job=%s tone=%r",
rid, recording_id, job_type, tone)
return rid
def trace_peek(recording_id, job_type):
return _trace_rid.get((str(recording_id), str(job_type)), "-")
def _accepts_on_progress(provider) -> bool:
"""True when the provider's transcribe() accepts on_progress=."""
import inspect
try:
return "on_progress" in inspect.signature(provider.transcribe).parameters
except (TypeError, ValueError): # builtins / exotic callables
return False
def _thread_progress_reporter(job_id):
"""Callback safe to invoke from a worker thread (faster-whisper runs
via asyncio.to_thread): schedules a tiny session update on the loop.
Progress writes are best-effort display state — failures are logged
and swallowed, never allowed to disturb the transcription itself."""
loop = asyncio.get_running_loop()
def report(pct: int) -> None:
async def _write() -> None:
from sqlalchemy import update
from shonar.db.session import session_factory
try:
async with session_factory()() as s:
await s.execute(
update(ProcessingJob)
.where(ProcessingJob.id == job_id)
.values(progress=int(pct))
)
await s.commit()
except Exception: # pragma: no cover - display state only
logger.debug("progress write failed for job %s", job_id, exc_info=True)
try:
asyncio.run_coroutine_threadsafe(_write(), loop)
except RuntimeError: # loop already gone (shutdown race)
logger.debug("progress dropped for job %s (loop gone)", job_id)
return report
def _async_progress_writer(job_id):
"""Progress callback for providers that stream on our own event loop
(the LLM adapters). Unlike _thread_progress_reporter there is no
worker thread: schedule the tiny write as a task on the running loop.
Same contract as the thread version — best-effort display state — and
identical values are coalesced so a chatty stream does not hammer
SQLite."""
loop = asyncio.get_running_loop()
last = -1
def report(pct: int) -> None:
nonlocal last
if int(pct) == last:
return
last = int(pct)
# TRACE (logging only): every distinct progress value the provider
# emits, with wall time — proves whether/when the 99 is emitted and
# written vs. the UI only ever showing 0%.
logger.info("TRACE progress job=%s pct=%s", job_id, pct)
async def _write() -> None:
from sqlalchemy import update
from shonar.db.session import session_factory
try:
async with session_factory()() as s:
await s.execute(
update(ProcessingJob)
.where(ProcessingJob.id == job_id)
.values(progress=last)
)
await s.commit()
except Exception: # pragma: no cover - display state only
logger.debug("summarize progress write failed for job %s",
job_id, exc_info=True)
try:
loop.create_task(_write())
except RuntimeError: # loop already gone (shutdown race)
logger.debug("summarize progress dropped for job %s", job_id)
return report
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 -----------------------------------------------------------------
async def enqueue_for_recording(session: AsyncSession, rec: Recording) -> list[JobType]:
"""Queue whatever AI stages apply. Safe to call repeatedly: completed
work is never redone, failed work is reset for another attempt."""
if rec.deleted_at is not None:
return []
settings = get_settings()
tprov = get_transcription_provider(settings)
lprov = get_llm_provider(settings)
queued: list[JobType] = []
if tprov is not None and await reset_or_create(session, rec, JobType.transcribe):
queued.append(JobType.transcribe)
if (
lprov is not None
and await latest_transcript_text(session, rec.id) is not None
and await reset_or_create(session, rec, JobType.summarize)
):
queued.append(JobType.summarize)
if tprov is None and lprov is None:
if rec.processing_status != ProcessingStatus.ai_disabled:
rec.processing_status = ProcessingStatus.ai_disabled
rec.processing_error = None
elif queued and rec.processing_status not in (
ProcessingStatus.processing,
ProcessingStatus.completed,
):
rec.processing_status = ProcessingStatus.queued
rec.processing_error = None
await session.flush()
for jt in queued:
await transport_enqueue(jt, rec.id)
return queued
async def reset_or_create(
session: AsyncSession, rec: Recording, job_type: JobType
) -> bool:
"""Ensure a queued job row. Returns True when (re)queued now: new rows,
plus failed/skipped rows (a re-upload or re-entry deserves another
attempt). Queued/running/succeeded rows are left alone."""
existing = await session.scalar(
select(ProcessingJob)
.where(
ProcessingJob.recording_id == rec.id,
ProcessingJob.job_type == job_type,
)
.order_by(ProcessingJob.id.desc())
)
if existing is None:
session.add(ProcessingJob(recording_id=rec.id, job_type=job_type))
return True
if existing.status in (JobStatus.queued, JobStatus.running, JobStatus.succeeded):
return False
existing.status = JobStatus.queued
existing.attempt = 0
existing.error = None
existing.stage = None
existing.progress = None
existing.started_at = None
existing.finished_at = None
return True
async def latest_transcript_text(
session: AsyncSession, recording_id: uuid.UUID
) -> str | None:
"""Newest non-superseded transcript; user-edited rows win over newer
machine rows (an edit is a verdict, not a draft)."""
rows = (
await session.scalars(
select(Transcript)
.where(
Transcript.recording_id == recording_id,
Transcript.superseded_at.is_(None),
)
.order_by(Transcript.version.desc())
)
).all()
if not rows:
return None
for r in rows:
if r.edited_by_user and r.text.strip():
return r.text
text = rows[0].text
return text if text.strip() else None
async def transport_enqueue(job_type: JobType, recording_id: uuid.UUID) -> None:
"""Best-effort trigger for the configured backend ('arq' or 'inline').
Failure only logs: the DB rows are the real queue and the sweep picks
up anything the transport missed."""
if get_settings().queue_backend.strip().lower() == "inline":
from shonar.services import inline_queue
await inline_queue.enqueue(job_type, str(recording_id))
return
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:
# 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
logger.warning("arq enqueue failed (%s); worker sweep will pick it up", e)
# --- tasks -----------------------------------------------------------------
def _job_try(ctx: dict) -> int:
try:
return int(ctx.get("job_try", 1))
except (TypeError, ValueError):
return 1
async def _load(session: AsyncSession, recording_id: str) -> Recording | None:
try:
rid = uuid.UUID(recording_id)
except ValueError:
return None
return await session.get(Recording, rid)
async def _job(
session: AsyncSession, rec: Recording, job_type: JobType
) -> ProcessingJob:
job = await session.scalar(
select(ProcessingJob)
.where(
ProcessingJob.recording_id == rec.id,
ProcessingJob.job_type == job_type,
)
.order_by(ProcessingJob.id.desc())
)
if job is None:
job = ProcessingJob(recording_id=rec.id, job_type=job_type)
session.add(job)
await session.flush()
return job
async def _fail(
session: AsyncSession,
rec: Recording,
job: ProcessingJob,
message: str,
ctx: dict,
exc: AIError | None = None,
) -> None:
"""Config errors fail now; transient errors fail only on the last try
(returning normally), otherwise they raise for arq retry."""
transient = exc is None or isinstance(exc, ProviderTransientError)
if transient and _job_try(ctx) < MAX_TRIES:
# Running state was committed before the long phase; put the row
# back to queued for the retry and persist that (a raise no longer
# rolls the pre-phase commit back).
job.status = JobStatus.queued
job.attempt = _job_try(ctx)
job.stage = None
await session.commit()
raise ProviderTransientError(message)
job.status = JobStatus.failed
job.error = message
job.stage = None
job.finished_at = utcnow()
if job.job_type == JobType.summarize and await latest_transcript_text(
session, rec.id
):
# The transcript is usable; a summary timeout must not mark the
# whole recording failed (the summary can be re-run separately).
rec.processing_status = ProcessingStatus.completed
rec.processing_error = f"Summary failed: {message}"
else:
rec.processing_status = ProcessingStatus.failed
rec.processing_error = message
await session.flush()
async def run_transcribe(ctx: dict, recording_id: str) -> None:
"""Transcribe the original audio; chain into summarization when an LLM
is configured."""
from shonar.db.session import session_factory
settings = get_settings()
async with session_factory()() as session:
rec = await _load(session, recording_id)
if rec is None or rec.deleted_at is not None:
# Finalize race: the API transaction may not have committed yet.
raise ProviderTransientError("Recording not ready; retrying.")
job = await _job(session, rec, JobType.transcribe)
provider = get_transcription_provider(settings)
if provider is None:
job.status = JobStatus.skipped
await session.flush()
await session.commit()
await _maybe_chain_summarize(session, rec)
return
# The recording's saved model is authoritative: a per-recording
# override wins, otherwise the global default in force at finalize
# time (changing the default never rewrites history).
from shonar.services.ai.model_registry import (
effective_model,
get_global_default_model,
)
model = effective_model(rec.transcription_model, await get_global_default_model(session))
if provider.name == "faster_whisper":
from shonar.services.ai.faster_whisper import FasterWhisperProvider
try:
provider = FasterWhisperProvider(model=model)
except AIError as e:
await _fail(session, rec, job, str(e), ctx, e)
await session.commit()
return
original = await session.scalar(
select(Asset).where(
Asset.recording_id == rec.id, Asset.kind == AssetKind.original
)
)
if original is None:
raise ProviderTransientError("Audio not ready; retrying.")
job.status = JobStatus.running
job.attempt = _job_try(ctx)
job.started_at = utcnow()
job.stage = "loading-model"
job.progress = None
rec.processing_status = ProcessingStatus.processing
rec.processing_error = None
# Commit before the long CPU phase: an open transaction is invisible
# to other readers (and on SQLite it locks out the progress writer).
await session.commit()
try:
audio = await get_storage().get(original.storage_key)
job.stage = "transcribing"
await session.commit()
kwargs = {}
if _accepts_on_progress(provider):
kwargs["on_progress"] = _thread_progress_reporter(job.id)
result = await provider.transcribe(audio, original.mime_type, **kwargs)
except AIError as e:
await _fail(session, rec, job, str(e), ctx, e)
await session.commit()
return
except asyncio.CancelledError:
raise
except Exception as e: # noqa: BLE001 — a raw provider crash (e.g. av.InvalidDataError
# on corrupt audio) must FAIL the job, not escape: an escaping
# exception leaves the row 'running' and the startup sweep
# requeues it forever (crash-loop seen live Sep 19).
logger.exception("transcribe provider raised for recording %s", rec.id)
job.status = JobStatus.failed
job.error = f"Transcription crashed: {type(e).__name__}: {e}"[:500]
job.stage = None
job.finished_at = utcnow()
rec.processing_status = ProcessingStatus.failed
rec.processing_error = job.error
await session.commit()
return
await store_transcript(
session, rec, result.text, result.segments, result.language,
provider.name, getattr(result, "model", ""),
)
job.status = JobStatus.succeeded
job.stage = None
job.progress = 100
job.finished_at = utcnow()
await session.flush()
await _maybe_chain_summarize(session, rec)
await session.commit()
async def _maybe_chain_summarize(session: AsyncSession, rec: Recording) -> None:
"""After transcription (or a skip): summarize when possible, else finish."""
settings = get_settings()
if get_llm_provider(settings) is None:
if rec.processing_status != ProcessingStatus.completed:
rec.processing_status = ProcessingStatus.completed
rec.processing_error = None
await session.flush()
return
if await latest_transcript_text(session, rec.id) is None:
# LLM configured but nothing to summarize (e.g. empty transcript).
existing = await session.scalar(
select(ProcessingJob).where(
ProcessingJob.recording_id == rec.id,
ProcessingJob.job_type == JobType.summarize,
)
)
if existing is not None and existing.status == JobStatus.queued:
existing.status = JobStatus.skipped
if rec.processing_status != ProcessingStatus.completed:
rec.processing_status = ProcessingStatus.completed
await session.flush()
return
if await reset_or_create(session, rec, JobType.summarize):
await transport_enqueue(JobType.summarize, rec.id)
# Status stays `processing` until the summarize task lands.
async def run_summarize(ctx: dict, recording_id: str) -> None:
"""Summarize the latest transcript into the structured summary shape."""
from shonar.db.session import session_factory
settings = get_settings()
async with session_factory()() as session:
rec = await _load(session, recording_id)
if rec is None or rec.deleted_at is not None:
raise ProviderTransientError("Recording not ready; retrying.")
job = await _job(session, rec, JobType.summarize)
provider = get_llm_provider(settings)
if provider is None:
job.status = JobStatus.skipped
await session.flush()
await session.commit()
return
text = await latest_transcript_text(session, rec.id)
if text is None:
job.status = JobStatus.skipped
await session.flush()
if rec.processing_status != ProcessingStatus.completed:
rec.processing_status = ProcessingStatus.completed
await session.commit()
return
job.status = JobStatus.running
job.attempt = _job_try(ctx)
job.started_at = utcnow()
job.stage = "summarizing"
rec.processing_status = ProcessingStatus.processing
await session.commit() # visible before the long LLM call
# TRACE (logging only): the tone the worker actually reads from the
# row NOW — differs from the reprocess line's tone if a later click
# overwrote it before pickup.
logger.info("TRACE summarize-start rid=%s job=%s rec=%s type=%s "
"tone=%r attempt=%s",
trace_peek(rec.id, job.job_type), job.id, rec.id,
job.job_type, job.tone, job.attempt)
fallback_note: str | None = None
result_provider = provider.name # overwritten only on the rescue path
try:
result = await provider.summarize(
text, title=rec.title, tone=job.tone,
on_progress=_async_progress_writer(job.id))
except AIError as e:
# Transient failures keep their normal retry budget first (the
# LAN GPU often recovers on try 2); the rescue summarizer runs
# when the primary is OUT of tries or misconfigured, and only
# if one is configured. The summary then comes from the
# fallback and job.error carries a note for the UI.
retryable = (isinstance(e, ProviderTransientError)
and _job_try(ctx) < MAX_TRIES)
# (result, provider_name, note) when the rescue succeeds.
rescue: tuple[object, str, str] | None = None
if not retryable:
fallback = get_llm_fallback_provider(settings)
if fallback is not None:
try:
fb_result = await fallback.summarize(
text, title=rec.title, tone=job.tone,
on_progress=_async_progress_writer(job.id))
except AIError as fe:
logger.warning(
"summarize fallback (%s) also failed: %s",
fallback.name, fe)
else:
rescue = (fb_result, fallback.name, (
f'primary "{provider.name}" failed ({e}). '
f'This summary was written by "{fallback.name}".'))
if rescue is None:
await _fail(session, rec, job, str(e), ctx, e)
await session.commit()
return
result, result_provider, fallback_note = rescue
except asyncio.CancelledError:
raise
except Exception as e: # noqa: BLE001 — same reason as transcribe: a raw
# provider crash must fail the row, not escape into the sweep loop.
logger.exception("summarize provider raised for recording %s", rec.id)
job.status = JobStatus.failed
job.error = f"Summarize crashed: {type(e).__name__}: {e}"[:500]
job.stage = None
job.finished_at = utcnow()
# A usable transcript survives: complete with the failure note
# (mirrors _fail's transcript-preserving rule).
if await latest_transcript_text(session, rec.id):
rec.processing_status = ProcessingStatus.completed
rec.processing_error = f"Summary failed: {job.error}"
else:
rec.processing_status = ProcessingStatus.failed
rec.processing_error = job.error
await session.commit()
return
await store_summary(session, rec, result.to_dict(), result_provider,
result.model, tone=job.tone)
job.status = JobStatus.succeeded
job.stage = None
job.progress = 100
job.error = fallback_note # None on the clean path: clears stale notes
job.finished_at = utcnow()
rec.processing_status = ProcessingStatus.completed
rec.processing_error = None
await session.flush()
await session.commit()
async def store_transcript(
session: AsyncSession,
rec: Recording,
text: str,
segments: list,
language: str | None,
provider_name: str,
model: str,
) -> None:
"""Insert a new auto version; a newest user-edited row wins instead and
nothing is inserted (edits are verdicts)."""
existing = (
await session.scalars(
select(Transcript)
.where(
Transcript.recording_id == rec.id,
Transcript.superseded_at.is_(None),
)
.order_by(Transcript.version.desc())
)
).all()
if existing and existing[0].edited_by_user:
return
now = utcnow()
max_version = await session.scalar(
select(func.max(Transcript.version)).where(Transcript.recording_id == rec.id)
)
for row in existing:
row.superseded_at = now
session.add(
Transcript(
recording_id=rec.id,
version=(max_version or 0) + 1,
language=language,
provider=provider_name,
model=model,
text=text,
segments=[
{"start": s.start, "end": s.end, "text": s.text, "speaker": s.speaker}
for s in segments
],
edited_by_user=False,
)
)
await session.flush()
async def store_summary(
session: AsyncSession,
rec: Recording,
content: dict,
provider_name: str,
model: str,
tone: str | None = None,
) -> None:
# TRACE (logging only): the tone arriving at storage. Compare with the
# summarize-start line's tone for the same rid: a mismatch here means
# the value changed between worker pickup and the DB write.
logger.info("TRACE store_summary rid=%s rec=%s tone=%r provider=%s",
trace_peek(rec.id, "summarize"), rec.id, tone,
provider_name)
existing = (
await session.scalars(
select(Summary)
.where(
Summary.recording_id == rec.id,
Summary.superseded_at.is_(None),
)
.order_by(Summary.version.desc())
)
).all()
if existing and existing[0].edited_by_user:
return
now = utcnow()
max_version = await session.scalar(
select(func.max(Summary.version)).where(Summary.recording_id == rec.id)
)
for row in existing:
row.superseded_at = now
session.add(
Summary(
recording_id=rec.id,
version=(max_version or 0) + 1,
provider=provider_name,
model=model,
tone=tone,
content=content,
edited_by_user=False,
)
)
await session.flush()
async def sweep_stale(limit: int = 100) -> int:
"""Crash recovery + transport-loss backstop: requeue jobs stuck running
or sitting queued, newest catastrophe first. Returns jobs re-enqueued."""
from shonar.db.session import session_factory
count = 0
async with session_factory()() as session:
rows = (
await session.scalars(
select(ProcessingJob)
.where(ProcessingJob.status.in_([JobStatus.queued, JobStatus.running]))
.order_by(ProcessingJob.id.desc())
.limit(limit)
)
).all()
for job in rows:
rec = await session.get(Recording, job.recording_id)
if rec is None or rec.deleted_at is not None:
job.status = JobStatus.skipped
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
await session.commit()
# Transport outside the transaction: rows are the queue, this just pokes.
async with session_factory()() as session:
rows = (
await session.scalars(
select(ProcessingJob)
.where(ProcessingJob.status == JobStatus.queued)
.order_by(ProcessingJob.id.desc())
.limit(limit)
)
).all()
for job in rows:
await transport_enqueue(job.job_type, job.recording_id)
return count