generated from john/python-template
Updated test suite
This commit is contained in:
@@ -39,6 +39,10 @@ dev = [
|
|||||||
|
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
addopts = "--strict-markers -q"
|
addopts = "--strict-markers -q"
|
||||||
|
asyncio_mode = "strict"
|
||||||
|
filterwarnings = [
|
||||||
|
"error:coroutine .* was never awaited:RuntimeWarning",
|
||||||
|
]
|
||||||
markers = [
|
markers = [
|
||||||
"unit: pure logic tests with no external dependencies",
|
"unit: pure logic tests with no external dependencies",
|
||||||
"integration: tests that touch framework or database contracts",
|
"integration: tests that touch framework or database contracts",
|
||||||
|
|||||||
@@ -4,6 +4,10 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from contextlib import AsyncExitStack
|
from contextlib import AsyncExitStack
|
||||||
from contextlib import asynccontextmanager
|
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 FastAPI
|
||||||
from fastapi import status
|
from fastapi import status
|
||||||
@@ -18,10 +22,14 @@ from .db import create_all
|
|||||||
from .db import dispose_database_runtime
|
from .db import dispose_database_runtime
|
||||||
from .db import initialize_database_runtime
|
from .db import initialize_database_runtime
|
||||||
from .services import ServiceBundle
|
from .services import ServiceBundle
|
||||||
|
from .services.jobs import JobService
|
||||||
from .ui import register_pages
|
from .ui import register_pages
|
||||||
from .worker import worker_consumer_lifespan
|
from .worker import worker_consumer_lifespan
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def _lifespan(app: FastAPI):
|
async def _lifespan(app: FastAPI):
|
||||||
configure_logging()
|
configure_logging()
|
||||||
@@ -37,6 +45,8 @@ async def _lifespan(app: FastAPI):
|
|||||||
settings.upload_dir.mkdir(parents=True, exist_ok=True)
|
settings.upload_dir.mkdir(parents=True, exist_ok=True)
|
||||||
settings.prompt_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:
|
async with AsyncExitStack() as stack:
|
||||||
stack.push_async_callback(dispose_database_runtime)
|
stack.push_async_callback(dispose_database_runtime)
|
||||||
stop_event, worker_notifier = await stack.enter_async_context(
|
stop_event, worker_notifier = await stack.enter_async_context(
|
||||||
@@ -50,6 +60,20 @@ async def _lifespan(app: FastAPI):
|
|||||||
yield
|
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:
|
def create_app() -> FastAPI:
|
||||||
"""Create and configure the FastAPI application."""
|
"""Create and configure the FastAPI application."""
|
||||||
app = FastAPI(title="Transcription", lifespan=_lifespan)
|
app = FastAPI(title="Transcription", lifespan=_lifespan)
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from enum import StrEnum
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
|
|
||||||
|
from pydantic import Field
|
||||||
from pydantic_settings import BaseSettings
|
from pydantic_settings import BaseSettings
|
||||||
from pydantic_settings import SettingsConfigDict
|
from pydantic_settings import SettingsConfigDict
|
||||||
|
|
||||||
@@ -50,6 +51,7 @@ class Settings(BaseSettings):
|
|||||||
# --- worker reliability ---
|
# --- worker reliability ---
|
||||||
worker_max_retries: int = 0
|
worker_max_retries: int = 0
|
||||||
worker_retry_backoff_seconds: float = 0.0
|
worker_retry_backoff_seconds: float = 0.0
|
||||||
|
worker_provider_timeout_seconds: float = Field(default=20.0, gt=0.0, le=20.0)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def should_bootstrap_schema(self) -> bool:
|
def should_bootstrap_schema(self) -> bool:
|
||||||
|
|||||||
@@ -54,7 +54,7 @@ class Source(SQLModel, table=True):
|
|||||||
# Relationships
|
# Relationships
|
||||||
document: Optional["Document"] = Relationship(back_populates="sources")
|
document: Optional["Document"] = Relationship(back_populates="sources")
|
||||||
job: Optional["Job"] = Relationship(back_populates="sources")
|
job: Optional["Job"] = Relationship(back_populates="sources")
|
||||||
revision: "Revision | None" = Relationship(
|
revision: Optional["Revision"] = Relationship(
|
||||||
back_populates="source",
|
back_populates="source",
|
||||||
sa_relationship_kwargs={"uselist": False},
|
sa_relationship_kwargs={"uselist": False},
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -149,8 +149,40 @@ class JobService(ServiceBase):
|
|||||||
async with self._session_scope(session) as _session:
|
async with self._session_scope(session) as _session:
|
||||||
query = (
|
query = (
|
||||||
select(Job)
|
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)
|
.where(Job.status == JobStatus.QUEUED)
|
||||||
.order_by(Job.date_created) # pyright: ignore[reportArgumentType]
|
.order_by(Job.date_created) # pyright: ignore[reportArgumentType]
|
||||||
)
|
)
|
||||||
return (await _session.exec(query)).first()
|
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 Settings
|
||||||
from ..config import get_settings
|
from ..config import get_settings
|
||||||
from ..errors import AppError
|
from ..errors import AppError
|
||||||
|
from ..errors import ErrorCategory
|
||||||
from ..errors import classify_unexpected_error
|
from ..errors import classify_unexpected_error
|
||||||
from ..errors import format_error_detail
|
from ..errors import format_error_detail
|
||||||
from ..models import Job
|
from ..models import Job
|
||||||
@@ -29,7 +30,7 @@ async def advance_job(
|
|||||||
settings = settings or get_settings()
|
settings = settings or get_settings()
|
||||||
match job.status:
|
match job.status:
|
||||||
case JobStatus.QUEUED:
|
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:
|
case JobStatus.FAILED:
|
||||||
if job.retry_count < settings.worker_max_retries:
|
if job.retry_count < settings.worker_max_retries:
|
||||||
return await services.jobs.update_job_state(
|
return await services.jobs.update_job_state(
|
||||||
@@ -49,9 +50,11 @@ async def process_queued_job(
|
|||||||
*,
|
*,
|
||||||
job: Job,
|
job: Job,
|
||||||
services: ServiceBundle,
|
services: ServiceBundle,
|
||||||
|
settings: Settings | None = None,
|
||||||
session: AsyncSession | None = None,
|
session: AsyncSession | None = None,
|
||||||
) -> Job | None:
|
) -> Job | None:
|
||||||
"""Process one complete transcription attempt for a queued job."""
|
"""Process one complete transcription attempt for a queued job."""
|
||||||
|
runtime_settings = settings or get_settings()
|
||||||
if job.status != JobStatus.QUEUED:
|
if job.status != JobStatus.QUEUED:
|
||||||
logger.warning(f"Job {job.id} is not queued. Current status: {job.status}")
|
logger.warning(f"Job {job.id} is not queued. Current status: {job.status}")
|
||||||
return
|
return
|
||||||
@@ -65,11 +68,15 @@ async def process_queued_job(
|
|||||||
job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING, session=session)
|
job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING, session=session)
|
||||||
await session.commit()
|
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."
|
assert source is not None, f"Job {job.id} has no associated source record."
|
||||||
|
|
||||||
try:
|
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)
|
job = await _finalize_transcribed(job=job, services=services, result=result, session=session)
|
||||||
logger.info(
|
logger.info(
|
||||||
"Job transcribed operation=worker.process_job job_id=%s document_id=%s source_id=%s provider=%s",
|
"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,
|
source.id,
|
||||||
result.provider,
|
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
|
except Exception as exc: # noqa: BLE001
|
||||||
match exc:
|
match exc:
|
||||||
case AppError() as error:
|
case AppError() as error:
|
||||||
|
|||||||
@@ -57,7 +57,18 @@ def register_page() -> None:
|
|||||||
transcription_service = TranscriptionService(session_factory=session_factory)
|
transcription_service = TranscriptionService(session_factory=session_factory)
|
||||||
render_navigation_header(current_path="/jobs")
|
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)
|
source = _resolve_primary_source(job)
|
||||||
|
|
||||||
with ui.splitter(value=30).classes("w-full h-[calc(100vh-64px)]") as splitter:
|
with ui.splitter(value=30).classes("w-full h-[calc(100vh-64px)]") as splitter:
|
||||||
@@ -92,7 +103,7 @@ def register_page() -> None:
|
|||||||
|
|
||||||
@ui.refreshable
|
@ui.refreshable
|
||||||
async def render_revision_panel() -> None:
|
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)
|
refreshed_source = _resolve_primary_source(refreshed_job)
|
||||||
if refreshed_source is None or refreshed_source.revision is None:
|
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")
|
ui.label("No revision exists for this source.").classes("text-body2 text-grey-3")
|
||||||
|
|||||||
@@ -33,7 +33,25 @@ class TestPipelineSuccessFlow:
|
|||||||
_ = (prompt_text, image_bytes, mime_type)
|
_ = (prompt_text, image_bytes, mime_type)
|
||||||
return TranscriptionResult(text="Pipeline transcript", provider="openrouter", model="test-model", prompt_name="transcribe_document.md")
|
return TranscriptionResult(text="Pipeline transcript", provider="openrouter", model="test-model", prompt_name="transcribe_document.md")
|
||||||
|
|
||||||
monkeypatch.setattr("transcription.services.transcription.OpenRouterTranscriptionProvider.transcribe", _fake_transcribe)
|
async def _fake_transcribe_document_image(
|
||||||
|
image_path,
|
||||||
|
*,
|
||||||
|
prompt_name="transcribe_document.md",
|
||||||
|
settings=None,
|
||||||
|
provider=None,
|
||||||
|
) -> TranscriptionResult:
|
||||||
|
_ = (image_path, prompt_name, settings, provider)
|
||||||
|
return TranscriptionResult(
|
||||||
|
text="Pipeline transcript",
|
||||||
|
provider="openrouter",
|
||||||
|
model="test-model",
|
||||||
|
prompt_name="transcribe_document.md",
|
||||||
|
)
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"transcription.services.workflows.transcribe_document_image",
|
||||||
|
_fake_transcribe_document_image,
|
||||||
|
)
|
||||||
|
|
||||||
processed = await process_next_queued_job(session=async_session)
|
processed = await process_next_queued_job(session=async_session)
|
||||||
job = await async_session.get(Job, upload_result.job_id)
|
job = await async_session.get(Job, upload_result.job_id)
|
||||||
@@ -60,11 +78,20 @@ class TestPipelineFailureFlow:
|
|||||||
settings=settings,
|
settings=settings,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _fake_transcribe(*, prompt_text: str, image_bytes: bytes, mime_type: str) -> TranscriptionResult:
|
async def _fake_transcribe_document_image(
|
||||||
_ = (prompt_text, image_bytes, mime_type)
|
image_path,
|
||||||
|
*,
|
||||||
|
prompt_name="transcribe_document.md",
|
||||||
|
settings=None,
|
||||||
|
provider=None,
|
||||||
|
) -> TranscriptionResult:
|
||||||
|
_ = (image_path, prompt_name, settings, provider)
|
||||||
raise RuntimeError("pipeline provider failure")
|
raise RuntimeError("pipeline provider failure")
|
||||||
|
|
||||||
monkeypatch.setattr("transcription.services.transcription.OpenRouterTranscriptionProvider.transcribe", _fake_transcribe)
|
monkeypatch.setattr(
|
||||||
|
"transcription.services.workflows.transcribe_document_image",
|
||||||
|
_fake_transcribe_document_image,
|
||||||
|
)
|
||||||
|
|
||||||
processed = await process_next_queued_job(session=async_session)
|
processed = await process_next_queued_job(session=async_session)
|
||||||
job = await async_session.get(Job, upload_result.job_id)
|
job = await async_session.get(Job, upload_result.job_id)
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ class _FakeChat:
|
|||||||
self._error = error
|
self._error = error
|
||||||
self.calls = []
|
self.calls = []
|
||||||
|
|
||||||
def send(self, **kwargs):
|
async def send_async(self, **kwargs):
|
||||||
self.calls.append(kwargs)
|
self.calls.append(kwargs)
|
||||||
if self._error:
|
if self._error:
|
||||||
raise self._error
|
raise self._error
|
||||||
@@ -48,7 +48,8 @@ class TestOpenRouterProviderInit:
|
|||||||
class TestOpenRouterProviderTranscribe:
|
class TestOpenRouterProviderTranscribe:
|
||||||
"""Verify OpenRouter request construction and response parsing."""
|
"""Verify OpenRouter request construction and response parsing."""
|
||||||
|
|
||||||
def test_includes_optional_referer_and_title_when_set(self):
|
@pytest.mark.asyncio
|
||||||
|
async def test_includes_optional_referer_and_title_when_set(self):
|
||||||
"""Transcribe sends app attribution fields when configured."""
|
"""Transcribe sends app attribution fields when configured."""
|
||||||
response = {"model": "vendor/model-a", "choices": [{"message": {"content": "Transcript text"}}]}
|
response = {"model": "vendor/model-a", "choices": [{"message": {"content": "Transcript text"}}]}
|
||||||
client = _FakeClient(response=response)
|
client = _FakeClient(response=response)
|
||||||
@@ -59,7 +60,7 @@ class TestOpenRouterProviderTranscribe:
|
|||||||
)
|
)
|
||||||
provider = OpenRouterTranscriptionProvider(settings=settings, client=client)
|
provider = OpenRouterTranscriptionProvider(settings=settings, client=client)
|
||||||
|
|
||||||
result = provider.transcribe(
|
result = await provider.transcribe(
|
||||||
prompt_text="Prompt body",
|
prompt_text="Prompt body",
|
||||||
image_bytes=b"img-bytes",
|
image_bytes=b"img-bytes",
|
||||||
mime_type="image/png",
|
mime_type="image/png",
|
||||||
@@ -70,7 +71,8 @@ class TestOpenRouterProviderTranscribe:
|
|||||||
assert send_call["x_open_router_title"] == "Transcription App"
|
assert send_call["x_open_router_title"] == "Transcription App"
|
||||||
assert result.text == "Transcript text"
|
assert result.text == "Transcript text"
|
||||||
|
|
||||||
def test_parses_successful_response_text(self):
|
@pytest.mark.asyncio
|
||||||
|
async def test_parses_successful_response_text(self):
|
||||||
"""Transcribe returns normalized text from a valid response payload."""
|
"""Transcribe returns normalized text from a valid response payload."""
|
||||||
response = {
|
response = {
|
||||||
"model": "vendor/model-b",
|
"model": "vendor/model-b",
|
||||||
@@ -81,7 +83,7 @@ class TestOpenRouterProviderTranscribe:
|
|||||||
client=_FakeClient(response=response),
|
client=_FakeClient(response=response),
|
||||||
)
|
)
|
||||||
|
|
||||||
result = provider.transcribe(
|
result = await provider.transcribe(
|
||||||
prompt_text="Prompt body",
|
prompt_text="Prompt body",
|
||||||
image_bytes=b"img-bytes",
|
image_bytes=b"img-bytes",
|
||||||
mime_type="image/jpeg",
|
mime_type="image/jpeg",
|
||||||
@@ -91,7 +93,8 @@ class TestOpenRouterProviderTranscribe:
|
|||||||
assert result.provider == "openrouter"
|
assert result.provider == "openrouter"
|
||||||
assert result.model == "vendor/model-b"
|
assert result.model == "vendor/model-b"
|
||||||
|
|
||||||
def test_maps_sdk_exception_to_provider_error(self):
|
@pytest.mark.asyncio
|
||||||
|
async def test_maps_sdk_exception_to_provider_error(self):
|
||||||
"""Transcribe converts SDK failures to ProviderError."""
|
"""Transcribe converts SDK failures to ProviderError."""
|
||||||
provider = OpenRouterTranscriptionProvider(
|
provider = OpenRouterTranscriptionProvider(
|
||||||
settings=Settings(openrouter_api_key="test-key"),
|
settings=Settings(openrouter_api_key="test-key"),
|
||||||
@@ -99,13 +102,14 @@ class TestOpenRouterProviderTranscribe:
|
|||||||
)
|
)
|
||||||
|
|
||||||
with pytest.raises(ProviderError):
|
with pytest.raises(ProviderError):
|
||||||
provider.transcribe(
|
await provider.transcribe(
|
||||||
prompt_text="Prompt body",
|
prompt_text="Prompt body",
|
||||||
image_bytes=b"img-bytes",
|
image_bytes=b"img-bytes",
|
||||||
mime_type="image/png",
|
mime_type="image/png",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_raises_on_empty_or_invalid_response(self):
|
@pytest.mark.asyncio
|
||||||
|
async def test_raises_on_empty_or_invalid_response(self):
|
||||||
"""Transcribe raises ProviderResponseError for missing completion text."""
|
"""Transcribe raises ProviderResponseError for missing completion text."""
|
||||||
provider = OpenRouterTranscriptionProvider(
|
provider = OpenRouterTranscriptionProvider(
|
||||||
settings=Settings(openrouter_api_key="test-key"),
|
settings=Settings(openrouter_api_key="test-key"),
|
||||||
@@ -113,7 +117,7 @@ class TestOpenRouterProviderTranscribe:
|
|||||||
)
|
)
|
||||||
|
|
||||||
with pytest.raises(ProviderResponseError):
|
with pytest.raises(ProviderResponseError):
|
||||||
provider.transcribe(
|
await provider.transcribe(
|
||||||
prompt_text="Prompt body",
|
prompt_text="Prompt body",
|
||||||
image_bytes=b"img-bytes",
|
image_bytes=b"img-bytes",
|
||||||
mime_type="image/png",
|
mime_type="image/png",
|
||||||
|
|||||||
@@ -4,85 +4,93 @@ import pytest
|
|||||||
|
|
||||||
from transcription.models import Document
|
from transcription.models import Document
|
||||||
from transcription.models import Job
|
from transcription.models import Job
|
||||||
|
from transcription.models import JobStatus
|
||||||
|
from transcription.models import Source
|
||||||
from transcription.services.documents import DocumentService
|
from transcription.services.documents import DocumentService
|
||||||
from transcription.services.jobs import JobService
|
from transcription.services.jobs import JobService
|
||||||
from transcription.services.jobs import JobStatus
|
|
||||||
|
|
||||||
|
|
||||||
class TestJobService:
|
class TestJobService:
|
||||||
class TestBasicCRUD:
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_create_job(self, job_service: JobService):
|
async def test_create_and_read_job(self, job_service: JobService, document_service: DocumentService):
|
||||||
"""Test creating a job."""
|
document = Document(id=uuid4(), name="test-bundle")
|
||||||
|
|
||||||
def fake_job_factory():
|
|
||||||
return Job(document_id=uuid4())
|
|
||||||
|
|
||||||
await job_service.create_job(job=fake_job_factory())
|
|
||||||
|
|
||||||
async with job_service._session_scope() as session:
|
|
||||||
for _ in range(10):
|
|
||||||
await job_service.create_job(job=fake_job_factory(), session=session)
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_backpropagation(self, job_service: JobService, document_service: DocumentService):
|
|
||||||
"""Test that creating a job backpropagates to the related document."""
|
|
||||||
doc_id = uuid4()
|
|
||||||
document = Document(
|
|
||||||
id=doc_id,
|
|
||||||
filename="test.txt",
|
|
||||||
file_path="/path/to/test.txt",
|
|
||||||
)
|
|
||||||
await document_service.create_document(document=document)
|
await document_service.create_document(document=document)
|
||||||
job = Job(document_id=doc_id)
|
|
||||||
|
job = Job(document_id=document.id)
|
||||||
await job_service.create_job(job=job)
|
await job_service.create_job(job=job)
|
||||||
|
|
||||||
read_job = await job_service.read_job(job_id=job.id)
|
fetched = await job_service.read_job(job_id=job.id)
|
||||||
assert isinstance(read_job.document, Document)
|
assert fetched.id == job.id
|
||||||
assert read_job.document.id == document.id
|
assert fetched.document is not None
|
||||||
|
assert fetched.document.id == document.id
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_reading_job(self, job_service: JobService):
|
async def test_update_job_state_updates_status_and_retry(self, job_service: JobService, document_service: DocumentService):
|
||||||
"""Test reading a job."""
|
document = Document(id=uuid4(), name="test-bundle")
|
||||||
uuid = uuid4()
|
await document_service.create_document(document=document)
|
||||||
await job_service.create_job(job=Job(id=uuid, document_id=uuid4()))
|
|
||||||
job = await job_service.read_job(job_id=uuid)
|
job = Job(document_id=document.id)
|
||||||
assert job.id == uuid
|
await job_service.create_job(job=job)
|
||||||
|
|
||||||
|
updated = await job_service.update_job_state(
|
||||||
|
job_id=job.id,
|
||||||
|
status=JobStatus.PROCESSING,
|
||||||
|
retry_count_increment=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert updated.status == JobStatus.PROCESSING
|
||||||
|
assert updated.retry_count == 1
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_updating_job(self, job_service: JobService):
|
async def test_query_jobs_by_status(self, job_service: JobService, document_service: DocumentService):
|
||||||
"""Test updating a job."""
|
document = Document(id=uuid4(), name="query-doc")
|
||||||
uuid = uuid4()
|
await document_service.create_document(document=document)
|
||||||
job = Job(id=uuid, document_id=uuid4())
|
|
||||||
async with job_service._session_scope() as session:
|
|
||||||
await job_service.create_job(job=job, session=session)
|
|
||||||
job.status = JobStatus.PROCESSING
|
|
||||||
await job_service.update_job(job=job, session=session)
|
|
||||||
read_job = await job_service.read_job(job_id=uuid, session=session)
|
|
||||||
assert read_job == job
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
await job_service.create_job(job=Job(document_id=document.id, status=JobStatus.PROCESSING))
|
||||||
async def test_deleting_job(self, job_service: JobService):
|
await job_service.create_job(job=Job(document_id=document.id, status=JobStatus.QUEUED))
|
||||||
"""Test deleting a job."""
|
|
||||||
|
|
||||||
class TestServiceMethods:
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_query_jobs(self, job_service: JobService):
|
|
||||||
"""Test querying jobs."""
|
|
||||||
await job_service.create_job(job=Job(document_id=uuid4(), status=JobStatus.PROCESSING))
|
|
||||||
result = await job_service.query_jobs(status=JobStatus.PROCESSING)
|
result = await job_service.query_jobs(status=JobStatus.PROCESSING)
|
||||||
jobs = {str(job.id).split("-")[0]: job.status for job in result}
|
assert len(result) == 1
|
||||||
assert len(jobs) == 1
|
assert result[0].status == JobStatus.PROCESSING
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_list_jobs(self, job_service: JobService):
|
async def test_query_jobs_by_source_filename(self, job_service: JobService, document_service: DocumentService):
|
||||||
"""Test listing jobs."""
|
document = Document(id=uuid4(), name="source-doc")
|
||||||
n = 5
|
await document_service.create_document(document=document)
|
||||||
for _ in range(n):
|
|
||||||
await job_service.create_job(job=Job(document_id=uuid4()))
|
job = Job(document_id=document.id)
|
||||||
jobs = await job_service.list_jobs()
|
await job_service.create_job(job=job)
|
||||||
assert len(jobs) == n
|
|
||||||
|
async with job_service._session_scope() as session:
|
||||||
|
session.add(
|
||||||
|
Source(
|
||||||
|
document_id=document.id,
|
||||||
|
job_id=job.id,
|
||||||
|
upload_name="letter.jpg",
|
||||||
|
filename="stored-letter.jpg",
|
||||||
|
file_path="/uploads/stored-letter.jpg",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
result = await job_service.query_jobs(filename="stored-letter.jpg")
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0].id == job.id
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_mark_job_status(self, job_service: JobService):
|
async def test_read_next_queued_job_orders_by_created_date(
|
||||||
"""Test marking a job with a new status."""
|
self,
|
||||||
|
job_service: JobService,
|
||||||
|
document_service: DocumentService,
|
||||||
|
):
|
||||||
|
document = Document(id=uuid4(), name="ordered-doc")
|
||||||
|
await document_service.create_document(document=document)
|
||||||
|
|
||||||
|
first = Job(document_id=document.id, status=JobStatus.QUEUED)
|
||||||
|
second = Job(document_id=document.id, status=JobStatus.QUEUED)
|
||||||
|
await job_service.create_job(job=first)
|
||||||
|
await job_service.create_job(job=second)
|
||||||
|
|
||||||
|
next_job = await job_service.read_next_queued_job()
|
||||||
|
assert next_job is not None
|
||||||
|
assert next_job.id == first.id
|
||||||
|
|||||||
@@ -49,10 +49,11 @@ class TestRealImageExternalTranscription:
|
|||||||
assert REAL_IMAGES_DIR.exists()
|
assert REAL_IMAGES_DIR.exists()
|
||||||
assert _real_image_paths()
|
assert _real_image_paths()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.parametrize("image_path", _real_image_paths(), ids=lambda p: p.name)
|
@pytest.mark.parametrize("image_path", _real_image_paths(), ids=lambda p: p.name)
|
||||||
def test_transcribes_real_image_fixture(self, image_path: Path):
|
async def test_transcribes_real_image_fixture(self, image_path: Path):
|
||||||
"""Real fixture image produces a non-empty transcription result."""
|
"""Real fixture image produces a non-empty transcription result."""
|
||||||
result = transcribe_document_image(image_path)
|
result = await transcribe_document_image(image_path)
|
||||||
assert result.provider == "openrouter"
|
assert result.provider == "openrouter"
|
||||||
assert isinstance(result.model, str) and result.model.strip()
|
assert isinstance(result.model, str) and result.model.strip()
|
||||||
assert isinstance(result.text, str) and result.text.strip()
|
assert isinstance(result.text, str) and result.text.strip()
|
||||||
|
|||||||
@@ -0,0 +1,63 @@
|
|||||||
|
"""Reliability tests for worker workflow timeout behavior."""
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from transcription.config import Settings
|
||||||
|
from transcription.models import Document
|
||||||
|
from transcription.models import Job
|
||||||
|
from transcription.models import JobStatus
|
||||||
|
from transcription.models import Source
|
||||||
|
from transcription.services import ServiceBundle
|
||||||
|
from transcription.services.workflows import process_queued_job
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
class TestWorkflowReliability:
|
||||||
|
"""Verify timeout and terminal-state reliability behavior."""
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_process_queued_job_timeout_marks_job_failed(self, default_session_factory, monkeypatch):
|
||||||
|
"""Provider timeout transitions a queued job to failed with error detail."""
|
||||||
|
services = ServiceBundle()
|
||||||
|
object.__setattr__(services, "documents", services.documents.__class__(session_factory=default_session_factory))
|
||||||
|
object.__setattr__(services, "jobs", services.jobs.__class__(session_factory=default_session_factory))
|
||||||
|
object.__setattr__(services, "transcriptions", services.transcriptions.__class__(session_factory=default_session_factory))
|
||||||
|
|
||||||
|
async with services.jobs._session_scope() as session:
|
||||||
|
document = Document(id=uuid4(), name="timeout-doc")
|
||||||
|
session.add(document)
|
||||||
|
await session.flush()
|
||||||
|
|
||||||
|
job = Job(document_id=document.id, status=JobStatus.QUEUED)
|
||||||
|
session.add(job)
|
||||||
|
await session.flush()
|
||||||
|
|
||||||
|
source = Source(
|
||||||
|
document_id=document.id,
|
||||||
|
job_id=job.id,
|
||||||
|
upload_name="timeout.jpg",
|
||||||
|
filename="timeout.jpg",
|
||||||
|
file_path=str(Path("tests/fixtures/images/real/Book Two - page 02.jpg")),
|
||||||
|
)
|
||||||
|
session.add(source)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
loaded = await services.jobs.read_job(job_id=job.id, session=session)
|
||||||
|
|
||||||
|
async def _never_returns(image_path, *, prompt_name="transcribe_document.md", settings=None, provider=None):
|
||||||
|
_ = (image_path, prompt_name, settings, provider)
|
||||||
|
raise TimeoutError("simulated provider timeout")
|
||||||
|
|
||||||
|
monkeypatch.setattr("transcription.services.workflows.transcribe_document_image", _never_returns)
|
||||||
|
|
||||||
|
timeout_settings = Settings(openrouter_api_key="test-key", worker_provider_timeout_seconds=20.0)
|
||||||
|
result = await process_queued_job(job=loaded, services=services, settings=timeout_settings)
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert result.status == JobStatus.FAILED
|
||||||
|
assert result.error_detail is not None
|
||||||
|
assert "timed out" in result.error_detail.lower()
|
||||||
|
assert "20.0s" in result.error_detail
|
||||||
+64
-29
@@ -1,5 +1,7 @@
|
|||||||
"""Tests for transcription.app."""
|
"""Tests for transcription.app."""
|
||||||
|
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
@@ -21,28 +23,43 @@ class TestAppFactory:
|
|||||||
class TestAppLifespan:
|
class TestAppLifespan:
|
||||||
"""Verify startup and shutdown lifecycle behavior."""
|
"""Verify startup and shutdown lifecycle behavior."""
|
||||||
|
|
||||||
def test_startup_initializes_runtime_dependencies(self, monkeypatch):
|
def test_startup_initializes_runtime_dependencies(self, monkeypatch, tmp_path):
|
||||||
"""Startup initializes logging, schema, directories, and worker resources."""
|
"""Startup initializes logging, schema, directories, and worker resources."""
|
||||||
calls = []
|
calls = []
|
||||||
|
|
||||||
monkeypatch.setattr("transcription.app.setup_logging", lambda: calls.append("logging"))
|
monkeypatch.setattr("transcription.app.configure_logging", lambda: calls.append("logging"))
|
||||||
monkeypatch.setattr("transcription.app.create_all", lambda **_kwargs: calls.append("schema"))
|
|
||||||
|
async def _create_all(**_kwargs):
|
||||||
|
calls.append("schema")
|
||||||
|
|
||||||
|
monkeypatch.setattr("transcription.app.create_all", _create_all)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"transcription.app.initialize_database_runtime",
|
"transcription.app.initialize_database_runtime",
|
||||||
lambda **_kwargs: type("_Runtime", (), {"engine": object()})(),
|
lambda **_kwargs: type("_Runtime", (), {"engine": object(), "session_factory": object()})(),
|
||||||
)
|
)
|
||||||
monkeypatch.setattr("transcription.app.dispose_database_runtime", lambda: calls.append("dispose_db"))
|
|
||||||
monkeypatch.setattr("transcription.app.should_bootstrap_schema", lambda _settings: True)
|
|
||||||
monkeypatch.setattr("transcription.app._start_worker", lambda _app: calls.append("start_worker"))
|
|
||||||
monkeypatch.setattr("transcription.app._stop_worker", lambda _app: calls.append("stop_worker"))
|
|
||||||
|
|
||||||
class _Dir:
|
async def _dispose_runtime():
|
||||||
def mkdir(self, parents: bool, exist_ok: bool):
|
calls.append("dispose_db")
|
||||||
calls.append("mkdir")
|
|
||||||
|
monkeypatch.setattr("transcription.app.dispose_database_runtime", _dispose_runtime)
|
||||||
|
|
||||||
|
async def _recover_stale(_app):
|
||||||
|
calls.append("recover")
|
||||||
|
|
||||||
|
monkeypatch.setattr("transcription.app._recover_stale_processing_jobs", _recover_stale)
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def _worker_lifespan(**_kwargs):
|
||||||
|
calls.append("worker_start")
|
||||||
|
yield object(), object()
|
||||||
|
calls.append("worker_stop")
|
||||||
|
|
||||||
|
monkeypatch.setattr("transcription.app.worker_consumer_lifespan", _worker_lifespan)
|
||||||
|
|
||||||
class _Settings:
|
class _Settings:
|
||||||
upload_dir = _Dir()
|
should_bootstrap_schema = True
|
||||||
prompt_dir = _Dir()
|
upload_dir = tmp_path / "uploads"
|
||||||
|
prompt_dir = tmp_path / "prompts"
|
||||||
|
|
||||||
monkeypatch.setattr("transcription.app.get_settings", lambda: _Settings())
|
monkeypatch.setattr("transcription.app.get_settings", lambda: _Settings())
|
||||||
|
|
||||||
@@ -52,32 +69,50 @@ class TestAppLifespan:
|
|||||||
|
|
||||||
assert "logging" in calls
|
assert "logging" in calls
|
||||||
assert "schema" in calls
|
assert "schema" in calls
|
||||||
assert "mkdir" in calls
|
assert "recover" in calls
|
||||||
assert "start_worker" in calls
|
assert "worker_start" in calls
|
||||||
|
assert "worker_stop" in calls
|
||||||
assert "dispose_db" in calls
|
assert "dispose_db" in calls
|
||||||
|
assert _Settings.upload_dir.exists()
|
||||||
|
assert _Settings.prompt_dir.exists()
|
||||||
|
|
||||||
def test_shutdown_stops_worker_resources(self, monkeypatch):
|
def test_shutdown_stops_worker_resources(self, monkeypatch, tmp_path):
|
||||||
"""Shutdown signals and stops worker resources cleanly."""
|
"""Shutdown signals and stops worker resources cleanly."""
|
||||||
calls = []
|
calls = []
|
||||||
|
|
||||||
monkeypatch.setattr("transcription.app.setup_logging", lambda: None)
|
monkeypatch.setattr("transcription.app.configure_logging", lambda: calls.append("logging"))
|
||||||
monkeypatch.setattr("transcription.app.create_all", lambda **_kwargs: None)
|
|
||||||
|
async def _create_all(**_kwargs):
|
||||||
|
calls.append("schema")
|
||||||
|
|
||||||
|
monkeypatch.setattr("transcription.app.create_all", _create_all)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"transcription.app.initialize_database_runtime",
|
"transcription.app.initialize_database_runtime",
|
||||||
lambda **_kwargs: type("_Runtime", (), {"engine": object()})(),
|
lambda **_kwargs: type("_Runtime", (), {"engine": object(), "session_factory": object()})(),
|
||||||
)
|
)
|
||||||
monkeypatch.setattr("transcription.app.dispose_database_runtime", lambda: calls.append("dispose_db"))
|
|
||||||
monkeypatch.setattr("transcription.app.should_bootstrap_schema", lambda _settings: True)
|
|
||||||
monkeypatch.setattr("transcription.app._start_worker", lambda _app: calls.append("start_worker"))
|
|
||||||
monkeypatch.setattr("transcription.app._stop_worker", lambda _app: calls.append("stop_worker"))
|
|
||||||
|
|
||||||
class _Dir:
|
async def _dispose_runtime():
|
||||||
def mkdir(self, parents: bool, exist_ok: bool):
|
calls.append("dispose_db")
|
||||||
return None
|
|
||||||
|
monkeypatch.setattr("transcription.app.dispose_database_runtime", _dispose_runtime)
|
||||||
|
|
||||||
|
async def _recover_stale(_app):
|
||||||
|
calls.append("recover")
|
||||||
|
|
||||||
|
monkeypatch.setattr("transcription.app._recover_stale_processing_jobs", _recover_stale)
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def _worker_lifespan(**_kwargs):
|
||||||
|
calls.append("worker_start")
|
||||||
|
yield object(), object()
|
||||||
|
calls.append("worker_stop")
|
||||||
|
|
||||||
|
monkeypatch.setattr("transcription.app.worker_consumer_lifespan", _worker_lifespan)
|
||||||
|
|
||||||
class _Settings:
|
class _Settings:
|
||||||
upload_dir = _Dir()
|
should_bootstrap_schema = True
|
||||||
prompt_dir = _Dir()
|
upload_dir = tmp_path / "uploads"
|
||||||
|
prompt_dir = tmp_path / "prompts"
|
||||||
|
|
||||||
monkeypatch.setattr("transcription.app.get_settings", lambda: _Settings())
|
monkeypatch.setattr("transcription.app.get_settings", lambda: _Settings())
|
||||||
|
|
||||||
@@ -85,4 +120,4 @@ class TestAppLifespan:
|
|||||||
with TestClient(app):
|
with TestClient(app):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
assert calls == ["start_worker", "stop_worker", "dispose_db"]
|
assert calls == ["logging", "schema", "recover", "worker_start", "worker_stop", "dispose_db"]
|
||||||
|
|||||||
+41
-71
@@ -1,97 +1,67 @@
|
|||||||
"""Tests for transcription.db — schema bootstrap and session factory."""
|
"""Tests for transcription.db runtime and schema bootstrap behavior."""
|
||||||
|
|
||||||
from sqlalchemy import inspect, text
|
import pytest
|
||||||
from sqlmodel import Session, SQLModel, create_engine
|
from sqlalchemy import inspect
|
||||||
from sqlmodel.pool import StaticPool
|
|
||||||
|
from transcription.config import Settings
|
||||||
|
from transcription.db import create_all
|
||||||
|
from transcription.db import dispose_database_runtime
|
||||||
|
from transcription.db import get_session
|
||||||
|
from transcription.db import initialize_database_runtime
|
||||||
|
|
||||||
|
|
||||||
def _in_memory_engine():
|
@pytest.mark.asyncio
|
||||||
"""Create a fresh in-memory SQLite engine for isolated db tests."""
|
async def test_create_all_creates_expected_tables(tmp_path):
|
||||||
return create_engine(
|
settings = Settings(
|
||||||
"sqlite://",
|
openrouter_api_key="test-key",
|
||||||
connect_args={"check_same_thread": False},
|
database_url=f"sqlite:///{tmp_path / 'schema.db'}",
|
||||||
poolclass=StaticPool,
|
environment="test",
|
||||||
)
|
)
|
||||||
|
runtime = initialize_database_runtime(settings=settings)
|
||||||
|
|
||||||
|
try:
|
||||||
|
await create_all(engine=runtime.engine)
|
||||||
|
async with runtime.engine.connect() as conn:
|
||||||
|
table_names = set(await conn.run_sync(lambda c: inspect(c).get_table_names()))
|
||||||
|
|
||||||
class TestSchemaBootstrap:
|
|
||||||
"""Verify create_all produces the expected table set."""
|
|
||||||
|
|
||||||
def test_create_all_creates_expected_tables(self):
|
|
||||||
"""After create_all(), document, source, job, and revision tables exist."""
|
|
||||||
engine = _in_memory_engine()
|
|
||||||
# Ensure models are imported so metadata is populated
|
|
||||||
from transcription.models import Document, Job, Revision, Source # noqa: F401
|
|
||||||
|
|
||||||
import transcription.db as db_module
|
|
||||||
|
|
||||||
db_module.create_all(engine=engine)
|
|
||||||
|
|
||||||
inspector = inspect(engine)
|
|
||||||
table_names = set(inspector.get_table_names())
|
|
||||||
assert "document" in table_names
|
assert "document" in table_names
|
||||||
assert "job" in table_names
|
assert "job" in table_names
|
||||||
assert "source" in table_names
|
assert "source" in table_names
|
||||||
assert "revision" in table_names
|
assert "revision" in table_names
|
||||||
|
finally:
|
||||||
|
await dispose_database_runtime()
|
||||||
|
|
||||||
|
|
||||||
class TestSessionFactory:
|
@pytest.mark.asyncio
|
||||||
"""Verify get_session yields and cleans up sessions."""
|
async def test_get_session_yields_async_session(tmp_path):
|
||||||
|
settings = Settings(
|
||||||
|
openrouter_api_key="test-key",
|
||||||
|
database_url=f"sqlite:///{tmp_path / 'session.db'}",
|
||||||
|
environment="test",
|
||||||
|
)
|
||||||
|
initialize_database_runtime(settings=settings)
|
||||||
|
|
||||||
def test_get_session_yields_session(self):
|
try:
|
||||||
"""get_session() yields a usable Session object."""
|
async with get_session(settings=settings) as session:
|
||||||
engine = _in_memory_engine()
|
assert session is not None
|
||||||
SQLModel.metadata.create_all(engine)
|
finally:
|
||||||
|
await dispose_database_runtime()
|
||||||
import transcription.db as db_module
|
|
||||||
|
|
||||||
with db_module.get_session(engine=engine) as session:
|
|
||||||
assert isinstance(session, Session)
|
|
||||||
|
|
||||||
def test_session_is_closed_after_generator_exit(self):
|
|
||||||
"""After the context manager exits, the session is closed."""
|
|
||||||
engine = _in_memory_engine()
|
|
||||||
SQLModel.metadata.create_all(engine)
|
|
||||||
|
|
||||||
import transcription.db as db_module
|
|
||||||
|
|
||||||
with db_module.get_session(engine=engine) as session:
|
|
||||||
# Session is usable inside the context
|
|
||||||
session.execute(text("SELECT 1"))
|
|
||||||
captured = session
|
|
||||||
|
|
||||||
# After exiting, the session's internal connection is released
|
|
||||||
# (no active transaction bound to the session)
|
|
||||||
assert captured._transaction is None
|
|
||||||
|
|
||||||
|
|
||||||
class TestBootstrapPolicy:
|
def test_bootstrap_policy_production_defaults_false():
|
||||||
"""Verify schema bootstrap policy defaults and overrides."""
|
|
||||||
|
|
||||||
def test_production_defaults_to_no_bootstrap(self):
|
|
||||||
"""Production defaults to explicit non-bootstrap startup behavior."""
|
|
||||||
from transcription.config import Settings
|
|
||||||
from transcription.db import should_bootstrap_schema
|
|
||||||
|
|
||||||
settings = Settings(openrouter_api_key="test-key", environment="production")
|
settings = Settings(openrouter_api_key="test-key", environment="production")
|
||||||
assert should_bootstrap_schema(settings) is False
|
assert settings.should_bootstrap_schema is False
|
||||||
|
|
||||||
def test_development_defaults_to_bootstrap(self):
|
|
||||||
"""Development defaults to schema bootstrap for local workflows."""
|
|
||||||
from transcription.config import Settings
|
|
||||||
from transcription.db import should_bootstrap_schema
|
|
||||||
|
|
||||||
|
def test_bootstrap_policy_development_defaults_true():
|
||||||
settings = Settings(openrouter_api_key="test-key", environment="development")
|
settings = Settings(openrouter_api_key="test-key", environment="development")
|
||||||
assert should_bootstrap_schema(settings) is True
|
assert settings.should_bootstrap_schema is True
|
||||||
|
|
||||||
def test_explicit_override_wins(self):
|
|
||||||
"""Explicit bootstrap_schema_on_startup overrides environment default."""
|
|
||||||
from transcription.config import Settings
|
|
||||||
from transcription.db import should_bootstrap_schema
|
|
||||||
|
|
||||||
|
def test_bootstrap_policy_explicit_override_true():
|
||||||
settings = Settings(
|
settings = Settings(
|
||||||
openrouter_api_key="test-key",
|
openrouter_api_key="test-key",
|
||||||
environment="production",
|
environment="production",
|
||||||
bootstrap_schema_on_startup=True,
|
bootstrap_schema_on_startup=True,
|
||||||
)
|
)
|
||||||
assert should_bootstrap_schema(settings) is True
|
assert settings.should_bootstrap_schema is True
|
||||||
|
|||||||
@@ -9,19 +9,19 @@ MVP_REQUIREMENT_TEST_MAP: dict[str, list[str]] = {
|
|||||||
"tests/integration/test_pipeline_flow.py",
|
"tests/integration/test_pipeline_flow.py",
|
||||||
],
|
],
|
||||||
"REQ-1": [
|
"REQ-1": [
|
||||||
"tests/services/test_upload.py",
|
"tests/integration/test_pipeline_flow.py",
|
||||||
"tests/ui/test_upload_page.py",
|
"tests/ui/test_upload_page.py",
|
||||||
],
|
],
|
||||||
"REQ-2": [
|
"REQ-2": [
|
||||||
"tests/services/test_worker.py",
|
"tests/services/test_workflows_reliability.py",
|
||||||
"tests/integration/test_pipeline_flow.py",
|
"tests/integration/test_pipeline_flow.py",
|
||||||
],
|
],
|
||||||
"REQ-3": [
|
"REQ-3": [
|
||||||
"tests/services/test_worker.py",
|
"tests/services/test_job_service.py",
|
||||||
"tests/ui/test_jobs_page.py",
|
"tests/ui/test_jobs_page.py",
|
||||||
],
|
],
|
||||||
"REQ-4": [
|
"REQ-4": [
|
||||||
"tests/services/test_worker.py",
|
"tests/services/test_workflows_reliability.py",
|
||||||
"tests/integration/test_pipeline_flow.py",
|
"tests/integration/test_pipeline_flow.py",
|
||||||
],
|
],
|
||||||
"REQ-5": [
|
"REQ-5": [
|
||||||
@@ -30,7 +30,7 @@ MVP_REQUIREMENT_TEST_MAP: dict[str, list[str]] = {
|
|||||||
],
|
],
|
||||||
"REQ-6": [
|
"REQ-6": [
|
||||||
"tests/test_app.py",
|
"tests/test_app.py",
|
||||||
"tests/services/test_worker.py",
|
"tests/services/test_workflows_reliability.py",
|
||||||
],
|
],
|
||||||
"REQ-8": [
|
"REQ-8": [
|
||||||
"tests/test_app.py",
|
"tests/test_app.py",
|
||||||
@@ -38,7 +38,7 @@ MVP_REQUIREMENT_TEST_MAP: dict[str, list[str]] = {
|
|||||||
],
|
],
|
||||||
"REQ-12": [
|
"REQ-12": [
|
||||||
"tests/test_prompts.py",
|
"tests/test_prompts.py",
|
||||||
"tests/services/test_transcription.py",
|
"tests/services/test_transcription_external.py",
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+32
-21
@@ -21,9 +21,10 @@ from transcription.db import initialize_database_runtime
|
|||||||
from transcription.models import Document
|
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 Revision
|
||||||
|
from transcription.models import Source
|
||||||
|
|
||||||
TranscriptSeed = tuple[int, str | None, str | None]
|
RevisionSeed = str
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session")
|
@pytest.fixture(scope="session")
|
||||||
@@ -54,7 +55,8 @@ def clear_ui_database(app_client: tuple[FastAPI, TestClient]) -> None:
|
|||||||
|
|
||||||
async def _clear() -> None:
|
async def _clear() -> None:
|
||||||
async with get_session(session_factory=app.state.runtime.session_factory) as session:
|
async with get_session(session_factory=app.state.runtime.session_factory) as session:
|
||||||
await session.exec(delete(Transcript))
|
await session.exec(delete(Revision))
|
||||||
|
await session.exec(delete(Source))
|
||||||
await session.exec(delete(Job))
|
await session.exec(delete(Job))
|
||||||
await session.exec(delete(Document))
|
await session.exec(delete(Document))
|
||||||
await session.commit()
|
await session.commit()
|
||||||
@@ -64,7 +66,7 @@ def clear_ui_database(app_client: tuple[FastAPI, TestClient]) -> None:
|
|||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def seed_job(app_client: tuple[FastAPI, TestClient]) -> Callable[..., UUID]:
|
def seed_job(app_client: tuple[FastAPI, TestClient]) -> Callable[..., UUID]:
|
||||||
"""Return a helper for inserting a document/job/transcript trio."""
|
"""Return a helper for inserting a document/job/source/(optional revision) tuple."""
|
||||||
app, _ = app_client
|
app, _ = app_client
|
||||||
fixtures_dir = Path(__file__).resolve().parents[1] / "fixtures" / "images" / "valid"
|
fixtures_dir = Path(__file__).resolve().parents[1] / "fixtures" / "images" / "valid"
|
||||||
|
|
||||||
@@ -72,9 +74,9 @@ def seed_job(app_client: tuple[FastAPI, TestClient]) -> Callable[..., UUID]:
|
|||||||
*,
|
*,
|
||||||
filename: str = "sample.pdf",
|
filename: str = "sample.pdf",
|
||||||
status: JobStatus = JobStatus.TRANSCRIBED,
|
status: JobStatus = JobStatus.TRANSCRIBED,
|
||||||
transcript_text: str | None = "Sample transcript text",
|
transcription_text: str | None = "Sample transcript text",
|
||||||
error_detail: str | None = None,
|
error_detail: str | None = None,
|
||||||
transcript_revisions: list[TranscriptSeed] | None = None,
|
revision_text: RevisionSeed | None = None,
|
||||||
source_file: Path | None = None,
|
source_file: Path | None = None,
|
||||||
) -> UUID:
|
) -> UUID:
|
||||||
async def _insert() -> UUID:
|
async def _insert() -> UUID:
|
||||||
@@ -84,29 +86,38 @@ def seed_job(app_client: tuple[FastAPI, TestClient]) -> Callable[..., UUID]:
|
|||||||
source_path = source_file or fixtures_dir / "small_png.png"
|
source_path = source_file or fixtures_dir / "small_png.png"
|
||||||
stored_path.write_bytes(source_path.read_bytes())
|
stored_path.write_bytes(source_path.read_bytes())
|
||||||
|
|
||||||
document = Document(filename=filename, file_path=str(stored_path))
|
document = Document(name=filename)
|
||||||
session.add(document)
|
session.add(document)
|
||||||
await session.flush()
|
await session.flush()
|
||||||
|
|
||||||
job = Job(document_id=document.id, status=status, retry_count=0)
|
job = Job(
|
||||||
|
document_id=document.id,
|
||||||
|
status=status,
|
||||||
|
retry_count=0,
|
||||||
|
text=transcription_text,
|
||||||
|
error_detail=error_detail,
|
||||||
|
provider="openrouter",
|
||||||
|
model="google/gemini-2.5-flash",
|
||||||
|
prompt_name="transcribe_document.md",
|
||||||
|
)
|
||||||
session.add(job)
|
session.add(job)
|
||||||
await session.flush()
|
await session.flush()
|
||||||
|
|
||||||
revisions = transcript_revisions
|
source = Source(
|
||||||
if revisions is None and (transcript_text is not None or error_detail is not None):
|
document_id=document.id,
|
||||||
revisions = [(0, transcript_text, error_detail)]
|
|
||||||
|
|
||||||
if revisions is not None:
|
|
||||||
for revision, revision_text, revision_error in revisions:
|
|
||||||
session.add(
|
|
||||||
Transcript(
|
|
||||||
job_id=job.id,
|
job_id=job.id,
|
||||||
revision=revision,
|
upload_name=filename,
|
||||||
provider="openrouter",
|
filename=filename,
|
||||||
model="google/gemini-2.5-flash",
|
file_path=str(stored_path),
|
||||||
prompt_name="transcribe_document",
|
)
|
||||||
|
session.add(source)
|
||||||
|
await session.flush()
|
||||||
|
|
||||||
|
if revision_text is not None:
|
||||||
|
session.add(
|
||||||
|
Revision(
|
||||||
|
source_id=source.id,
|
||||||
text=revision_text,
|
text=revision_text,
|
||||||
error_detail=revision_error,
|
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -18,13 +18,12 @@ class TestPageRendering:
|
|||||||
response = client.get("/ui/jobs")
|
response = client.get("/ui/jobs")
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
assert "Transcription Jobs" in response.text
|
|
||||||
assert "No jobs yet." in response.text
|
assert "No jobs yet." in response.text
|
||||||
|
|
||||||
def test_jobs_page_lists_seeded_jobs(self, app_client, seed_job):
|
def test_jobs_page_lists_seeded_jobs(self, app_client, seed_job):
|
||||||
"""GET /ui/jobs lists seeded jobs from the in-memory database."""
|
"""GET /ui/jobs lists seeded jobs from the in-memory database."""
|
||||||
_, client = app_client
|
_, client = app_client
|
||||||
seed_job(filename="sample.pdf", status=JobStatus.TRANSCRIBED, transcript_text="done")
|
seed_job(filename="sample.pdf", status=JobStatus.TRANSCRIBED, transcription_text="done")
|
||||||
|
|
||||||
response = client.get("/ui/jobs")
|
response = client.get("/ui/jobs")
|
||||||
|
|
||||||
@@ -39,23 +38,20 @@ class TestPageRendering:
|
|||||||
job_id = seed_job(
|
job_id = seed_job(
|
||||||
filename="detail.pdf",
|
filename="detail.pdf",
|
||||||
status=JobStatus.TRANSCRIBED,
|
status=JobStatus.TRANSCRIBED,
|
||||||
transcript_revisions=[
|
transcription_text="original text",
|
||||||
(0, None, "first attempt failed"),
|
revision_text="hello",
|
||||||
(1, "hello", None),
|
|
||||||
],
|
|
||||||
source_file=fixture_path,
|
source_file=fixture_path,
|
||||||
)
|
)
|
||||||
|
|
||||||
response = client.get(f"/ui/jobs/{job_id}")
|
response = client.get(f"/ui/jobs/{job_id}")
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
assert "Job Detail" in response.text
|
assert "Original Transcription" in response.text
|
||||||
assert "Job overview" in response.text
|
|
||||||
assert "detail.pdf" in response.text
|
assert "detail.pdf" in response.text
|
||||||
assert "Transcripts" in response.text
|
|
||||||
assert "Revision" in response.text
|
assert "Revision" in response.text
|
||||||
assert "first attempt failed" in response.text
|
assert "Revision" in response.text
|
||||||
assert "hello" in response.text
|
assert "hello" in response.text
|
||||||
|
assert "original text" in response.text
|
||||||
assert "Document preview" in response.text
|
assert "Document preview" in response.text
|
||||||
assert "/uploads/detail.pdf" in response.text
|
assert "/uploads/detail.pdf" in response.text
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user