generated from john/python-template
worker updates
This commit is contained in:
+27
-18
@@ -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
@@ -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
@@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user