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
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()
+14 -12
View File
@@ -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:
+45 -49
View File
@@ -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,
)
)