generated from john/python-template
Updated test suite
This commit is contained in:
@@ -4,6 +4,10 @@ from __future__ import annotations
|
||||
|
||||
from contextlib import AsyncExitStack
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import UTC
|
||||
from datetime import datetime
|
||||
from datetime import timedelta
|
||||
import logging
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi import status
|
||||
@@ -18,10 +22,14 @@ from .db import create_all
|
||||
from .db import dispose_database_runtime
|
||||
from .db import initialize_database_runtime
|
||||
from .services import ServiceBundle
|
||||
from .services.jobs import JobService
|
||||
from .ui import register_pages
|
||||
from .worker import worker_consumer_lifespan
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _lifespan(app: FastAPI):
|
||||
configure_logging()
|
||||
@@ -37,6 +45,8 @@ async def _lifespan(app: FastAPI):
|
||||
settings.upload_dir.mkdir(parents=True, exist_ok=True)
|
||||
settings.prompt_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
await _recover_stale_processing_jobs(app)
|
||||
|
||||
async with AsyncExitStack() as stack:
|
||||
stack.push_async_callback(dispose_database_runtime)
|
||||
stop_event, worker_notifier = await stack.enter_async_context(
|
||||
@@ -50,6 +60,20 @@ async def _lifespan(app: FastAPI):
|
||||
yield
|
||||
|
||||
|
||||
async def _recover_stale_processing_jobs(app: FastAPI) -> None:
|
||||
"""Re-queue stale processing jobs at startup.
|
||||
|
||||
Any job left in PROCESSING longer than the configured provider timeout is
|
||||
assumed orphaned and moved back to QUEUED before the worker starts.
|
||||
"""
|
||||
settings = app.state.settings
|
||||
stale_before = datetime.now(UTC) - timedelta(seconds=settings.worker_provider_timeout_seconds)
|
||||
job_service = JobService(session_factory=app.state.runtime.session_factory)
|
||||
recovered = await job_service.requeue_stale_processing_jobs(stale_before=stale_before)
|
||||
if recovered > 0:
|
||||
logger.warning("Recovered %s stale processing job(s) at startup", recovered)
|
||||
|
||||
|
||||
def create_app() -> FastAPI:
|
||||
"""Create and configure the FastAPI application."""
|
||||
app = FastAPI(title="Transcription", lifespan=_lifespan)
|
||||
|
||||
@@ -11,6 +11,7 @@ from enum import StrEnum
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import Field
|
||||
from pydantic_settings import BaseSettings
|
||||
from pydantic_settings import SettingsConfigDict
|
||||
|
||||
@@ -50,6 +51,7 @@ class Settings(BaseSettings):
|
||||
# --- worker reliability ---
|
||||
worker_max_retries: int = 0
|
||||
worker_retry_backoff_seconds: float = 0.0
|
||||
worker_provider_timeout_seconds: float = Field(default=20.0, gt=0.0, le=20.0)
|
||||
|
||||
@property
|
||||
def should_bootstrap_schema(self) -> bool:
|
||||
|
||||
@@ -54,7 +54,7 @@ class Source(SQLModel, table=True):
|
||||
# Relationships
|
||||
document: Optional["Document"] = Relationship(back_populates="sources")
|
||||
job: Optional["Job"] = Relationship(back_populates="sources")
|
||||
revision: "Revision | None" = Relationship(
|
||||
revision: Optional["Revision"] = Relationship(
|
||||
back_populates="source",
|
||||
sa_relationship_kwargs={"uselist": False},
|
||||
)
|
||||
|
||||
@@ -149,8 +149,40 @@ class JobService(ServiceBase):
|
||||
async with self._session_scope(session) as _session:
|
||||
query = (
|
||||
select(Job)
|
||||
.options(selectinload(Job.document)) # pyright: ignore[reportArgumentType]
|
||||
.options(
|
||||
selectinload(Job.document), # pyright: ignore[reportArgumentType]
|
||||
selectinload(Job.sources).selectinload(Source.revision), # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
.where(Job.status == JobStatus.QUEUED)
|
||||
.order_by(Job.date_created) # pyright: ignore[reportArgumentType]
|
||||
)
|
||||
return (await _session.exec(query)).first()
|
||||
|
||||
async def requeue_stale_processing_jobs(
|
||||
self,
|
||||
*,
|
||||
stale_before: datetime,
|
||||
session: AsyncSession | None = None,
|
||||
) -> int:
|
||||
"""Move stale processing jobs back to queued state.
|
||||
|
||||
Jobs with ``status=PROCESSING`` and ``date_updated`` older than
|
||||
``stale_before`` are considered stale and re-queued.
|
||||
"""
|
||||
async with self._session_scope(session) as _session:
|
||||
query = (
|
||||
select(Job)
|
||||
.where(Job.status == JobStatus.PROCESSING)
|
||||
.where(Job.date_updated < stale_before)
|
||||
)
|
||||
stale_jobs = (await _session.exec(query)).all()
|
||||
if not stale_jobs:
|
||||
return 0
|
||||
|
||||
now = datetime.now(UTC)
|
||||
for job in stale_jobs:
|
||||
job.status = JobStatus.QUEUED
|
||||
job.date_updated = now
|
||||
|
||||
await self._finalize(session=_session, caller_session=session, refresh=stale_jobs)
|
||||
return len(stale_jobs)
|
||||
|
||||
@@ -6,6 +6,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
from ..config import Settings
|
||||
from ..config import get_settings
|
||||
from ..errors import AppError
|
||||
from ..errors import ErrorCategory
|
||||
from ..errors import classify_unexpected_error
|
||||
from ..errors import format_error_detail
|
||||
from ..models import Job
|
||||
@@ -29,7 +30,7 @@ async def advance_job(
|
||||
settings = settings or get_settings()
|
||||
match job.status:
|
||||
case JobStatus.QUEUED:
|
||||
return await process_queued_job(job=job, services=services, session=session)
|
||||
return await process_queued_job(job=job, services=services, settings=settings, session=session)
|
||||
case JobStatus.FAILED:
|
||||
if job.retry_count < settings.worker_max_retries:
|
||||
return await services.jobs.update_job_state(
|
||||
@@ -49,9 +50,11 @@ async def process_queued_job(
|
||||
*,
|
||||
job: Job,
|
||||
services: ServiceBundle,
|
||||
settings: Settings | None = None,
|
||||
session: AsyncSession | None = None,
|
||||
) -> Job | None:
|
||||
"""Process one complete transcription attempt for a queued job."""
|
||||
runtime_settings = settings or get_settings()
|
||||
if job.status != JobStatus.QUEUED:
|
||||
logger.warning(f"Job {job.id} is not queued. Current status: {job.status}")
|
||||
return
|
||||
@@ -65,11 +68,15 @@ async def process_queued_job(
|
||||
job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING, session=session)
|
||||
await session.commit()
|
||||
|
||||
source = _resolve_primary_source(job)
|
||||
source_job = await services.jobs.read_job(job_id=job.id, session=session)
|
||||
source = _resolve_primary_source(source_job)
|
||||
assert source is not None, f"Job {job.id} has no associated source record."
|
||||
|
||||
try:
|
||||
result = await transcribe_document_image(source.file_path)
|
||||
result = await asyncio.wait_for(
|
||||
transcribe_document_image(source.file_path),
|
||||
timeout=runtime_settings.worker_provider_timeout_seconds,
|
||||
)
|
||||
job = await _finalize_transcribed(job=job, services=services, result=result, session=session)
|
||||
logger.info(
|
||||
"Job transcribed operation=worker.process_job job_id=%s document_id=%s source_id=%s provider=%s",
|
||||
@@ -78,6 +85,22 @@ async def process_queued_job(
|
||||
source.id,
|
||||
result.provider,
|
||||
)
|
||||
except TimeoutError as exc:
|
||||
error = AppError(
|
||||
f"Provider call timed out after {runtime_settings.worker_provider_timeout_seconds:.1f}s",
|
||||
category=ErrorCategory.EXTERNAL_PROVIDER,
|
||||
suggestion="Retry the job. If this repeats, verify provider latency and request payload size.",
|
||||
retriable=True,
|
||||
)
|
||||
job = await _finalize_failed(job=job, services=services, error=error, session=session)
|
||||
logger.error(
|
||||
"Job failed operation=worker.process_job job_id=%s document_id=%s source_id=%s error_id=%s category=%s",
|
||||
job.id,
|
||||
job.document_id,
|
||||
source.id,
|
||||
error.error_id,
|
||||
error.category.value,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
match exc:
|
||||
case AppError() as error:
|
||||
|
||||
@@ -57,7 +57,18 @@ def register_page() -> None:
|
||||
transcription_service = TranscriptionService(session_factory=session_factory)
|
||||
render_navigation_header(current_path="/jobs")
|
||||
|
||||
job = await jobs_service.read_job(job_id=UUID(job_id))
|
||||
try:
|
||||
parsed_job_id = UUID(job_id)
|
||||
except ValueError:
|
||||
ui.label("Invalid job id").classes("text-h6 text-negative")
|
||||
return
|
||||
|
||||
try:
|
||||
job = await jobs_service.read_job(job_id=parsed_job_id)
|
||||
except ValueError:
|
||||
ui.label("Job not found").classes("text-h6 text-negative")
|
||||
return
|
||||
|
||||
source = _resolve_primary_source(job)
|
||||
|
||||
with ui.splitter(value=30).classes("w-full h-[calc(100vh-64px)]") as splitter:
|
||||
@@ -92,7 +103,7 @@ def register_page() -> None:
|
||||
|
||||
@ui.refreshable
|
||||
async def render_revision_panel() -> None:
|
||||
refreshed_job = await jobs_service.read_job(job_id=UUID(job_id))
|
||||
refreshed_job = await jobs_service.read_job(job_id=parsed_job_id)
|
||||
refreshed_source = _resolve_primary_source(refreshed_job)
|
||||
if refreshed_source is None or refreshed_source.revision is None:
|
||||
ui.label("No revision exists for this source.").classes("text-body2 text-grey-3")
|
||||
|
||||
Reference in New Issue
Block a user