Implement Rec - Phase 3 complete
Quality Gate / gate (push) Successful in 2m37s

This commit is contained in:
Jim Lancaster
2026-09-02 16:15:33 -05:00
parent 4410d23f5c
commit 27ec81ca5b
13 changed files with 353 additions and 116 deletions
+1 -1
View File
@@ -50,7 +50,7 @@ BACKUP_RETENTION_DAYS=14
# --- worker reliability --- # --- worker reliability ---
WORKER_MAX_RETRIES=0 WORKER_MAX_RETRIES=0
WORKER_PROVIDER_TIMEOUT_SECONDS=30.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_RETRY_BACKOFF_SECONDS=1.0
WORKER_SHUTDOWN_GRACE_SECONDS=5.0 WORKER_SHUTDOWN_GRACE_SECONDS=5.0
WORKER_POLL_INTERVAL_SECONDS=1.0 WORKER_POLL_INTERVAL_SECONDS=1.0
@@ -31,9 +31,9 @@ Enforced by `tests/test_provider_boundaries.py`.
## Contract Surface ## Contract Surface
- Every adapter satisfies the `TranscriptionProvider` protocol in `base.py`, including - Every adapter satisfies the `TranscriptionProvider` protocol in `base.py`. Failed-call evidence is
`current_request_manifest` and `current_transport_evidence`, which exist so a *failed* call still returned through the caller-owned `ProviderCallEvidence` sink passed to `transcribe()`, so
yields evidence. evidence stays scoped to one invocation instead of living on mutable adapter instance state.
- `TranscriptionResult`, `RequestManifest`, and `TransportEvidence` are `extra="forbid"` and frozen. - `TranscriptionResult`, `RequestManifest`, and `TransportEvidence` are `extra="forbid"` and frozen.
Add a field to the contract rather than smuggling data through an untyped dict. 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 - 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 - 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 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`. `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` - Handle the streamed-body case (`httpx.ResponseNotRead`) rather than assuming `response.content`
is always available. is always available.
- When no response arrives — timeout, DNS, connection reset — emit - When no response arrives — timeout, DNS, connection reset — emit
+30 -1
View File
@@ -37,6 +37,7 @@ PromptFilename = Annotated[str, StringConstraints(strip_whitespace=True, min_len
Probability = Annotated[float, Field(ge=0.0, le=1.0)] Probability = Annotated[float, Field(ge=0.0, le=1.0)]
Temperature = Annotated[float, Field(ge=0.0, le=2.0)] Temperature = Annotated[float, Field(ge=0.0, le=2.0)]
DEFAULT_PROVIDER_MODEL = "google/gemini-2.5-flash" DEFAULT_PROVIDER_MODEL = "google/gemini-2.5-flash"
WORKER_STALE_TIMEOUT_MULTIPLIER = 3.0
class SqliteSettings(BaseModel): class SqliteSettings(BaseModel):
@@ -114,7 +115,7 @@ class Settings(BaseSettings):
# Bounded only from below. Vision transcription of a dense page routinely runs # 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. # 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_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_retry_backoff_seconds: float = Field(default=1.0, ge=0.0)
worker_shutdown_grace_seconds: float = Field(default=5.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) 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)} 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 @property
def should_bootstrap_schema(self) -> bool: def should_bootstrap_schema(self) -> bool:
"""Return whether startup should auto-create schema for this environment.""" """Return whether startup should auto-create schema for this environment."""
+2
View File
@@ -4,6 +4,7 @@ from transcription.config import Provider
from transcription.config import Settings from transcription.config import Settings
from transcription.config import get_settings from transcription.config import get_settings
from transcription.providers.base import ProviderAuthError from transcription.providers.base import ProviderAuthError
from transcription.providers.base import ProviderCallEvidence
from transcription.providers.base import ProviderError from transcription.providers.base import ProviderError
from transcription.providers.base import ProviderResponseError from transcription.providers.base import ProviderResponseError
from transcription.providers.base import TranscriptionMetadata from transcription.providers.base import TranscriptionMetadata
@@ -27,6 +28,7 @@ def get_transcription_provider(*, settings: Settings | None = None) -> Transcrip
__all__ = [ __all__ = [
"OpenRouterTranscriptionProvider", "OpenRouterTranscriptionProvider",
"ProviderAuthError", "ProviderAuthError",
"ProviderCallEvidence",
"ProviderError", "ProviderError",
"ProviderResponseError", "ProviderResponseError",
"RequestManifest", "RequestManifest",
+11 -11
View File
@@ -1,5 +1,6 @@
"""Provider interfaces and validated shared contracts for transcription adapters.""" """Provider interfaces and validated shared contracts for transcription adapters."""
from dataclasses import dataclass
from typing import Protocol from typing import Protocol
from pydantic import BaseModel from pydantic import BaseModel
@@ -37,6 +38,14 @@ class ProviderResponseError(ProviderError):
"""Raised when provider responses are malformed or unusable.""" """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): class ProviderUsage(BaseModel):
"""Normalized provider token accounting.""" """Normalized provider token accounting."""
@@ -107,16 +116,6 @@ class TranscriptionProvider(Protocol):
"""Return the resolved model slug this adapter will call.""" """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( async def transcribe(
self, self,
*, *,
@@ -127,8 +126,9 @@ class TranscriptionProvider(Protocol):
top_p: float | None = None, top_p: float | None = None,
source_reference: SourceEvidenceReference | None = None, source_reference: SourceEvidenceReference | None = None,
requested_model: str | None = None, requested_model: str | None = None,
evidence_capture: ProviderCallEvidence | None = None,
) -> TranscriptionResult: ) -> 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: async def aclose(self) -> None:
+107 -96
View File
@@ -3,11 +3,13 @@
from __future__ import annotations from __future__ import annotations
import base64 import base64
import contextvars
import hashlib import hashlib
import json import json
import logging import logging
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from collections.abc import Callable from collections.abc import Callable
from dataclasses import dataclass
from typing import Annotated from typing import Annotated
from typing import Any from typing import Any
from typing import Literal from typing import Literal
@@ -25,6 +27,7 @@ from pydantic import ValidationError
from transcription.config import Settings from transcription.config import Settings
from transcription.config import get_settings from transcription.config import get_settings
from transcription.providers.base import ProviderAuthError from transcription.providers.base import ProviderAuthError
from transcription.providers.base import ProviderCallEvidence
from transcription.providers.base import ProviderError from transcription.providers.base import ProviderError
from transcription.providers.base import ProviderResponseError from transcription.providers.base import ProviderResponseError
from transcription.providers.base import ProviderUsage from transcription.providers.base import ProviderUsage
@@ -65,19 +68,24 @@ class _CapturingAsyncClient:
def __init__(self, client: httpx.AsyncClient): def __init__(self, client: httpx.AsyncClient):
self._client = client self._client = client
self.last_response: httpx.Response | None = None self._active_capture: contextvars.ContextVar[_TransportCapture | None] = contextvars.ContextVar(
self.last_body: bytes | None = None "openrouter_transport_capture",
default=None,
)
async def send(self, request: httpx.Request, **kwargs: Any) -> httpx.Response: async def send(self, request: httpx.Request, **kwargs: Any) -> httpx.Response:
capture = self._active_capture.get()
response = await self._client.send(request, **kwargs) response = await self._client.send(request, **kwargs)
self.last_response = response if capture is None:
return response
capture.response = response
try: try:
self.last_body = response.content capture.body = response.content
except httpx.ResponseNotRead: except httpx.ResponseNotRead:
stream = response.stream stream = response.stream
if not isinstance(stream, httpx.AsyncByteStream): if not isinstance(stream, httpx.AsyncByteStream):
raise raise
response.stream = _CapturingAsyncByteStream(stream, self._capture_body) response.stream = _CapturingAsyncByteStream(stream, lambda body: self._capture_body(capture, body))
return response return response
def build_request(self, *args: Any, **kwargs: Any) -> httpx.Request: def build_request(self, *args: Any, **kwargs: Any) -> httpx.Request:
@@ -86,12 +94,21 @@ class _CapturingAsyncClient:
async def aclose(self) -> None: async def aclose(self) -> None:
await self._client.aclose() await self._client.aclose()
def reset(self) -> None: def begin_capture(self, capture: _TransportCapture) -> contextvars.Token[_TransportCapture | None]:
self.last_response = None return self._active_capture.set(capture)
self.last_body = None
def _capture_body(self, body: bytes) -> None: def end_capture(self, token: contextvars.Token[_TransportCapture | None]) -> None:
self.last_body = body 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): class _ProviderModel(BaseModel):
@@ -195,8 +212,6 @@ class OpenRouterTranscriptionProvider:
self._settings = settings or get_settings() self._settings = settings or get_settings()
self._model = self._settings.provider_model or DEFAULT_OPENROUTER_MODEL self._model = self._settings.provider_model or DEFAULT_OPENROUTER_MODEL
self._capturing_client: _CapturingAsyncClient | None = None self._capturing_client: _CapturingAsyncClient | None = None
self._current_request_manifest: RequestManifest | None = None
self._current_transport_evidence: TransportEvidence | None = None
if client is None: if client is None:
# httpx defaults every phase to 5s, which silently caps provider calls far # httpx defaults every phase to 5s, which silently caps provider calls far
# below worker_provider_timeout_seconds. Track the configured budget instead. # below worker_provider_timeout_seconds. Track the configured budget instead.
@@ -218,18 +233,6 @@ class OpenRouterTranscriptionProvider:
"""Return the resolved OpenRouter model slug.""" """Return the resolved OpenRouter model slug."""
return self._model 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: async def aclose(self) -> None:
if self._capturing_client is not None: if self._capturing_client is not None:
await self._capturing_client.aclose() await self._capturing_client.aclose()
@@ -244,6 +247,7 @@ class OpenRouterTranscriptionProvider:
top_p: float | None = None, top_p: float | None = None,
source_reference: SourceEvidenceReference | None = None, source_reference: SourceEvidenceReference | None = None,
requested_model: str | None = None, requested_model: str | None = None,
evidence_capture: ProviderCallEvidence | None = None,
) -> TranscriptionResult: ) -> TranscriptionResult:
"""Send prompt + image to OpenRouter and return normalized text output.""" """Send prompt + image to OpenRouter and return normalized text output."""
request = self._build_request( request = self._build_request(
@@ -261,79 +265,87 @@ class OpenRouterTranscriptionProvider:
temperature=temperature, temperature=temperature,
top_p=top_p, top_p=top_p,
) )
self._current_request_manifest = manifest if evidence_capture is not None:
self._current_transport_evidence = None evidence_capture.request_manifest = manifest
if self._capturing_client is not None: evidence_capture.transport_evidence = None
self._capturing_client.reset() 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( try:
**request.model_dump(mode="json", exclude_none=True), response = await self._client.chat.send_async(
retries=None, **request.model_dump(mode="json", exclude_none=True),
) retries=None,
except Exception as exc: )
transport = self._captured_transport_evidence() except Exception as exc:
self._current_transport_evidence = transport transport = self._captured_transport_evidence(transport_capture)
if isinstance(exc, openrouter_errors.UnauthorizedResponseError): if evidence_capture is not None:
raise ProviderAuthError( evidence_capture.transport_evidence = transport
"OpenRouter authentication failed", if isinstance(exc, openrouter_errors.UnauthorizedResponseError):
raise ProviderAuthError(
"OpenRouter authentication failed",
request_manifest=manifest,
transport_evidence=transport,
failure_phase="http_response" if transport.response_received else "connection",
) from exc
failure_phase = (
"response_validation"
if isinstance(exc, openrouter_errors.ResponseValidationError)
else "http_response"
if transport.response_received
else "connection"
)
raise ProviderError(
self._transport_error_message(transport),
request_manifest=manifest, request_manifest=manifest,
transport_evidence=transport, transport_evidence=transport,
failure_phase="http_response" if transport.response_received else "connection", failure_phase=failure_phase,
) from exc ) from exc
failure_phase = (
"response_validation" transport = self._captured_transport_evidence(transport_capture)
if isinstance(exc, openrouter_errors.ResponseValidationError) if evidence_capture is not None:
else "http_response" evidence_capture.transport_evidence = transport
if transport.response_received raw_api_response = self._coerce_raw_response(response)
else "connection" try:
validated_response = OpenRouterResponse.model_validate(raw_api_response)
except ValidationError as exc:
raise ProviderResponseError(
"OpenRouter response failed schema validation",
request_manifest=manifest,
transport_evidence=transport,
failure_phase="response_validation",
) from exc
try:
text = self._extract_text(validated_response)
except ProviderResponseError as exc:
raise ProviderResponseError(
str(exc),
request_manifest=manifest,
transport_evidence=transport,
failure_phase="response_validation",
) from exc
model = validated_response.model or requested_model or self.model
metadata = self._build_metadata(validated_response)
logger.info("OpenRouter transcription completed using model=%s", model)
return TranscriptionResult(
text=text,
provider="openrouter",
prompt_name=None,
prompt_hash=None,
system_prompt=None,
user_prompt=prompt_text,
temperature=temperature,
top_p=top_p,
model=model,
metadata=metadata,
raw_api_response=raw_api_response,
request_manifest=manifest,
transport_evidence=transport,
) )
raise ProviderError( finally:
self._transport_error_message(transport), client = self._capturing_client
request_manifest=manifest, if token is not None and client is not None:
transport_evidence=transport, client.end_capture(token)
failure_phase=failure_phase,
) from exc
transport = self._captured_transport_evidence()
self._current_transport_evidence = transport
raw_api_response = self._coerce_raw_response(response)
try:
validated_response = OpenRouterResponse.model_validate(raw_api_response)
except ValidationError as exc:
raise ProviderResponseError(
"OpenRouter response failed schema validation",
request_manifest=manifest,
transport_evidence=transport,
failure_phase="response_validation",
) from exc
try:
text = self._extract_text(validated_response)
except ProviderResponseError as exc:
raise ProviderResponseError(
str(exc),
request_manifest=manifest,
transport_evidence=transport,
failure_phase="response_validation",
) from exc
model = validated_response.model or requested_model or self.model
metadata = self._build_metadata(validated_response)
logger.info("OpenRouter transcription completed using model=%s", model)
return TranscriptionResult(
text=text,
provider="openrouter",
prompt_name=None,
prompt_hash=None,
system_prompt=None,
user_prompt=prompt_text,
temperature=temperature,
top_p=top_p,
model=model,
metadata=metadata,
raw_api_response=raw_api_response,
request_manifest=manifest,
transport_evidence=transport,
)
def _build_request_manifest( def _build_request_manifest(
self, self,
@@ -394,16 +406,15 @@ class OpenRouterTranscriptionProvider:
return [self._replace_embedded_media(item, source_reference=source_reference) for item in value] return [self._replace_embedded_media(item, source_reference=source_reference) for item in value]
return value return value
def _captured_transport_evidence(self) -> TransportEvidence: def _captured_transport_evidence(self, capture: _TransportCapture) -> TransportEvidence:
response = self._capturing_client.last_response if self._capturing_client is not None else None response = capture.response
if response is None: if response is None:
return TransportEvidence(response_received=False) return TransportEvidence(response_received=False)
headers = filter_safe_response_headers(response.headers) headers = filter_safe_response_headers(response.headers)
body = self._capturing_client.last_body if self._capturing_client is not None else None
return TransportEvidence( return TransportEvidence(
response_received=True, response_received=True,
status_code=response.status_code, status_code=response.status_code,
body=body, body=capture.body,
safe_headers=headers, safe_headers=headers,
content_type=headers.get("content-type"), content_type=headers.get("content-type"),
content_encoding=headers.get("content-encoding"), content_encoding=headers.get("content-encoding"),
+10
View File
@@ -183,6 +183,16 @@ class JobService(ServiceBase):
await self._finalize(session=_session, caller_session=session, refresh=(job,)) await self._finalize(session=_session, caller_session=session, refresh=(job,))
return 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( async def claim_next_queued_job(
self, self,
*, *,
+3
View File
@@ -38,6 +38,7 @@ from transcription.db.models import JobSourceStatus
from transcription.db.models import Source from transcription.db.models import Source
from transcription.errors import ErrorCategory from transcription.errors import ErrorCategory
from transcription.providers import ProviderAuthError from transcription.providers import ProviderAuthError
from transcription.providers import ProviderCallEvidence
from transcription.providers import ProviderError from transcription.providers import ProviderError
from transcription.providers import ProviderResponseError from transcription.providers import ProviderResponseError
from transcription.providers import RequestManifest from transcription.providers import RequestManifest
@@ -769,6 +770,7 @@ async def transcribe_document_image(
provider: TranscriptionProvider | None = None, provider: TranscriptionProvider | None = None,
source_reference: SourceEvidenceReference | None = None, source_reference: SourceEvidenceReference | None = None,
requested_model: str | None = None, requested_model: str | None = None,
evidence_capture: ProviderCallEvidence | None = None,
) -> TranscriptionResult: ) -> TranscriptionResult:
"""Transcribe a local image using the configured prompt and provider.""" """Transcribe a local image using the configured prompt and provider."""
runtime_settings = settings or get_settings() runtime_settings = settings or get_settings()
@@ -802,6 +804,7 @@ async def transcribe_document_image(
top_p=prompt_execution.top_p, top_p=prompt_execution.top_p,
source_reference=source_reference, source_reference=source_reference,
requested_model=requested_model, requested_model=requested_model,
evidence_capture=evidence_capture,
) )
finally: finally:
if owns_adapter: if owns_adapter:
+7 -2
View File
@@ -21,6 +21,7 @@ from ..errors import AppError
from ..errors import ErrorCategory 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 ..providers import ProviderCallEvidence
from ..providers import ProviderError from ..providers import ProviderError
from ..providers import RequestManifest from ..providers import RequestManifest
from ..providers import SourceEvidenceReference from ..providers import SourceEvidenceReference
@@ -284,6 +285,7 @@ async def process_queued_job( # noqa: PLR0915
result: TranscriptionResult | None = None result: TranscriptionResult | None = None
provider_input = None provider_input = None
page_outcome: _SuccessfulPage | _FailedPage page_outcome: _SuccessfulPage | _FailedPage
provider_call_evidence = ProviderCallEvidence()
try: try:
provider_input = build_provider_input(source, upload_dir=runtime_settings.upload_dir) provider_input = build_provider_input(source, upload_dir=runtime_settings.upload_dir)
source_reference = SourceEvidenceReference( source_reference = SourceEvidenceReference(
@@ -308,6 +310,7 @@ async def process_queued_job( # noqa: PLR0915
provider=provider, provider=provider,
source_reference=source_reference, source_reference=source_reference,
requested_model=source_job.model, requested_model=source_job.model,
evidence_capture=provider_call_evidence,
), ),
timeout=runtime_settings.worker_provider_timeout_seconds, 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), _duration_ms_between(started_at, finished_at),
max(0, int((asyncio.get_running_loop().time() - monotonic_started_at) * 1000)), max(0, int((asyncio.get_running_loop().time() - monotonic_started_at) * 1000)),
), ),
request_manifest=provider.current_request_manifest, request_manifest=provider_call_evidence.request_manifest,
transport_evidence=provider.current_transport_evidence, transport_evidence=provider_call_evidence.transport_evidence,
failure_phase="local_timeout", failure_phase="local_timeout",
) )
failed_pages.append(page_outcome) failed_pages.append(page_outcome)
@@ -673,10 +676,12 @@ async def _persist_page_outcome(
if session is None: if session is None:
async with unit_of_work(services=services, session=session) as local_session: 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 _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() await local_session.commit()
return return
await _write_page_outcome(job=job, services=services, page=page, session=session) 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() await session.commit()
+20 -1
View File
@@ -99,6 +99,7 @@ class TestPipelineSuccessFlow:
provider=None, provider=None,
source_reference=None, source_reference=None,
requested_model=None, requested_model=None,
evidence_capture=None,
) -> TranscriptionResult: ) -> TranscriptionResult:
_ = ( _ = (
image_path, image_path,
@@ -110,6 +111,7 @@ class TestPipelineSuccessFlow:
provider, provider,
source_reference, source_reference,
requested_model, requested_model,
evidence_capture,
) )
return TranscriptionResult( return TranscriptionResult(
text="Pipeline transcript", text="Pipeline transcript",
@@ -191,9 +193,20 @@ class TestPipelineSuccessFlow:
provider=None, provider=None,
source_reference=None, source_reference=None,
requested_model=None, requested_model=None,
evidence_capture=None,
) -> TranscriptionResult: ) -> TranscriptionResult:
page_name = Path(image_path).name 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( return TranscriptionResult(
text=f"Transcript for {page_name}", text=f"Transcript for {page_name}",
provider="openrouter", provider="openrouter",
@@ -260,6 +273,7 @@ class TestPipelineSuccessFlow:
provider=None, provider=None,
source_reference=None, source_reference=None,
requested_model=None, requested_model=None,
evidence_capture=None,
) -> TranscriptionResult: ) -> TranscriptionResult:
nonlocal call_count nonlocal call_count
call_count += 1 call_count += 1
@@ -273,6 +287,7 @@ class TestPipelineSuccessFlow:
provider, provider,
source_reference, source_reference,
requested_model, requested_model,
evidence_capture,
) )
if call_count == 2: if call_count == 2:
raise RuntimeError("simulated page failure") raise RuntimeError("simulated page failure")
@@ -351,6 +366,7 @@ class TestPipelineSuccessFlow:
provider=None, provider=None,
source_reference=None, source_reference=None,
requested_model=None, requested_model=None,
evidence_capture=None,
) -> TranscriptionResult: ) -> TranscriptionResult:
nonlocal call_count nonlocal call_count
_ = ( _ = (
@@ -363,6 +379,7 @@ class TestPipelineSuccessFlow:
provider, provider,
source_reference, source_reference,
requested_model, requested_model,
evidence_capture,
) )
call_count += 1 call_count += 1
return TranscriptionResult( return TranscriptionResult(
@@ -419,6 +436,7 @@ class TestPipelineFailureFlow:
provider=None, provider=None,
source_reference=None, source_reference=None,
requested_model=None, requested_model=None,
evidence_capture=None,
) -> TranscriptionResult: ) -> TranscriptionResult:
_ = ( _ = (
image_path, image_path,
@@ -430,6 +448,7 @@ class TestPipelineFailureFlow:
provider, provider,
source_reference, source_reference,
requested_model, requested_model,
evidence_capture,
) )
raise RuntimeError("pipeline provider failure") raise RuntimeError("pipeline provider failure")
+78
View File
@@ -1,16 +1,21 @@
"""Tests for transcription.providers.openrouter.""" """Tests for transcription.providers.openrouter."""
import asyncio import asyncio
import json
import logging import logging
from types import SimpleNamespace from types import SimpleNamespace
from typing import cast from typing import cast
from uuid import uuid4
import httpx
import pytest import pytest
from openrouter import OpenRouter from openrouter import OpenRouter
from transcription.config import Settings from transcription.config import Settings
from transcription.providers.base import ProviderCallEvidence
from transcription.providers.base import ProviderError from transcription.providers.base import ProviderError
from transcription.providers.base import ProviderResponseError 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 DEFAULT_OPENROUTER_MODEL
from transcription.providers.openrouter import OpenRouterTranscriptionProvider from transcription.providers.openrouter import OpenRouterTranscriptionProvider
@@ -245,3 +250,76 @@ class TestOpenRouterProviderTranscribe:
assert result.request_manifest is None assert result.request_manifest is None
assert "request manifest omitted" in caplog.text.lower() 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, provider=None,
source_reference=None, source_reference=None,
requested_model=None, requested_model=None,
evidence_capture=None,
): ):
_ = ( _ = (
image_path, image_path,
@@ -97,6 +98,7 @@ class TestWorkflowReliability:
provider, provider,
source_reference, source_reference,
requested_model, requested_model,
evidence_capture,
) )
raise TimeoutError("simulated provider timeout") raise TimeoutError("simulated provider timeout")
@@ -370,6 +372,64 @@ class TestWorkflowReliability:
assert result is not None assert result is not None
assert result.status == JobStatus.TRANSCRIBED 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 @pytest.mark.asyncio
async def test_failed_job_with_validation_category_is_not_requeued(self, default_session_factory): async def test_failed_job_with_validation_category_is_not_requeued(self, default_session_factory):
services = ServiceBundle.from_session_factory(default_session_factory) services = ServiceBundle.from_session_factory(default_session_factory)
+19 -1
View File
@@ -167,11 +167,29 @@ class TestWorkerReliabilitySettings:
settings = _make_settings() settings = _make_settings()
assert settings.run_embedded_worker is True assert settings.run_embedded_worker is True
assert settings.worker_max_retries == 0 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_retry_backoff_seconds == 1.0
assert settings.worker_shutdown_grace_seconds == 5.0 assert settings.worker_shutdown_grace_seconds == 5.0
assert settings.worker_poll_interval_seconds == 1.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(): def test_provider_timeout_is_not_capped_at_twenty_seconds():
"""HIGH-03: vision transcription regularly runs past the old le=20.0 ceiling.""" """HIGH-03: vision transcription regularly runs past the old le=20.0 ceiling."""