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