From 755f908b6add16291cc461707dd718db4467b317 Mon Sep 17 00:00:00 2001 From: John Lancaster <32917998+jsl12@users.noreply.github.com> Date: Fri, 26 Jun 2026 19:17:44 -0500 Subject: [PATCH] worker updates --- src/transcription/app.py | 45 +++++++++++------- src/transcription/db.py | 26 +++++----- src/transcription/worker.py | 94 ++++++++++++++++++------------------- 3 files changed, 86 insertions(+), 79 deletions(-) diff --git a/src/transcription/app.py b/src/transcription/app.py index 9aa28d6..72a8815 100644 --- a/src/transcription/app.py +++ b/src/transcription/app.py @@ -2,9 +2,9 @@ from __future__ import annotations +import asyncio from contextlib import asynccontextmanager -from threading import Event -from threading import Thread +from contextlib import suppress from fastapi import FastAPI from sqlalchemy.ext.asyncio import async_sessionmaker @@ -23,29 +23,38 @@ from .worker import run_worker_loop def _start_worker(app: FastAPI) -> None: session_factory: async_sessionmaker[AsyncSession] = app.state.db_session_factory - stop_event = Event() - worker_thread = Thread( - target=run_worker_loop, - kwargs={ - "session_factory": session_factory, - "stop_event": stop_event, - "poll_interval_seconds": 1.0, - }, - daemon=True, + stop_event = asyncio.Event() + wake_queue: asyncio.Queue[None] = asyncio.Queue() + worker_task = asyncio.create_task( + run_worker_loop( + session_factory=session_factory, + stop_event=stop_event, + wake_queue=wake_queue, + poll_interval_seconds=1.0, + ) ) - worker_thread.start() + wake_queue.put_nowait(None) app.state.worker_stop_event = stop_event - app.state.worker_thread = worker_thread + app.state.worker_wake_queue = wake_queue + app.state.worker_task = worker_task -def _stop_worker(app: FastAPI) -> None: +async def _stop_worker(app: FastAPI) -> None: stop_event = getattr(app.state, "worker_stop_event", None) - worker_thread = getattr(app.state, "worker_thread", None) + wake_queue = getattr(app.state, "worker_wake_queue", None) + worker_task = getattr(app.state, "worker_task", None) if stop_event is not None: stop_event.set() - if worker_thread is not None: - worker_thread.join(timeout=2.0) + if wake_queue is not None: + wake_queue.put_nowait(None) + if worker_task is not None: + try: + await asyncio.wait_for(worker_task, timeout=2.0) + except TimeoutError: + worker_task.cancel() + with suppress(asyncio.CancelledError): + await worker_task @asynccontextmanager @@ -68,7 +77,7 @@ async def _lifespan(app: FastAPI): try: yield finally: - _stop_worker(app) + await _stop_worker(app) await cleanup_database() diff --git a/src/transcription/db.py b/src/transcription/db.py index 0f3ce16..ec11fd7 100644 --- a/src/transcription/db.py +++ b/src/transcription/db.py @@ -9,6 +9,7 @@ from __future__ import annotations import contextlib import logging from collections.abc import AsyncGenerator +from contextvars import ContextVar from dataclasses import dataclass from sqlalchemy import inspect @@ -34,7 +35,7 @@ class DatabaseRuntime: session_factory: async_sessionmaker[AsyncSession] -_runtime: DatabaseRuntime | None = None +_runtime: ContextVar[DatabaseRuntime | None] = ContextVar("database_runtime", default=None) def _to_async_database_url(database_url: str) -> str: @@ -61,27 +62,28 @@ def _build_engine(settings: Settings) -> AsyncEngine: def initialize_database_runtime(*, settings: Settings | None = None) -> DatabaseRuntime: """Initialize lifespan-owned async DB resources once per process.""" - global _runtime - if _runtime is not None: - return _runtime + runtime = _runtime.get() + if runtime is not None: + return runtime active_settings = settings or get_settings() engine = _build_engine(active_settings) session_factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) - _runtime = DatabaseRuntime(engine=engine, session_factory=session_factory) + runtime = DatabaseRuntime(engine=engine, session_factory=session_factory) + _runtime.set(runtime) logger.debug("Initialized async database runtime for database_url=%s", engine.url) - return _runtime + return runtime def get_engine() -> AsyncEngine: """Return the current async SQLAlchemy engine.""" - runtime = _runtime or initialize_database_runtime() + runtime = _runtime.get() or initialize_database_runtime() return runtime.engine def get_session_factory() -> async_sessionmaker[AsyncSession]: """Return the shared async session factory.""" - runtime = _runtime or initialize_database_runtime() + runtime = _runtime.get() or initialize_database_runtime() return runtime.session_factory @@ -92,11 +94,11 @@ async def cleanup_database() -> None: async def dispose_database_runtime() -> None: """Dispose lifespan-owned async database resources.""" - global _runtime - if _runtime is None: + runtime = _runtime.get() + if runtime is None: return - await _runtime.engine.dispose() - _runtime = None + await runtime.engine.dispose() + _runtime.set(None) async def create_all(*, engine: AsyncEngine | None = None) -> None: diff --git a/src/transcription/worker.py b/src/transcription/worker.py index fbe8d39..33db0c6 100644 --- a/src/transcription/worker.py +++ b/src/transcription/worker.py @@ -4,12 +4,12 @@ from __future__ import annotations import asyncio import logging +from contextlib import suppress from datetime import UTC from datetime import datetime -from threading import Event -from pydantic import ValidationError from sqlalchemy.ext.asyncio import async_sessionmaker +from sqlalchemy.orm import selectinload from sqlmodel import select from sqlmodel.ext.asyncio.session import AsyncSession @@ -20,7 +20,6 @@ from transcription.errors import AppError from transcription.errors import ErrorCategory from transcription.errors import classify_unexpected_error from transcription.errors import format_error_detail -from transcription.models import Document from transcription.models import Job from transcription.models import JobStatus from transcription.models import Transcript @@ -29,6 +28,35 @@ from transcription.services.transcription import transcribe_document_image logger = logging.getLogger(__name__) +async def run_worker_loop( + *, + session_factory: async_sessionmaker[AsyncSession] | None = None, + stop_event: asyncio.Event | None = None, + wake_queue: asyncio.Queue[None] | None = None, + poll_interval_seconds: float = 1.0, +) -> None: + """Run worker loop until stop_event is set. + + If wake_queue is provided, queue activity wakes the loop immediately while + timeout-based wakeups preserve current polling behavior. + """ + while True: + if stop_event is not None and stop_event.is_set(): + logger.info("Worker stop event received") + return + + if wake_queue is not None: + with suppress(TimeoutError): + await asyncio.wait_for(wake_queue.get(), timeout=poll_interval_seconds) + + processed_any = False + while await process_next_queued_job(session_factory=session_factory): + processed_any = True + + if wake_queue is None and not processed_any: + await asyncio.sleep(poll_interval_seconds) + + async def process_next_queued_job( *, session: AsyncSession | None = None, @@ -45,7 +73,13 @@ async def process_next_queued_job( async def _process_next_queued_job(*, session: AsyncSession) -> bool: - job = (await session.exec(select(Job).where(Job.status == JobStatus.QUEUED).order_by(Job.created_at))).first() + query = ( + select(Job) + .options(selectinload(Job.document)) # pyright: ignore[reportArgumentType] + .where(Job.status == JobStatus.QUEUED) + .order_by(Job.created_at) # pyright: ignore[reportArgumentType] + ) + job = (await session.exec(query)).first() if job is None: return False @@ -57,14 +91,14 @@ async def _process_next_queued_job(*, session: AsyncSession) -> bool: await session.commit() await session.refresh(job) - document = await session.get(Document, job.document_id) + document = job.document if document is None: error = AppError( "Document not found", category=ErrorCategory.NOT_FOUND, suggestion="Re-upload the source document and retry processing.", ) - _finalize_failed_job(session=session, job=job, error=error) + await _finalize_failed_job(session=session, job=job, error=error) logger.error( "Job failed operation=worker.process_job job_id=%s error_id=%s category=%s", job.id, @@ -74,7 +108,8 @@ async def _process_next_queued_job(*, session: AsyncSession) -> bool: return True try: - result = transcribe_document_image(document.file_path) + # Provider SDK calls are synchronous and should not block the event loop. + result = await asyncio.to_thread(transcribe_document_image, document.file_path) await _upsert_transcript(session=session, job_id=job.id, text=result.text, error_detail=None) job.status = JobStatus.TRANSCRIBED job.updated_at = datetime.now(UTC) @@ -88,11 +123,12 @@ async def _process_next_queued_job(*, session: AsyncSession) -> bool: ) except Exception as exc: error = exc if isinstance(exc, AppError) else classify_unexpected_error(exc, operation="worker.process_job") - settings = _get_worker_settings() + settings = get_settings() if _should_retry(job=job, error=error, settings=settings): await _requeue_for_retry(session=session, job=job, error=error, settings=settings) logger.warning( - "Job retried operation=worker.process_job job_id=%s document_id=%s retry_count=%s error_id=%s category=%s", + "Job retried operation=worker.process_job " + "job_id=%s document_id=%s retry_count=%s error_id=%s category=%s", job.id, document.id, job.retry_count, @@ -127,13 +163,6 @@ async def _upsert_transcript( return transcript -def _get_worker_settings() -> Settings: - try: - return get_settings() - except ValidationError: - return Settings(openrouter_api_key="test-key") - - def _should_retry(*, job: Job, error: AppError, settings: Settings) -> bool: return error.retriable and job.retry_count < settings.worker_max_retries @@ -155,36 +184,3 @@ async def _finalize_failed_job(*, session: AsyncSession, job: Job, error: AppErr job.updated_at = datetime.now(UTC) session.add(job) await session.commit() - - -async def _run_worker_loop_async( - *, - session_factory: async_sessionmaker[AsyncSession] | None = None, - stop_event: Event | None = None, - poll_interval_seconds: float = 1.0, -) -> None: - """Run worker polling loop until stop_event is set.""" - while True: - if stop_event is not None and stop_event.is_set(): - logger.info("Worker stop event received") - return - - processed = await process_next_queued_job(session_factory=session_factory) - if not processed: - await asyncio.sleep(poll_interval_seconds) - - -def run_worker_loop( - *, - session_factory: async_sessionmaker[AsyncSession] | None = None, - stop_event: Event | None = None, - poll_interval_seconds: float = 1.0, -) -> None: - """Synchronous thread entrypoint that runs the async worker loop.""" - asyncio.run( - _run_worker_loop_async( - session_factory=session_factory, - stop_event=stop_event, - poll_interval_seconds=poll_interval_seconds, - ) - )