From 7b9715b3f139d6d996599beae45ae2421fbfc62f Mon Sep 17 00:00:00 2001 From: zoltan57 <40281233+zoltan57@users.noreply.github.com> Date: Mon, 17 Aug 2026 16:26:31 -0500 Subject: [PATCH] V4.6 Phase 3: worker and provider reliability Claim jobs atomically [CRIT-01] - Replace JobService.read_next_queued_job with claim_next_queued_job, which selects and transitions QUEUED -> PROCESSING inside one transaction. The old read-then-write sequence left a window in which two workers could observe the same QUEUED row. - Add the missing .limit(1). The poll previously ordered the entire queued set and discarded all but the first row. - Drop the eager loads from the hot poll entirely. They were pure waste: process_queued_job immediately re-reads the job through read_job with the relationships it actually needs. - Guard the row with with_for_update(skip_locked=True) on PostgreSQL so the claim stays correct once more than one worker exists. On SQLite the claim is a bounded single-writer transaction. - Correct the comment at the remaining direct-call claim site, which described the hazard rather than the guarantee. Reuse the provider connection [HIGH-02] - Build the ServiceBundle once per worker loop instead of once per job, and close it at loop shutdown. Every job previously constructed a new SourceService, and with it a new provider adapter and a new httpx.AsyncClient, paying a full TLS handshake per page and discarding the connection pool. - process_next_queued_job now accepts an optional caller-owned bundle and only closes bundles it created itself. Uncap the provider timeout [HIGH-03] - Remove le=20.0 from worker_provider_timeout_seconds. The cap equalled the default, so the ceiling could never be raised, and dense-page vision transcription routinely needs longer. Default raised to 180s. - Pass an explicit httpx.Timeout to the OpenRouter AsyncClient. httpx defaults every phase to 5 seconds, so the real read budget was 5s regardless of the configured value; the outer asyncio.wait_for could never be the binding constraint. Connect stays at 10s. Tighten the provider boundary [MED-03] - Declare model, current_request_manifest, current_transport_evidence, and aclose on the TranscriptionProvider Protocol. - Delete the per-call inspect.signature(adapter.transcribe).parameters reflection and the untyped kwargs dict it fed. The Protocol had declared requested_model all along, so the reflection was dead defensive weight on the hot path. - Replace the three getattr probes for aclose and the evidence attributes with direct typed access. Deduplicate bundle construction [MED-06] - Add ServiceBundle.from_session_factory and ServiceBundle.aclose, replacing the duplicated four-service instantiation blocks in app.py and worker.py. - _recover_stale_processing_jobs now uses the bundle built moments earlier instead of constructing a second JobService. Tests - Claiming returns the oldest job, marks it PROCESSING, never hands the same job out twice, and emits exactly one unadorned SELECT carrying LIMIT and no JOIN. - The worker loop threads one bundle through consecutive jobs and closes it once at shutdown; a caller-owned bundle is left open. - Settings accepts a timeout above 20 seconds and still rejects zero. - The OpenRouter client's read, write, and pool timeouts track the configured budget rather than the httpx default. Note: .env in this checkout still pins WORKER_PROVIDER_TIMEOUT_SECONDS=20 and should be raised to pick up this fix. Verified: 268 passed, 4 skipped; ruff check clean. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .env.example | 4 +- src/transcription/app.py | 14 +--- src/transcription/config.py | 4 +- src/transcription/providers/base.py | 19 +++++ src/transcription/providers/openrouter.py | 10 ++- src/transcription/services/__init__.py | 26 +++++++ src/transcription/services/jobs.py | 27 +++++-- src/transcription/services/sources.py | 29 +++---- src/transcription/services/workflows.py | 23 +++--- src/transcription/worker.py | 92 +++++++++++++---------- tests/services/test_job_service.py | 42 ++++++++++- tests/test_config.py | 25 ++++++ tests/test_worker.py | 70 +++++++++++++++-- 13 files changed, 287 insertions(+), 98 deletions(-) diff --git a/.env.example b/.env.example index a66cae7..6f30590 100644 --- a/.env.example +++ b/.env.example @@ -49,8 +49,8 @@ PROMPT_DIR="./prompts" # --- worker reliability --- WORKER_MAX_RETRIES=0 WORKER_RETRY_BACKOFF_SECONDS=0 -# WORKER_PROVIDER_TIMEOUT_SECONDS=[0-20] -WORKER_PROVIDER_TIMEOUT_SECONDS=20 +# WORKER_PROVIDER_TIMEOUT_SECONDS=180 +WORKER_PROVIDER_TIMEOUT_SECONDS=180 WORKER_MIN_TRANSCRIPTION_CHARS=0 WORKER_MIN_TRANSCRIPTION_LINES=0 WORKER_FAIL_ON_FINISH_REASON_LENGTH=false diff --git a/src/transcription/app.py b/src/transcription/app.py index cd4cc31..e770893 100644 --- a/src/transcription/app.py +++ b/src/transcription/app.py @@ -25,10 +25,6 @@ from .db import create_all from .db import dispose_database_runtime from .db import initialize_database_runtime from .services import ServiceBundle -from .services.documents import DocumentService -from .services.jobs import JobService -from .services.people import PeopleService -from .services.sources import SourceService from .ui import register_pages from .worker import worker_consumer_lifespan @@ -42,12 +38,7 @@ async def _lifespan(app: FastAPI): app.state.settings = settings app.state.runtime = initialize_database_runtime(settings=settings) session_factory = app.state.runtime.session_factory - app.state.services = ServiceBundle( - documents=DocumentService(session_factory=session_factory, settings=settings), - sources=SourceService(session_factory=session_factory, settings=settings), - jobs=JobService(session_factory=session_factory, settings=settings), - people=PeopleService(session_factory=session_factory, settings=settings), - ) + app.state.services = ServiceBundle.from_session_factory(session_factory, settings=settings) if settings.should_bootstrap_schema: await create_all(engine=app.state.runtime.engine) @@ -78,8 +69,7 @@ async def _recover_stale_processing_jobs(app: FastAPI) -> None: """ settings = app.state.settings stale_before = datetime.now(UTC) - timedelta(seconds=settings.worker_provider_timeout_seconds) - job_service = JobService(session_factory=app.state.runtime.session_factory) - recovered = await job_service.requeue_stale_processing_jobs(stale_before=stale_before) + recovered = await app.state.services.jobs.requeue_stale_processing_jobs(stale_before=stale_before) if recovered > 0: logger.warning("Recovered %s stale processing job(s) at startup", recovered) diff --git a/src/transcription/config.py b/src/transcription/config.py index 7d3514f..aff7f74 100644 --- a/src/transcription/config.py +++ b/src/transcription/config.py @@ -106,7 +106,9 @@ class Settings(BaseSettings): # --- worker reliability --- worker_max_retries: int = Field(default=0, ge=0) - worker_provider_timeout_seconds: float = Field(default=20.0, gt=0.0, le=20.0) + # 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=180.0, gt=0.0) worker_min_transcription_chars: int = Field(default=0, ge=0) worker_min_transcription_lines: int = Field(default=0, ge=0) worker_fail_on_finish_reason_length: bool = False diff --git a/src/transcription/providers/base.py b/src/transcription/providers/base.py index c2bac4a..d109a1b 100644 --- a/src/transcription/providers/base.py +++ b/src/transcription/providers/base.py @@ -102,6 +102,21 @@ class TranscriptionResult(BaseModel): class TranscriptionProvider(Protocol): """Contract every transcription provider adapter must satisfy.""" + @property + def model(self) -> str: + """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, *, @@ -115,3 +130,7 @@ class TranscriptionProvider(Protocol): ) -> TranscriptionResult: """Transcribe the provided image according to the prompt text.""" ... + + async def aclose(self) -> None: + """Release any pooled network resources held by the adapter.""" + ... diff --git a/src/transcription/providers/openrouter.py b/src/transcription/providers/openrouter.py index 1b7cadc..2b5d72a 100644 --- a/src/transcription/providers/openrouter.py +++ b/src/transcription/providers/openrouter.py @@ -195,7 +195,15 @@ class OpenRouterTranscriptionProvider: self._current_request_manifest: RequestManifest | None = None self._current_transport_evidence: TransportEvidence | None = None if client is None: - self._capturing_client = _CapturingAsyncClient(async_client or httpx.AsyncClient(follow_redirects=True)) + # httpx defaults every phase to 5s, which silently caps provider calls far + # below worker_provider_timeout_seconds. Track the configured budget instead. + timeout = httpx.Timeout( + self._settings.worker_provider_timeout_seconds, + connect=10.0, + ) + self._capturing_client = _CapturingAsyncClient( + async_client or httpx.AsyncClient(follow_redirects=True, timeout=timeout) + ) client = OpenRouter( api_key=self._settings.openrouter_api_key.get_secret_value(), async_client=self._capturing_client, diff --git a/src/transcription/services/__init__.py b/src/transcription/services/__init__.py index 7ed9a35..256f49a 100644 --- a/src/transcription/services/__init__.py +++ b/src/transcription/services/__init__.py @@ -2,7 +2,12 @@ from dataclasses import dataclass from dataclasses import field +from typing import Self +from sqlalchemy.ext.asyncio import async_sessionmaker +from sqlmodel.ext.asyncio.session import AsyncSession + +from ..config import Settings from .documents import DocumentService from .jobs import JobService from .people import PeopleService @@ -20,3 +25,24 @@ class ServiceBundle: sources: SourceService = field(default_factory=SourceService) jobs: JobService = field(default_factory=JobService) people: PeopleService = field(default_factory=PeopleService) + + @classmethod + def from_session_factory( + cls, + session_factory: async_sessionmaker[AsyncSession] | None = None, + *, + settings: Settings | None = None, + ) -> Self: + """Build a bundle whose services all share one session factory and settings.""" + if session_factory is None: + return cls() + return cls( + documents=DocumentService(session_factory=session_factory, settings=settings), + sources=SourceService(session_factory=session_factory, settings=settings), + jobs=JobService(session_factory=session_factory, settings=settings), + people=PeopleService(session_factory=session_factory, settings=settings), + ) + + async def aclose(self) -> None: + """Release provider resources held by the bundle.""" + await self.sources.aclose() diff --git a/src/transcription/services/jobs.py b/src/transcription/services/jobs.py index 84fe9ad..c1789bb 100644 --- a/src/transcription/services/jobs.py +++ b/src/transcription/services/jobs.py @@ -168,24 +168,37 @@ class JobService(ServiceBase): await self._finalize(session=_session, caller_session=session, refresh=(job,)) return job - async def read_next_queued_job( + async def claim_next_queued_job( self, *, session: AsyncSession | None = None, ) -> Job | None: - """Read the next queued job ordered by creation time.""" + """Atomically claim the oldest queued job by transitioning it to PROCESSING. + + The selection is deliberately unadorned: no eager loads are applied to the + hot poll, because callers re-read the claimed job with the relationships + they actually need. On PostgreSQL the row is locked with ``SKIP LOCKED`` so + concurrent workers never contend for the same job. + """ async with self._session_scope(session) as _session: query = ( select(Job) - .options( - selectinload(Job.document), # pyright: ignore[reportArgumentType] - selectinload(Job.job_sources).selectinload(JobSource.source), # pyright: ignore[reportArgumentType] - ) .where(Job.status == JobStatus.QUEUED) # Break ties by id so "next" is stable when two rows share close timestamps. .order_by(Job.date_created, Job.id) # pyright: ignore[reportArgumentType] + .limit(1) ) - return (await _session.exec(query)).first() + if _session.get_bind().dialect.name == "postgresql": + query = query.with_for_update(skip_locked=True) + + job = (await _session.exec(query)).first() + if job is None: + return None + + job.status = JobStatus.PROCESSING + job.date_updated = datetime.now(UTC) + await self._finalize(session=_session, caller_session=session, refresh=(job,)) + return job async def requeue_stale_processing_jobs( self, diff --git a/src/transcription/services/sources.py b/src/transcription/services/sources.py index 1b349c8..273395d 100644 --- a/src/transcription/services/sources.py +++ b/src/transcription/services/sources.py @@ -4,7 +4,6 @@ from __future__ import annotations import base64 import hashlib -import inspect import logging import os from collections.abc import Sequence @@ -1142,8 +1141,8 @@ class SourceService(ServiceBase): def _resolve_transcript_model(*, provider: TranscriptionProvider, settings: Settings) -> str: - provider_model = getattr(provider, "model", None) - if isinstance(provider_model, str) and provider_model.strip(): + provider_model = provider.model + if provider_model and provider_model.strip(): return provider_model if settings.provider_model and settings.provider_model.strip(): @@ -1226,22 +1225,18 @@ async def transcribe_document_image( try: with handle_transcription_errors(): - transcribe_kwargs = { - "prompt_text": prompt_execution.user_prompt, - "image_bytes": image_bytes, - "mime_type": mime_type, - "temperature": prompt_execution.temperature, - "top_p": prompt_execution.top_p, - "source_reference": source_reference, - } - if "requested_model" in inspect.signature(adapter.transcribe).parameters: - transcribe_kwargs["requested_model"] = requested_model - result = await adapter.transcribe(**transcribe_kwargs) + result = await adapter.transcribe( + prompt_text=prompt_execution.user_prompt, + image_bytes=image_bytes, + mime_type=mime_type, + temperature=prompt_execution.temperature, + top_p=prompt_execution.top_p, + source_reference=source_reference, + requested_model=requested_model, + ) finally: if owns_adapter: - close = getattr(adapter, "aclose", None) - if close is not None: - await close() + await adapter.aclose() logger.info("Transcription completed for image=%s provider=%s", image_path, result.provider) return TranscriptionResult( text=result.text, diff --git a/src/transcription/services/workflows.py b/src/transcription/services/workflows.py index 00efe5d..f69702c 100644 --- a/src/transcription/services/workflows.py +++ b/src/transcription/services/workflows.py @@ -185,13 +185,15 @@ async def process_queued_job( # noqa: PLR0915 logger.warning(f"Job {job.id} is not queued. Current status: {job.status}") return - # Transaction A: claim job for processing. + # Transaction A: claim job for processing. Reached only when a caller hands us a + # still-QUEUED job directly; the worker path already claimed it atomically in + # JobService.claim_next_queued_job. if current_status == JobStatus.QUEUED: if session is None: job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING) else: - # If we're sharing the session, need to make sure setting the Job to PROCESSING is committed before we start - # the transcription, otherwise other workers may see the job as still QUEUED and try to process it. + # Commit the PROCESSING transition before transcription starts so the claim + # is durable and visible to any other worker before the long provider call. job = await services.jobs.mark_job_status(job.id, JobStatus.PROCESSING, session=session) await session.commit() @@ -294,12 +296,8 @@ async def process_queued_job( # noqa: PLR0915 0, int((asyncio.get_running_loop().time() - monotonic_started_at) * 1000), ), - request_manifest=getattr(services.sources.provider, "current_request_manifest", None), - transport_evidence=getattr( - services.sources.provider, - "current_transport_evidence", - None, - ), + request_manifest=services.sources.provider.current_request_manifest, + transport_evidence=services.sources.provider.current_transport_evidence, failure_phase="local_timeout", model_input_artifact_id=( provider_input.derivative_id if provider_input is not None else None @@ -417,11 +415,16 @@ async def process_next_queued_job( session: AsyncSession | None = None, ) -> bool: """Process the next queued job if one exists.""" - job = await services.jobs.read_next_queued_job(session=session) + job = await services.jobs.claim_next_queued_job(session=session) if job is None: return False + # The claim must be durable before the provider call starts, otherwise another + # worker could observe the job as still QUEUED and process it a second time. + if session is not None: + await session.commit() + await advance_job(job=job, services=services, settings=settings, session=session) return True diff --git a/src/transcription/worker.py b/src/transcription/worker.py index 798111b..d79b060 100644 --- a/src/transcription/worker.py +++ b/src/transcription/worker.py @@ -19,10 +19,6 @@ from transcription.errors import AppError from transcription.errors import classify_unexpected_error from .services import ServiceBundle -from .services.documents import DocumentService -from .services.jobs import JobService -from .services.people import PeopleService -from .services.sources import SourceService from .services.workflows import process_next_queued_job as process_next_queued_job_workflow logger = logging.getLogger(__name__) @@ -121,56 +117,74 @@ async def run_worker_loop( If wake_event is provided, signal activity wakes the loop immediately while timeout-based wakeups preserve current polling behavior. + + The service bundle — and with it the provider's pooled HTTP client — is built + once for the lifetime of the loop, so consecutive jobs reuse one connection + instead of paying a fresh TLS handshake each time. """ - while True: - if stop_event is not None and stop_event.is_set(): - logger.info("Worker stop event received") - return - - if wake_event is not None: - with suppress(TimeoutError): - await asyncio.wait_for(wake_event.wait(), timeout=poll_interval_seconds) - wake_event.clear() - - processed_any = False + services = ServiceBundle.from_session_factory(session_factory) + try: while True: - with handle_worker_exceptions(operation="worker.process_next_queued_job"): - processed = await process_next_queued_job(session_factory=session_factory) - if not processed: - break - processed_any = True - continue + if stop_event is not None and stop_event.is_set(): + logger.info("Worker stop event received") + return - break + if wake_event is not None: + with suppress(TimeoutError): + await asyncio.wait_for(wake_event.wait(), timeout=poll_interval_seconds) + wake_event.clear() - if wake_event is None and not processed_any: - await asyncio.sleep(poll_interval_seconds) + processed_any = False + while True: + with handle_worker_exceptions(operation="worker.process_next_queued_job"): + processed = await process_next_queued_job( + session_factory=session_factory, + services=services, + ) + if not processed: + break + processed_any = True + continue + + break + + if wake_event is None and not processed_any: + await asyncio.sleep(poll_interval_seconds) + finally: + await services.aclose() async def process_next_queued_job( *, session: AsyncSession | None = None, session_factory: async_sessionmaker[AsyncSession] | None = None, + services: ServiceBundle | None = None, ) -> bool: """Process the next queued job and persist terminal outcome. Returns True when a job was processed, False when no queued job exists. + + When ``services`` is supplied the caller owns its lifecycle; otherwise a + bundle is created and closed here. """ - if session_factory is None: - services = ServiceBundle() - else: - services = ServiceBundle( - documents=DocumentService(session_factory=session_factory), - sources=SourceService(session_factory=session_factory), - jobs=JobService(session_factory=session_factory), - people=PeopleService(session_factory=session_factory), - ) + if services is not None: + return await _process_next_queued_job(services=services, session=session, session_factory=session_factory) + owned = ServiceBundle.from_session_factory(session_factory) try: - if session is None: - async with session_scope(session_factory=session_factory) as local_session: - return await process_next_queued_job_workflow(services=services, session=local_session) - - return await process_next_queued_job_workflow(services=services, session=session) + return await _process_next_queued_job(services=owned, session=session, session_factory=session_factory) finally: - await services.sources.aclose() + await owned.aclose() + + +async def _process_next_queued_job( + *, + services: ServiceBundle, + session: AsyncSession | None, + session_factory: async_sessionmaker[AsyncSession] | None, +) -> bool: + if session is None: + async with session_scope(session_factory=session_factory) as local_session: + return await process_next_queued_job_workflow(services=services, session=local_session) + + return await process_next_queued_job_workflow(services=services, session=session) diff --git a/tests/services/test_job_service.py b/tests/services/test_job_service.py index 966cfa2..648bbcf 100644 --- a/tests/services/test_job_service.py +++ b/tests/services/test_job_service.py @@ -4,6 +4,7 @@ from datetime import timedelta from uuid import uuid4 import pytest +from sqlalchemy import event from transcription.db.models import Document from transcription.db.models import Job @@ -101,7 +102,7 @@ class TestJobService: assert result[0].id == job.id @pytest.mark.asyncio - async def test_read_next_queued_job_orders_by_created_date( + async def test_claim_next_queued_job_claims_oldest_and_marks_processing( self, job_service: JobService, document_service: DocumentService, @@ -119,9 +120,42 @@ class TestJobService: await job_service.create_job(job=first) await job_service.create_job(job=second) - next_job = await job_service.read_next_queued_job() - assert next_job is not None - assert next_job.id == first.id + claimed = await job_service.claim_next_queued_job() + assert claimed is not None + assert claimed.id == first.id + assert claimed.status == JobStatus.PROCESSING + + # The claim is exclusive: the same job is never handed out twice. + next_claim = await job_service.claim_next_queued_job() + assert next_claim is not None + assert next_claim.id == second.id + + assert await job_service.claim_next_queued_job() is None + + @pytest.mark.asyncio + async def test_claim_next_queued_job_emits_a_bounded_unadorned_query( + self, + job_service: JobService, + ): + """CRIT-01: the hot poll must not select a subgraph or scan the queue.""" + statements: list[str] = [] + + async with job_service._session_scope() as session: + bind = session.get_bind() + + def capture(_conn, _cursor, statement, *_rest): + statements.append(statement) + + event.listen(bind, "before_cursor_execute", capture) + try: + await job_service.claim_next_queued_job(session=session) + finally: + event.remove(bind, "before_cursor_execute", capture) + + selects = [item for item in statements if item.lstrip().upper().startswith("SELECT")] + assert len(selects) == 1, selects + assert "LIMIT" in selects[0].upper() + assert "JOIN" not in selects[0].upper() @pytest.mark.asyncio async def test_create_job_persists_provider_and_model( diff --git a/tests/test_config.py b/tests/test_config.py index e8d9c8f..ce4e39a 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -149,3 +149,28 @@ class TestWorkerReliabilitySettings: """worker retry settings default to no retries.""" settings = _make_settings() assert settings.worker_max_retries == 0 + + +def test_provider_timeout_is_not_capped_at_twenty_seconds(): + """HIGH-03: vision transcription regularly runs past the old le=20.0 ceiling.""" + settings = Settings(openrouter_api_key="test-key", worker_provider_timeout_seconds=300.0) + assert settings.worker_provider_timeout_seconds == 300.0 + + +def test_provider_timeout_must_still_be_positive(): + with pytest.raises(ValidationError): + Settings(openrouter_api_key="test-key", worker_provider_timeout_seconds=0.0) + + +def test_openrouter_client_timeout_tracks_the_configured_budget(): + """HIGH-03: httpx defaults every phase to 5s, silently capping the provider call.""" + from transcription.providers.openrouter import OpenRouterTranscriptionProvider + + settings = Settings(openrouter_api_key="test-key", worker_provider_timeout_seconds=123.0) + provider = OpenRouterTranscriptionProvider(settings=settings) + timeout = provider._capturing_client._client.timeout + + assert timeout.read == 123.0 + assert timeout.write == 123.0 + assert timeout.pool == 123.0 + assert timeout.connect == 10.0 diff --git a/tests/test_worker.py b/tests/test_worker.py index 6d5127b..e12fdf9 100644 --- a/tests/test_worker.py +++ b/tests/test_worker.py @@ -3,6 +3,7 @@ import logging import pytest +from transcription.services import ServiceBundle from transcription.worker import process_next_queued_job from transcription.worker import run_worker_loop @@ -12,9 +13,9 @@ async def test_run_worker_loop_survives_process_next_exception(monkeypatch, capl calls = 0 stop_event = asyncio.Event() - async def _fake_process_next_queued_job(*, session=None, session_factory=None): + async def _fake_process_next_queued_job(*, session=None, session_factory=None, services=None): nonlocal calls - _ = (session, session_factory) + _ = (session, session_factory, services) calls += 1 if calls == 1: raise RuntimeError("boom") @@ -31,7 +32,10 @@ async def test_run_worker_loop_survives_process_next_exception(monkeypatch, capl @pytest.mark.asyncio -async def test_process_next_closes_initialized_provider(monkeypatch): +async def test_run_worker_loop_reuses_one_bundle_across_jobs(monkeypatch): + """HIGH-02: the provider client is built once per loop, not once per job.""" + stop_event = asyncio.Event() + seen: list[object] = [] closed = False class _Sources: @@ -39,9 +43,43 @@ async def test_process_next_closes_initialized_provider(monkeypatch): nonlocal closed closed = True - services = type("_Services", (), {"sources": _Sources()})() + bundle = ServiceBundle(sources=_Sources()) # type: ignore[arg-type] + monkeypatch.setattr( + "transcription.worker.ServiceBundle.from_session_factory", + classmethod(lambda _cls, _factory=None, **_kwargs: bundle), + ) - monkeypatch.setattr("transcription.worker.ServiceBundle", lambda: services) + async def _fake_process_next_queued_job(*, session=None, session_factory=None, services=None): + _ = (session, session_factory) + seen.append(services) + if len(seen) >= 3: + stop_event.set() + return False + return True + + monkeypatch.setattr("transcription.worker.process_next_queued_job", _fake_process_next_queued_job) + + await run_worker_loop(stop_event=stop_event, poll_interval_seconds=0) + + assert len(seen) == 3 + assert all(item is bundle for item in seen) + assert closed is True + + +@pytest.mark.asyncio +async def test_process_next_closes_provider_for_the_bundle_it_owns(monkeypatch): + closed = False + + class _Sources: + async def aclose(self): + nonlocal closed + closed = True + + bundle = ServiceBundle(sources=_Sources()) # type: ignore[arg-type] + monkeypatch.setattr( + "transcription.worker.ServiceBundle.from_session_factory", + classmethod(lambda _cls, _factory=None, **_kwargs: bundle), + ) async def _no_job(*, services, session): _ = (services, session) @@ -51,3 +89,25 @@ async def test_process_next_closes_initialized_provider(monkeypatch): assert await process_next_queued_job() is False assert closed is True + + +@pytest.mark.asyncio +async def test_process_next_leaves_a_caller_owned_bundle_open(monkeypatch): + """A bundle passed in belongs to the caller and must outlive one job.""" + closed = False + + class _Sources: + async def aclose(self): + nonlocal closed + closed = True + + bundle = ServiceBundle(sources=_Sources()) # type: ignore[arg-type] + + async def _no_job(*, services, session): + _ = (services, session) + return False + + monkeypatch.setattr("transcription.worker.process_next_queued_job_workflow", _no_job) + + assert await process_next_queued_job(services=bundle) is False + assert closed is False