worker updates

This commit is contained in:
John Lancaster
2026-06-26 19:17:44 -05:00
parent a16c6f5ecd
commit 755f908b6a
3 changed files with 86 additions and 79 deletions
+27 -18
View File
@@ -2,9 +2,9 @@
from __future__ import annotations from __future__ import annotations
import asyncio
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from threading import Event from contextlib import suppress
from threading import Thread
from fastapi import FastAPI from fastapi import FastAPI
from sqlalchemy.ext.asyncio import async_sessionmaker from sqlalchemy.ext.asyncio import async_sessionmaker
@@ -23,29 +23,38 @@ from .worker import run_worker_loop
def _start_worker(app: FastAPI) -> None: def _start_worker(app: FastAPI) -> None:
session_factory: async_sessionmaker[AsyncSession] = app.state.db_session_factory session_factory: async_sessionmaker[AsyncSession] = app.state.db_session_factory
stop_event = Event() stop_event = asyncio.Event()
worker_thread = Thread( wake_queue: asyncio.Queue[None] = asyncio.Queue()
target=run_worker_loop, worker_task = asyncio.create_task(
kwargs={ run_worker_loop(
"session_factory": session_factory, session_factory=session_factory,
"stop_event": stop_event, stop_event=stop_event,
"poll_interval_seconds": 1.0, wake_queue=wake_queue,
}, poll_interval_seconds=1.0,
daemon=True,
) )
worker_thread.start() )
wake_queue.put_nowait(None)
app.state.worker_stop_event = stop_event 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) 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: if stop_event is not None:
stop_event.set() stop_event.set()
if worker_thread is not None: if wake_queue is not None:
worker_thread.join(timeout=2.0) 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 @asynccontextmanager
@@ -68,7 +77,7 @@ async def _lifespan(app: FastAPI):
try: try:
yield yield
finally: finally:
_stop_worker(app) await _stop_worker(app)
await cleanup_database() await cleanup_database()
+14 -12
View File
@@ -9,6 +9,7 @@ from __future__ import annotations
import contextlib import contextlib
import logging import logging
from collections.abc import AsyncGenerator from collections.abc import AsyncGenerator
from contextvars import ContextVar
from dataclasses import dataclass from dataclasses import dataclass
from sqlalchemy import inspect from sqlalchemy import inspect
@@ -34,7 +35,7 @@ class DatabaseRuntime:
session_factory: async_sessionmaker[AsyncSession] 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: 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: def initialize_database_runtime(*, settings: Settings | None = None) -> DatabaseRuntime:
"""Initialize lifespan-owned async DB resources once per process.""" """Initialize lifespan-owned async DB resources once per process."""
global _runtime runtime = _runtime.get()
if _runtime is not None: if runtime is not None:
return _runtime return runtime
active_settings = settings or get_settings() active_settings = settings or get_settings()
engine = _build_engine(active_settings) engine = _build_engine(active_settings)
session_factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) 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) logger.debug("Initialized async database runtime for database_url=%s", engine.url)
return _runtime return runtime
def get_engine() -> AsyncEngine: def get_engine() -> AsyncEngine:
"""Return the current async SQLAlchemy engine.""" """Return the current async SQLAlchemy engine."""
runtime = _runtime or initialize_database_runtime() runtime = _runtime.get() or initialize_database_runtime()
return runtime.engine return runtime.engine
def get_session_factory() -> async_sessionmaker[AsyncSession]: def get_session_factory() -> async_sessionmaker[AsyncSession]:
"""Return the shared async session factory.""" """Return the shared async session factory."""
runtime = _runtime or initialize_database_runtime() runtime = _runtime.get() or initialize_database_runtime()
return runtime.session_factory return runtime.session_factory
@@ -92,11 +94,11 @@ async def cleanup_database() -> None:
async def dispose_database_runtime() -> None: async def dispose_database_runtime() -> None:
"""Dispose lifespan-owned async database resources.""" """Dispose lifespan-owned async database resources."""
global _runtime runtime = _runtime.get()
if _runtime is None: if runtime is None:
return return
await _runtime.engine.dispose() await runtime.engine.dispose()
_runtime = None _runtime.set(None)
async def create_all(*, engine: AsyncEngine | None = None) -> None: async def create_all(*, engine: AsyncEngine | None = None) -> None:
+45 -49
View File
@@ -4,12 +4,12 @@ from __future__ import annotations
import asyncio import asyncio
import logging import logging
from contextlib import suppress
from datetime import UTC from datetime import UTC
from datetime import datetime from datetime import datetime
from threading import Event
from pydantic import ValidationError
from sqlalchemy.ext.asyncio import async_sessionmaker from sqlalchemy.ext.asyncio import async_sessionmaker
from sqlalchemy.orm import selectinload
from sqlmodel import select from sqlmodel import select
from sqlmodel.ext.asyncio.session import AsyncSession 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 ErrorCategory
from transcription.errors import classify_unexpected_error from transcription.errors import classify_unexpected_error
from transcription.errors import format_error_detail from transcription.errors import format_error_detail
from transcription.models import Document
from transcription.models import Job from transcription.models import Job
from transcription.models import JobStatus from transcription.models import JobStatus
from transcription.models import Transcript from transcription.models import Transcript
@@ -29,6 +28,35 @@ from transcription.services.transcription import transcribe_document_image
logger = logging.getLogger(__name__) 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( async def process_next_queued_job(
*, *,
session: AsyncSession | None = None, session: AsyncSession | None = None,
@@ -45,7 +73,13 @@ async def process_next_queued_job(
async def _process_next_queued_job(*, session: AsyncSession) -> bool: 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: if job is None:
return False return False
@@ -57,14 +91,14 @@ async def _process_next_queued_job(*, session: AsyncSession) -> bool:
await session.commit() await session.commit()
await session.refresh(job) await session.refresh(job)
document = await session.get(Document, job.document_id) document = job.document
if document is None: if document is None:
error = AppError( error = AppError(
"Document not found", "Document not found",
category=ErrorCategory.NOT_FOUND, category=ErrorCategory.NOT_FOUND,
suggestion="Re-upload the source document and retry processing.", 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( logger.error(
"Job failed operation=worker.process_job job_id=%s error_id=%s category=%s", "Job failed operation=worker.process_job job_id=%s error_id=%s category=%s",
job.id, job.id,
@@ -74,7 +108,8 @@ async def _process_next_queued_job(*, session: AsyncSession) -> bool:
return True return True
try: 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) await _upsert_transcript(session=session, job_id=job.id, text=result.text, error_detail=None)
job.status = JobStatus.TRANSCRIBED job.status = JobStatus.TRANSCRIBED
job.updated_at = datetime.now(UTC) job.updated_at = datetime.now(UTC)
@@ -88,11 +123,12 @@ async def _process_next_queued_job(*, session: AsyncSession) -> bool:
) )
except Exception as exc: except Exception as exc:
error = exc if isinstance(exc, AppError) else classify_unexpected_error(exc, operation="worker.process_job") 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): if _should_retry(job=job, error=error, settings=settings):
await _requeue_for_retry(session=session, job=job, error=error, settings=settings) await _requeue_for_retry(session=session, job=job, error=error, settings=settings)
logger.warning( 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, job.id,
document.id, document.id,
job.retry_count, job.retry_count,
@@ -127,13 +163,6 @@ async def _upsert_transcript(
return 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: def _should_retry(*, job: Job, error: AppError, settings: Settings) -> bool:
return error.retriable and job.retry_count < settings.worker_max_retries 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) job.updated_at = datetime.now(UTC)
session.add(job) session.add(job)
await session.commit() 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,
)
)