generated from john/python-template
This commit is contained in:
@@ -50,7 +50,7 @@ BACKUP_RETENTION_DAYS=14
|
||||
# --- worker reliability ---
|
||||
WORKER_MAX_RETRIES=0
|
||||
WORKER_PROVIDER_TIMEOUT_SECONDS=30.0
|
||||
WORKER_STALE_JOB_SECONDS=30.0
|
||||
WORKER_STALE_JOB_SECONDS=90.0
|
||||
WORKER_RETRY_BACKOFF_SECONDS=1.0
|
||||
WORKER_SHUTDOWN_GRACE_SECONDS=5.0
|
||||
WORKER_POLL_INTERVAL_SECONDS=1.0
|
||||
|
||||
@@ -31,9 +31,9 @@ Enforced by `tests/test_provider_boundaries.py`.
|
||||
|
||||
## Contract Surface
|
||||
|
||||
- Every adapter satisfies the `TranscriptionProvider` protocol in `base.py`, including
|
||||
`current_request_manifest` and `current_transport_evidence`, which exist so a *failed* call still
|
||||
yields evidence.
|
||||
- Every adapter satisfies the `TranscriptionProvider` protocol in `base.py`. Failed-call evidence is
|
||||
returned through the caller-owned `ProviderCallEvidence` sink passed to `transcribe()`, so
|
||||
evidence stays scoped to one invocation instead of living on mutable adapter instance state.
|
||||
- `TranscriptionResult`, `RequestManifest`, and `TransportEvidence` are `extra="forbid"` and frozen.
|
||||
Add a field to the contract rather than smuggling data through an untyped dict.
|
||||
- Evidence contracts in `evidence.py` are versioned (`schema_name` + `schema_version`). A change to
|
||||
@@ -51,6 +51,8 @@ wins if this file drifts from it.
|
||||
- Reset per-call capture state at the start of every call. Without it, a connection failure can
|
||||
attach the *previous* call's response as evidence for this one. Guarded by
|
||||
`tests/test_v42_evidence.py::test_openrouter_does_not_reuse_prior_response_on_connection_failure`.
|
||||
- Keep transport capture scoped to the call, not the adapter instance. Concurrent `transcribe()`
|
||||
calls on one adapter must not be able to overwrite each other's response evidence.
|
||||
- Handle the streamed-body case (`httpx.ResponseNotRead`) rather than assuming `response.content`
|
||||
is always available.
|
||||
- When no response arrives — timeout, DNS, connection reset — emit
|
||||
|
||||
@@ -37,6 +37,7 @@ PromptFilename = Annotated[str, StringConstraints(strip_whitespace=True, min_len
|
||||
Probability = Annotated[float, Field(ge=0.0, le=1.0)]
|
||||
Temperature = Annotated[float, Field(ge=0.0, le=2.0)]
|
||||
DEFAULT_PROVIDER_MODEL = "google/gemini-2.5-flash"
|
||||
WORKER_STALE_TIMEOUT_MULTIPLIER = 3.0
|
||||
|
||||
|
||||
class SqliteSettings(BaseModel):
|
||||
@@ -114,7 +115,7 @@ class Settings(BaseSettings):
|
||||
# Bounded only from below. Vision transcription of a dense page routinely runs
|
||||
# well past twenty seconds, so an upper cap here would silently fail real work.
|
||||
worker_provider_timeout_seconds: float = Field(default=30.0, gt=0.0)
|
||||
worker_stale_job_seconds: float = Field(default=30.0, gt=0.0)
|
||||
worker_stale_job_seconds: float = Field(default=90.0, gt=0.0)
|
||||
worker_retry_backoff_seconds: float = Field(default=1.0, ge=0.0)
|
||||
worker_shutdown_grace_seconds: float = Field(default=5.0, ge=0.0)
|
||||
worker_poll_interval_seconds: float = Field(default=1.0, gt=0.0)
|
||||
@@ -169,6 +170,34 @@ class Settings(BaseSettings):
|
||||
|
||||
return {**data, "provider_model": default_model, "provider_models": tuple(deduplicated)}
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def _derive_worker_stale_job_seconds(cls, data: object) -> object:
|
||||
"""Default stale-job recovery with margin over one provider timeout."""
|
||||
if not isinstance(data, dict):
|
||||
return data
|
||||
if data.get("worker_stale_job_seconds") is not None:
|
||||
return data
|
||||
|
||||
timeout = data.get("worker_provider_timeout_seconds", 30.0)
|
||||
if not isinstance(timeout, (str, int, float)):
|
||||
return data
|
||||
try:
|
||||
timeout_seconds = float(timeout)
|
||||
except ValueError:
|
||||
return data
|
||||
return {
|
||||
**data,
|
||||
"worker_stale_job_seconds": timeout_seconds * WORKER_STALE_TIMEOUT_MULTIPLIER,
|
||||
}
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_worker_stale_job_seconds(self) -> "Settings":
|
||||
"""Reject stale recovery that can fire before one provider timeout expires."""
|
||||
if self.worker_stale_job_seconds <= self.worker_provider_timeout_seconds:
|
||||
raise ValueError("WORKER_STALE_JOB_SECONDS must exceed WORKER_PROVIDER_TIMEOUT_SECONDS")
|
||||
return self
|
||||
|
||||
@property
|
||||
def should_bootstrap_schema(self) -> bool:
|
||||
"""Return whether startup should auto-create schema for this environment."""
|
||||
|
||||
@@ -4,6 +4,7 @@ from transcription.config import Provider
|
||||
from transcription.config import Settings
|
||||
from transcription.config import get_settings
|
||||
from transcription.providers.base import ProviderAuthError
|
||||
from transcription.providers.base import ProviderCallEvidence
|
||||
from transcription.providers.base import ProviderError
|
||||
from transcription.providers.base import ProviderResponseError
|
||||
from transcription.providers.base import TranscriptionMetadata
|
||||
@@ -27,6 +28,7 @@ def get_transcription_provider(*, settings: Settings | None = None) -> Transcrip
|
||||
__all__ = [
|
||||
"OpenRouterTranscriptionProvider",
|
||||
"ProviderAuthError",
|
||||
"ProviderCallEvidence",
|
||||
"ProviderError",
|
||||
"ProviderResponseError",
|
||||
"RequestManifest",
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Provider interfaces and validated shared contracts for transcription adapters."""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol
|
||||
|
||||
from pydantic import BaseModel
|
||||
@@ -37,6 +38,14 @@ class ProviderResponseError(ProviderError):
|
||||
"""Raised when provider responses are malformed or unusable."""
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ProviderCallEvidence:
|
||||
"""Caller-owned evidence sink for one provider invocation."""
|
||||
|
||||
request_manifest: RequestManifest | None = None
|
||||
transport_evidence: TransportEvidence | None = None
|
||||
|
||||
|
||||
class ProviderUsage(BaseModel):
|
||||
"""Normalized provider token accounting."""
|
||||
|
||||
@@ -107,16 +116,6 @@ class TranscriptionProvider(Protocol):
|
||||
"""Return the resolved model slug this adapter will call."""
|
||||
...
|
||||
|
||||
@property
|
||||
def current_request_manifest(self) -> RequestManifest | None:
|
||||
"""Return the manifest for the most recent call, for failure evidence."""
|
||||
...
|
||||
|
||||
@property
|
||||
def current_transport_evidence(self) -> TransportEvidence | None:
|
||||
"""Return transport-level evidence for the most recent call."""
|
||||
...
|
||||
|
||||
async def transcribe(
|
||||
self,
|
||||
*,
|
||||
@@ -127,8 +126,9 @@ class TranscriptionProvider(Protocol):
|
||||
top_p: float | None = None,
|
||||
source_reference: SourceEvidenceReference | None = None,
|
||||
requested_model: str | None = None,
|
||||
evidence_capture: ProviderCallEvidence | None = None,
|
||||
) -> TranscriptionResult:
|
||||
"""Transcribe the provided image according to the prompt text."""
|
||||
"""Transcribe one source and write failure evidence into the provided capture sink."""
|
||||
...
|
||||
|
||||
async def aclose(self) -> None:
|
||||
|
||||
@@ -3,11 +3,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import contextvars
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import AsyncIterator
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Annotated
|
||||
from typing import Any
|
||||
from typing import Literal
|
||||
@@ -25,6 +27,7 @@ from pydantic import ValidationError
|
||||
from transcription.config import Settings
|
||||
from transcription.config import get_settings
|
||||
from transcription.providers.base import ProviderAuthError
|
||||
from transcription.providers.base import ProviderCallEvidence
|
||||
from transcription.providers.base import ProviderError
|
||||
from transcription.providers.base import ProviderResponseError
|
||||
from transcription.providers.base import ProviderUsage
|
||||
@@ -65,19 +68,24 @@ class _CapturingAsyncClient:
|
||||
|
||||
def __init__(self, client: httpx.AsyncClient):
|
||||
self._client = client
|
||||
self.last_response: httpx.Response | None = None
|
||||
self.last_body: bytes | None = None
|
||||
self._active_capture: contextvars.ContextVar[_TransportCapture | None] = contextvars.ContextVar(
|
||||
"openrouter_transport_capture",
|
||||
default=None,
|
||||
)
|
||||
|
||||
async def send(self, request: httpx.Request, **kwargs: Any) -> httpx.Response:
|
||||
capture = self._active_capture.get()
|
||||
response = await self._client.send(request, **kwargs)
|
||||
self.last_response = response
|
||||
if capture is None:
|
||||
return response
|
||||
capture.response = response
|
||||
try:
|
||||
self.last_body = response.content
|
||||
capture.body = response.content
|
||||
except httpx.ResponseNotRead:
|
||||
stream = response.stream
|
||||
if not isinstance(stream, httpx.AsyncByteStream):
|
||||
raise
|
||||
response.stream = _CapturingAsyncByteStream(stream, self._capture_body)
|
||||
response.stream = _CapturingAsyncByteStream(stream, lambda body: self._capture_body(capture, body))
|
||||
return response
|
||||
|
||||
def build_request(self, *args: Any, **kwargs: Any) -> httpx.Request:
|
||||
@@ -86,12 +94,21 @@ class _CapturingAsyncClient:
|
||||
async def aclose(self) -> None:
|
||||
await self._client.aclose()
|
||||
|
||||
def reset(self) -> None:
|
||||
self.last_response = None
|
||||
self.last_body = None
|
||||
def begin_capture(self, capture: _TransportCapture) -> contextvars.Token[_TransportCapture | None]:
|
||||
return self._active_capture.set(capture)
|
||||
|
||||
def _capture_body(self, body: bytes) -> None:
|
||||
self.last_body = body
|
||||
def end_capture(self, token: contextvars.Token[_TransportCapture | None]) -> None:
|
||||
self._active_capture.reset(token)
|
||||
|
||||
@staticmethod
|
||||
def _capture_body(capture: _TransportCapture, body: bytes) -> None:
|
||||
capture.body = body
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _TransportCapture:
|
||||
response: httpx.Response | None = None
|
||||
body: bytes | None = None
|
||||
|
||||
|
||||
class _ProviderModel(BaseModel):
|
||||
@@ -195,8 +212,6 @@ class OpenRouterTranscriptionProvider:
|
||||
self._settings = settings or get_settings()
|
||||
self._model = self._settings.provider_model or DEFAULT_OPENROUTER_MODEL
|
||||
self._capturing_client: _CapturingAsyncClient | None = None
|
||||
self._current_request_manifest: RequestManifest | None = None
|
||||
self._current_transport_evidence: TransportEvidence | None = None
|
||||
if client is None:
|
||||
# httpx defaults every phase to 5s, which silently caps provider calls far
|
||||
# below worker_provider_timeout_seconds. Track the configured budget instead.
|
||||
@@ -218,18 +233,6 @@ class OpenRouterTranscriptionProvider:
|
||||
"""Return the resolved OpenRouter model slug."""
|
||||
return self._model
|
||||
|
||||
@property
|
||||
def current_request_manifest(self) -> RequestManifest | None:
|
||||
return self._current_request_manifest
|
||||
|
||||
@property
|
||||
def current_transport_evidence(self) -> TransportEvidence | None:
|
||||
if self._current_transport_evidence is not None:
|
||||
return self._current_transport_evidence
|
||||
if self._current_request_manifest is None:
|
||||
return None
|
||||
return self._captured_transport_evidence()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
if self._capturing_client is not None:
|
||||
await self._capturing_client.aclose()
|
||||
@@ -244,6 +247,7 @@ class OpenRouterTranscriptionProvider:
|
||||
top_p: float | None = None,
|
||||
source_reference: SourceEvidenceReference | None = None,
|
||||
requested_model: str | None = None,
|
||||
evidence_capture: ProviderCallEvidence | None = None,
|
||||
) -> TranscriptionResult:
|
||||
"""Send prompt + image to OpenRouter and return normalized text output."""
|
||||
request = self._build_request(
|
||||
@@ -261,18 +265,21 @@ class OpenRouterTranscriptionProvider:
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
)
|
||||
self._current_request_manifest = manifest
|
||||
self._current_transport_evidence = None
|
||||
if self._capturing_client is not None:
|
||||
self._capturing_client.reset()
|
||||
if evidence_capture is not None:
|
||||
evidence_capture.request_manifest = manifest
|
||||
evidence_capture.transport_evidence = None
|
||||
transport_capture = _TransportCapture()
|
||||
token = self._capturing_client.begin_capture(transport_capture) if self._capturing_client is not None else None
|
||||
try:
|
||||
try:
|
||||
response = await self._client.chat.send_async(
|
||||
**request.model_dump(mode="json", exclude_none=True),
|
||||
retries=None,
|
||||
)
|
||||
except Exception as exc:
|
||||
transport = self._captured_transport_evidence()
|
||||
self._current_transport_evidence = transport
|
||||
transport = self._captured_transport_evidence(transport_capture)
|
||||
if evidence_capture is not None:
|
||||
evidence_capture.transport_evidence = transport
|
||||
if isinstance(exc, openrouter_errors.UnauthorizedResponseError):
|
||||
raise ProviderAuthError(
|
||||
"OpenRouter authentication failed",
|
||||
@@ -294,8 +301,9 @@ class OpenRouterTranscriptionProvider:
|
||||
failure_phase=failure_phase,
|
||||
) from exc
|
||||
|
||||
transport = self._captured_transport_evidence()
|
||||
self._current_transport_evidence = transport
|
||||
transport = self._captured_transport_evidence(transport_capture)
|
||||
if evidence_capture is not None:
|
||||
evidence_capture.transport_evidence = transport
|
||||
raw_api_response = self._coerce_raw_response(response)
|
||||
try:
|
||||
validated_response = OpenRouterResponse.model_validate(raw_api_response)
|
||||
@@ -334,6 +342,10 @@ class OpenRouterTranscriptionProvider:
|
||||
request_manifest=manifest,
|
||||
transport_evidence=transport,
|
||||
)
|
||||
finally:
|
||||
client = self._capturing_client
|
||||
if token is not None and client is not None:
|
||||
client.end_capture(token)
|
||||
|
||||
def _build_request_manifest(
|
||||
self,
|
||||
@@ -394,16 +406,15 @@ class OpenRouterTranscriptionProvider:
|
||||
return [self._replace_embedded_media(item, source_reference=source_reference) for item in value]
|
||||
return value
|
||||
|
||||
def _captured_transport_evidence(self) -> TransportEvidence:
|
||||
response = self._capturing_client.last_response if self._capturing_client is not None else None
|
||||
def _captured_transport_evidence(self, capture: _TransportCapture) -> TransportEvidence:
|
||||
response = capture.response
|
||||
if response is None:
|
||||
return TransportEvidence(response_received=False)
|
||||
headers = filter_safe_response_headers(response.headers)
|
||||
body = self._capturing_client.last_body if self._capturing_client is not None else None
|
||||
return TransportEvidence(
|
||||
response_received=True,
|
||||
status_code=response.status_code,
|
||||
body=body,
|
||||
body=capture.body,
|
||||
safe_headers=headers,
|
||||
content_type=headers.get("content-type"),
|
||||
content_encoding=headers.get("content-encoding"),
|
||||
|
||||
@@ -183,6 +183,16 @@ class JobService(ServiceBase):
|
||||
await self._finalize(session=_session, caller_session=session, refresh=(job,))
|
||||
return job
|
||||
|
||||
async def note_processing_progress(self, *, job_id: UUID, session: AsyncSession | None = None) -> Job:
|
||||
"""Refresh job liveness while a multi-page batch is still in progress."""
|
||||
async with self._session_scope(session) as _session:
|
||||
job = await _session.get(Job, job_id)
|
||||
if job is None:
|
||||
raise self._not_found(job_id)
|
||||
job.date_updated = _utc_now_naive()
|
||||
await self._finalize(session=_session, caller_session=session, refresh=(job,))
|
||||
return job
|
||||
|
||||
async def claim_next_queued_job(
|
||||
self,
|
||||
*,
|
||||
|
||||
@@ -38,6 +38,7 @@ from transcription.db.models import JobSourceStatus
|
||||
from transcription.db.models import Source
|
||||
from transcription.errors import ErrorCategory
|
||||
from transcription.providers import ProviderAuthError
|
||||
from transcription.providers import ProviderCallEvidence
|
||||
from transcription.providers import ProviderError
|
||||
from transcription.providers import ProviderResponseError
|
||||
from transcription.providers import RequestManifest
|
||||
@@ -769,6 +770,7 @@ async def transcribe_document_image(
|
||||
provider: TranscriptionProvider | None = None,
|
||||
source_reference: SourceEvidenceReference | None = None,
|
||||
requested_model: str | None = None,
|
||||
evidence_capture: ProviderCallEvidence | None = None,
|
||||
) -> TranscriptionResult:
|
||||
"""Transcribe a local image using the configured prompt and provider."""
|
||||
runtime_settings = settings or get_settings()
|
||||
@@ -802,6 +804,7 @@ async def transcribe_document_image(
|
||||
top_p=prompt_execution.top_p,
|
||||
source_reference=source_reference,
|
||||
requested_model=requested_model,
|
||||
evidence_capture=evidence_capture,
|
||||
)
|
||||
finally:
|
||||
if owns_adapter:
|
||||
|
||||
@@ -21,6 +21,7 @@ from ..errors import AppError
|
||||
from ..errors import ErrorCategory
|
||||
from ..errors import classify_unexpected_error
|
||||
from ..errors import format_error_detail
|
||||
from ..providers import ProviderCallEvidence
|
||||
from ..providers import ProviderError
|
||||
from ..providers import RequestManifest
|
||||
from ..providers import SourceEvidenceReference
|
||||
@@ -284,6 +285,7 @@ async def process_queued_job( # noqa: PLR0915
|
||||
result: TranscriptionResult | None = None
|
||||
provider_input = None
|
||||
page_outcome: _SuccessfulPage | _FailedPage
|
||||
provider_call_evidence = ProviderCallEvidence()
|
||||
try:
|
||||
provider_input = build_provider_input(source, upload_dir=runtime_settings.upload_dir)
|
||||
source_reference = SourceEvidenceReference(
|
||||
@@ -308,6 +310,7 @@ async def process_queued_job( # noqa: PLR0915
|
||||
provider=provider,
|
||||
source_reference=source_reference,
|
||||
requested_model=source_job.model,
|
||||
evidence_capture=provider_call_evidence,
|
||||
),
|
||||
timeout=runtime_settings.worker_provider_timeout_seconds,
|
||||
)
|
||||
@@ -366,8 +369,8 @@ async def process_queued_job( # noqa: PLR0915
|
||||
_duration_ms_between(started_at, finished_at),
|
||||
max(0, int((asyncio.get_running_loop().time() - monotonic_started_at) * 1000)),
|
||||
),
|
||||
request_manifest=provider.current_request_manifest,
|
||||
transport_evidence=provider.current_transport_evidence,
|
||||
request_manifest=provider_call_evidence.request_manifest,
|
||||
transport_evidence=provider_call_evidence.transport_evidence,
|
||||
failure_phase="local_timeout",
|
||||
)
|
||||
failed_pages.append(page_outcome)
|
||||
@@ -673,10 +676,12 @@ async def _persist_page_outcome(
|
||||
if session is None:
|
||||
async with unit_of_work(services=services, session=session) as local_session:
|
||||
await _write_page_outcome(job=job, services=services, page=page, session=local_session)
|
||||
await services.jobs.note_processing_progress(job_id=job.id, session=local_session)
|
||||
await local_session.commit()
|
||||
return
|
||||
|
||||
await _write_page_outcome(job=job, services=services, page=page, session=session)
|
||||
await services.jobs.note_processing_progress(job_id=job.id, session=session)
|
||||
await session.commit()
|
||||
|
||||
|
||||
|
||||
@@ -99,6 +99,7 @@ class TestPipelineSuccessFlow:
|
||||
provider=None,
|
||||
source_reference=None,
|
||||
requested_model=None,
|
||||
evidence_capture=None,
|
||||
) -> TranscriptionResult:
|
||||
_ = (
|
||||
image_path,
|
||||
@@ -110,6 +111,7 @@ class TestPipelineSuccessFlow:
|
||||
provider,
|
||||
source_reference,
|
||||
requested_model,
|
||||
evidence_capture,
|
||||
)
|
||||
return TranscriptionResult(
|
||||
text="Pipeline transcript",
|
||||
@@ -191,9 +193,20 @@ class TestPipelineSuccessFlow:
|
||||
provider=None,
|
||||
source_reference=None,
|
||||
requested_model=None,
|
||||
evidence_capture=None,
|
||||
) -> TranscriptionResult:
|
||||
page_name = Path(image_path).name
|
||||
_ = (prompt_name, prompt_text, temperature, top_p, settings, provider, source_reference, requested_model)
|
||||
_ = (
|
||||
prompt_name,
|
||||
prompt_text,
|
||||
temperature,
|
||||
top_p,
|
||||
settings,
|
||||
provider,
|
||||
source_reference,
|
||||
requested_model,
|
||||
evidence_capture,
|
||||
)
|
||||
return TranscriptionResult(
|
||||
text=f"Transcript for {page_name}",
|
||||
provider="openrouter",
|
||||
@@ -260,6 +273,7 @@ class TestPipelineSuccessFlow:
|
||||
provider=None,
|
||||
source_reference=None,
|
||||
requested_model=None,
|
||||
evidence_capture=None,
|
||||
) -> TranscriptionResult:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
@@ -273,6 +287,7 @@ class TestPipelineSuccessFlow:
|
||||
provider,
|
||||
source_reference,
|
||||
requested_model,
|
||||
evidence_capture,
|
||||
)
|
||||
if call_count == 2:
|
||||
raise RuntimeError("simulated page failure")
|
||||
@@ -351,6 +366,7 @@ class TestPipelineSuccessFlow:
|
||||
provider=None,
|
||||
source_reference=None,
|
||||
requested_model=None,
|
||||
evidence_capture=None,
|
||||
) -> TranscriptionResult:
|
||||
nonlocal call_count
|
||||
_ = (
|
||||
@@ -363,6 +379,7 @@ class TestPipelineSuccessFlow:
|
||||
provider,
|
||||
source_reference,
|
||||
requested_model,
|
||||
evidence_capture,
|
||||
)
|
||||
call_count += 1
|
||||
return TranscriptionResult(
|
||||
@@ -419,6 +436,7 @@ class TestPipelineFailureFlow:
|
||||
provider=None,
|
||||
source_reference=None,
|
||||
requested_model=None,
|
||||
evidence_capture=None,
|
||||
) -> TranscriptionResult:
|
||||
_ = (
|
||||
image_path,
|
||||
@@ -430,6 +448,7 @@ class TestPipelineFailureFlow:
|
||||
provider,
|
||||
source_reference,
|
||||
requested_model,
|
||||
evidence_capture,
|
||||
)
|
||||
raise RuntimeError("pipeline provider failure")
|
||||
|
||||
|
||||
@@ -1,16 +1,21 @@
|
||||
"""Tests for transcription.providers.openrouter."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from types import SimpleNamespace
|
||||
from typing import cast
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openrouter import OpenRouter
|
||||
|
||||
from transcription.config import Settings
|
||||
from transcription.providers.base import ProviderCallEvidence
|
||||
from transcription.providers.base import ProviderError
|
||||
from transcription.providers.base import ProviderResponseError
|
||||
from transcription.providers.evidence import SourceEvidenceReference
|
||||
from transcription.providers.openrouter import DEFAULT_OPENROUTER_MODEL
|
||||
from transcription.providers.openrouter import OpenRouterTranscriptionProvider
|
||||
|
||||
@@ -245,3 +250,76 @@ class TestOpenRouterProviderTranscribe:
|
||||
|
||||
assert result.request_manifest is None
|
||||
assert "request manifest omitted" in caplog.text.lower()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_calls_keep_evidence_scoped_to_their_own_capture(self):
|
||||
first_release = asyncio.Event()
|
||||
|
||||
async def handler(request: httpx.Request) -> httpx.Response:
|
||||
payload = json.loads(request.content.decode("utf-8"))
|
||||
prompt = payload["messages"][0]["content"][0]["text"]
|
||||
if prompt == "First prompt":
|
||||
await first_release.wait()
|
||||
body = b'{"error":{"message":"first failure"}}'
|
||||
else:
|
||||
body = b'{"error":{"message":"second failure"}}'
|
||||
return httpx.Response(
|
||||
500,
|
||||
content=body,
|
||||
headers={"Content-Type": "application/json"},
|
||||
request=request,
|
||||
)
|
||||
|
||||
provider = OpenRouterTranscriptionProvider(
|
||||
settings=Settings(openrouter_api_key="test-key"),
|
||||
async_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
first_capture = ProviderCallEvidence()
|
||||
second_capture = ProviderCallEvidence()
|
||||
|
||||
first_task = asyncio.create_task(
|
||||
provider.transcribe(
|
||||
prompt_text="First prompt",
|
||||
image_bytes=b"one",
|
||||
mime_type="image/png",
|
||||
evidence_capture=first_capture,
|
||||
source_reference=SourceEvidenceReference(
|
||||
source_id=uuid4(),
|
||||
digest_sha256="1" * 64,
|
||||
byte_size=3,
|
||||
media_type="image/png",
|
||||
page_number=1,
|
||||
),
|
||||
)
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
second_task = asyncio.create_task(
|
||||
provider.transcribe(
|
||||
prompt_text="Second prompt",
|
||||
image_bytes=b"two",
|
||||
mime_type="image/png",
|
||||
evidence_capture=second_capture,
|
||||
source_reference=SourceEvidenceReference(
|
||||
source_id=uuid4(),
|
||||
digest_sha256="2" * 64,
|
||||
byte_size=3,
|
||||
media_type="image/png",
|
||||
page_number=2,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(ProviderError):
|
||||
await second_task
|
||||
first_release.set()
|
||||
with pytest.raises(ProviderError):
|
||||
await first_task
|
||||
|
||||
assert first_capture.request_manifest is not None
|
||||
assert first_capture.request_manifest.prompt_content == "First prompt"
|
||||
assert first_capture.transport_evidence is not None
|
||||
assert first_capture.transport_evidence.body == b'{"error":{"message":"first failure"}}'
|
||||
assert second_capture.request_manifest is not None
|
||||
assert second_capture.request_manifest.prompt_content == "Second prompt"
|
||||
assert second_capture.transport_evidence is not None
|
||||
assert second_capture.transport_evidence.body == b'{"error":{"message":"second failure"}}'
|
||||
|
||||
@@ -86,6 +86,7 @@ class TestWorkflowReliability:
|
||||
provider=None,
|
||||
source_reference=None,
|
||||
requested_model=None,
|
||||
evidence_capture=None,
|
||||
):
|
||||
_ = (
|
||||
image_path,
|
||||
@@ -97,6 +98,7 @@ class TestWorkflowReliability:
|
||||
provider,
|
||||
source_reference,
|
||||
requested_model,
|
||||
evidence_capture,
|
||||
)
|
||||
raise TimeoutError("simulated provider timeout")
|
||||
|
||||
@@ -370,6 +372,64 @@ class TestWorkflowReliability:
|
||||
assert result is not None
|
||||
assert result.status == JobStatus.TRANSCRIBED
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_intermediate_page_commit_advances_job_liveness_timestamp(
|
||||
self,
|
||||
default_session_factory,
|
||||
monkeypatch,
|
||||
):
|
||||
services = ServiceBundle.from_session_factory(default_session_factory)
|
||||
async with services.jobs._session_scope() as session:
|
||||
document = Document(id=uuid4(), name="heartbeat-doc")
|
||||
session.add(document)
|
||||
await session.flush()
|
||||
job = Job(document_id=document.id, status=JobStatus.QUEUED)
|
||||
session.add(job)
|
||||
await session.flush()
|
||||
for page_number in (1, 2, 3):
|
||||
source = Source(
|
||||
document_id=document.id,
|
||||
page_number=page_number,
|
||||
upload_name=f"page-{page_number}.jpg",
|
||||
filename=f"page-{page_number}.jpg",
|
||||
file_path=str(Path("tests/fixtures/images/real/Book Two - page 02.jpg")),
|
||||
file_hash=str(page_number) * 64,
|
||||
file_size_bytes=1,
|
||||
)
|
||||
session.add(source)
|
||||
await session.flush()
|
||||
session.add(JobSource(job_id=job.id, source_id=source.id))
|
||||
await session.commit()
|
||||
loaded = await services.jobs.read_job(job_id=job.id, session=session)
|
||||
|
||||
third_started = asyncio.Event()
|
||||
release_third = asyncio.Event()
|
||||
call_count = 0
|
||||
|
||||
async def _transcribe(image_path, **kwargs):
|
||||
nonlocal call_count
|
||||
_ = (image_path, kwargs)
|
||||
call_count += 1
|
||||
if call_count == 3:
|
||||
third_started.set()
|
||||
await release_third.wait()
|
||||
return TranscriptionResult(text=f"page {call_count}", provider="fixture", model="model")
|
||||
|
||||
monkeypatch.setattr("transcription.services.workflows.transcribe_document_image", _transcribe)
|
||||
task = asyncio.create_task(process_queued_job(job=loaded, services=services))
|
||||
await asyncio.wait_for(third_started.wait(), timeout=2)
|
||||
|
||||
async with services.jobs._session_scope() as session:
|
||||
current_job = await session.get(Job, loaded.id)
|
||||
assert current_job is not None
|
||||
attempts = await services.evidence.list_execution_attempts(job_id=loaded.id, session=session)
|
||||
|
||||
assert len(attempts) == 2
|
||||
assert current_job.date_updated >= attempts[1].finished_at
|
||||
|
||||
release_third.set()
|
||||
await task
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_job_with_validation_category_is_not_requeued(self, default_session_factory):
|
||||
services = ServiceBundle.from_session_factory(default_session_factory)
|
||||
|
||||
+19
-1
@@ -167,11 +167,29 @@ class TestWorkerReliabilitySettings:
|
||||
settings = _make_settings()
|
||||
assert settings.run_embedded_worker is True
|
||||
assert settings.worker_max_retries == 0
|
||||
assert settings.worker_stale_job_seconds == 30.0
|
||||
assert settings.worker_stale_job_seconds == 90.0
|
||||
assert settings.worker_retry_backoff_seconds == 1.0
|
||||
assert settings.worker_shutdown_grace_seconds == 5.0
|
||||
assert settings.worker_poll_interval_seconds == 1.0
|
||||
|
||||
def test_worker_stale_threshold_defaults_to_three_times_provider_timeout(self):
|
||||
settings = _make_settings(worker_provider_timeout_seconds=45.0)
|
||||
|
||||
assert settings.worker_stale_job_seconds == 135.0
|
||||
|
||||
def test_worker_stale_threshold_derives_from_environment_timeout(self, monkeypatch):
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "test-key")
|
||||
monkeypatch.setenv("WORKER_PROVIDER_TIMEOUT_SECONDS", "45.0")
|
||||
monkeypatch.delenv("WORKER_STALE_JOB_SECONDS", raising=False)
|
||||
|
||||
settings = Settings(_env_file=None)
|
||||
|
||||
assert settings.worker_stale_job_seconds == 135.0
|
||||
|
||||
def test_worker_stale_threshold_must_exceed_provider_timeout(self):
|
||||
with pytest.raises(ValidationError):
|
||||
_make_settings(worker_provider_timeout_seconds=45.0, worker_stale_job_seconds=45.0)
|
||||
|
||||
|
||||
def test_provider_timeout_is_not_capped_at_twenty_seconds():
|
||||
"""HIGH-03: vision transcription regularly runs past the old le=20.0 ceiling."""
|
||||
|
||||
Reference in New Issue
Block a user