diff --git a/.env.production.example b/.env.production.example index cf01413..4ae6df6 100644 --- a/.env.production.example +++ b/.env.production.example @@ -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 diff --git a/.github/instructions/providers.instructions.md b/.github/instructions/providers.instructions.md index 2e60a6b..02d4b70 100644 --- a/.github/instructions/providers.instructions.md +++ b/.github/instructions/providers.instructions.md @@ -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 diff --git a/src/transcription/config.py b/src/transcription/config.py index 8d0bb14..a5cf204 100644 --- a/src/transcription/config.py +++ b/src/transcription/config.py @@ -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.""" diff --git a/src/transcription/providers/__init__.py b/src/transcription/providers/__init__.py index ff917c7..cab7ee3 100644 --- a/src/transcription/providers/__init__.py +++ b/src/transcription/providers/__init__.py @@ -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", diff --git a/src/transcription/providers/base.py b/src/transcription/providers/base.py index d109a1b..f12b1e9 100644 --- a/src/transcription/providers/base.py +++ b/src/transcription/providers/base.py @@ -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: diff --git a/src/transcription/providers/openrouter.py b/src/transcription/providers/openrouter.py index 768451d..69fde8e 100644 --- a/src/transcription/providers/openrouter.py +++ b/src/transcription/providers/openrouter.py @@ -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,79 +265,87 @@ 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: - 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 - if isinstance(exc, openrouter_errors.UnauthorizedResponseError): - raise ProviderAuthError( - "OpenRouter authentication failed", + 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(transport_capture) + if evidence_capture is not None: + evidence_capture.transport_evidence = transport + 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, transport_evidence=transport, - failure_phase="http_response" if transport.response_received else "connection", + failure_phase=failure_phase, ) from exc - failure_phase = ( - "response_validation" - if isinstance(exc, openrouter_errors.ResponseValidationError) - else "http_response" - if transport.response_received - else "connection" + + 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) + 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( - self._transport_error_message(transport), - request_manifest=manifest, - transport_evidence=transport, - 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, - ) + 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"), diff --git a/src/transcription/services/jobs.py b/src/transcription/services/jobs.py index 822a236..82b99fe 100644 --- a/src/transcription/services/jobs.py +++ b/src/transcription/services/jobs.py @@ -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, *, diff --git a/src/transcription/services/sources.py b/src/transcription/services/sources.py index eb8513c..1208792 100644 --- a/src/transcription/services/sources.py +++ b/src/transcription/services/sources.py @@ -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: diff --git a/src/transcription/services/workflows.py b/src/transcription/services/workflows.py index da98c04..a65c7c3 100644 --- a/src/transcription/services/workflows.py +++ b/src/transcription/services/workflows.py @@ -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() diff --git a/tests/integration/test_pipeline_flow.py b/tests/integration/test_pipeline_flow.py index 697c040..a4a7f9a 100644 --- a/tests/integration/test_pipeline_flow.py +++ b/tests/integration/test_pipeline_flow.py @@ -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") diff --git a/tests/providers/test_openrouter.py b/tests/providers/test_openrouter.py index 592f8f4..0848f05 100644 --- a/tests/providers/test_openrouter.py +++ b/tests/providers/test_openrouter.py @@ -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"}}' diff --git a/tests/services/test_workflows_reliability.py b/tests/services/test_workflows_reliability.py index 5b2a128..b8f5518 100644 --- a/tests/services/test_workflows_reliability.py +++ b/tests/services/test_workflows_reliability.py @@ -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) diff --git a/tests/test_config.py b/tests/test_config.py index 7486d71..2ae3294 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -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."""