From f1fb45e0d234bbe65b4edd2ce92b79d7ca2ae478 Mon Sep 17 00:00:00 2001 From: John Lancaster <32917998+jsl12@users.noreply.github.com> Date: Sun, 28 Jun 2026 09:35:14 -0500 Subject: [PATCH] async provider --- src/transcription/providers/base.py | 2 +- src/transcription/providers/openrouter.py | 6 +- src/transcription/services/transcription.py | 6 +- src/transcription/services/workflows.py | 68 +++++++++++++++++++++ src/transcription/worker.py | 42 +------------ 5 files changed, 78 insertions(+), 46 deletions(-) create mode 100644 src/transcription/services/workflows.py diff --git a/src/transcription/providers/base.py b/src/transcription/providers/base.py index 86a6e25..aa69a73 100644 --- a/src/transcription/providers/base.py +++ b/src/transcription/providers/base.py @@ -41,6 +41,6 @@ class TranscriptionResult: class TranscriptionProvider(Protocol): """Contract every transcription provider adapter must satisfy.""" - def transcribe(self, *, prompt_text: str, image_bytes: bytes, mime_type: str) -> TranscriptionResult: + async def transcribe(self, *, prompt_text: str, image_bytes: bytes, mime_type: str) -> TranscriptionResult: """Transcribe the provided image according to the prompt text.""" ... diff --git a/src/transcription/providers/openrouter.py b/src/transcription/providers/openrouter.py index e8ec702..7382a0a 100644 --- a/src/transcription/providers/openrouter.py +++ b/src/transcription/providers/openrouter.py @@ -44,11 +44,11 @@ class OpenRouterTranscriptionProvider: """Return the resolved OpenRouter model slug.""" return self._model - def transcribe(self, *, prompt_text: str, image_bytes: bytes, mime_type: str) -> TranscriptionResult: + async def transcribe(self, *, prompt_text: str, image_bytes: bytes, mime_type: str) -> TranscriptionResult: """Send prompt + image to OpenRouter and return normalized text output.""" request = self._build_request(prompt_text=prompt_text, image_bytes=image_bytes, mime_type=mime_type) try: - response = self._client.chat.send( + response = await self._client.chat.send_async( messages=request.messages, model=request.model, http_referer=request.http_referer, @@ -63,7 +63,7 @@ class OpenRouterTranscriptionProvider: text = self._extract_text(response) model = self._get_optional_attr(response, "model") or self.model logger.info("OpenRouter transcription completed using model=%s", model) - return TranscriptionResult(text=text, provider="openrouter", model=model) + return TranscriptionResult(text=text, provider="openrouter", prompt_name="", model=model) def _build_request(self, *, prompt_text: str, image_bytes: bytes, mime_type: str) -> OpenRouterRequest: image_b64 = base64.b64encode(image_bytes).decode("ascii") diff --git a/src/transcription/services/transcription.py b/src/transcription/services/transcription.py index 18e5c67..6a6e1c8 100644 --- a/src/transcription/services/transcription.py +++ b/src/transcription/services/transcription.py @@ -101,7 +101,7 @@ class TranscriptionService(ServiceBase): session: AsyncSession | None = None, ): """Transcribe a local image using the configured prompt and provider.""" - result = transcribe_document_image( + result = await transcribe_document_image( image_path=image_path, prompt_name=prompt_name, settings=self.settings, @@ -110,7 +110,7 @@ class TranscriptionService(ServiceBase): await self.create_transcript(transcript=result.to_transcript(job_id=job_id), session=session) -def transcribe_document_image( +async def transcribe_document_image( image_path: str | Path, *, prompt_name: str = DEFAULT_PROMPT_FILE, @@ -126,7 +126,7 @@ def transcribe_document_image( logger.info("Starting transcription for image=%s mime_type=%s", image_path, mime_type) with handle_transcription_errors(): - result = adapter.transcribe( + result = await adapter.transcribe( prompt_text=prompt_text, image_bytes=image_bytes, mime_type=mime_type, diff --git a/src/transcription/services/workflows.py b/src/transcription/services/workflows.py new file mode 100644 index 0000000..5564cc0 --- /dev/null +++ b/src/transcription/services/workflows.py @@ -0,0 +1,68 @@ +import logging +from functools import partial + +from sqlmodel.ext.asyncio.session import AsyncSession + +from ..config import Settings +from ..config import get_settings +from ..models import Job +from ..models import JobStatus +from . import ServiceBundle +from .transcription import TranscriptionError + +logger = logging.getLogger(__name__) + + +async def advance_job( + job: Job, + services: ServiceBundle, + settings: Settings | None = None, + session: AsyncSession | None = None, +) -> Job | None: + """Process a single job using the service bundle.""" + settings = settings or get_settings() + match job.status: + case JobStatus.QUEUED: + processed_job = await _process_queued(job, services, settings=settings, session=session) + case JobStatus.FAILED: + if job.retry_count < settings.worker_max_retries: + job.retry_count += 1 + job.status = JobStatus.QUEUED + logger.info(f"Queuing job {job.id} for retry {job.retry_count}/{settings.worker_max_retries}") + else: + logger.error(f"Job {job.id} has failed and reached max retries.") + return + case _: + return + + +async def _process_queued( + job: Job, + services: ServiceBundle, + settings: Settings | None = None, + session: AsyncSession | None = None, +) -> Job | None: + """Process a queued job using the service bundle.""" + settings = settings or get_settings() + job_id = job.id + updater = partial(services.jobs.mark_job_status, job_id, session=session) + + if job.status != JobStatus.QUEUED: + logger.warning(f"Job {job.id} is not queued. Current status: {job.status}") + return + + job = await updater(JobStatus.PROCESSING) + try: + await services.transcriptions.transcribe_document(job.document.file_path, job.id) + await updater(JobStatus.TRANSCRIBED) + except TranscriptionError as exc: + job.status = JobStatus.FAILED + job.error_message = str(exc) + logger.error( + "Job failed operation=worker.process_job job_id=%s document_id=%s error=%s", + job.id, + job.document.id, + exc, + ) + finally: + return job diff --git a/src/transcription/worker.py b/src/transcription/worker.py index 33a0930..dd10347 100644 --- a/src/transcription/worker.py +++ b/src/transcription/worker.py @@ -27,11 +27,11 @@ from transcription.errors import format_error_detail from transcription.models import Job from transcription.models import JobStatus from transcription.models import Transcript -from transcription.services.transcription import TranscriptionError from transcription.services.transcription import transcribe_document_image from .services import ServiceBundle from .services.jobs import JobService +from .services.workflows import process_job logger = logging.getLogger(__name__) @@ -50,12 +50,12 @@ async def queue_consumer_loop(queue: asyncio.Queue[UUID], stop_event: asyncio.Ev @contextmanager -def handle_worker_exceptions(): +def handle_worker_exceptions(operation: str = "worker.loop"): """Context manager to log and suppress exceptions in the worker loop.""" try: yield except Exception as exc: - error = exc if isinstance(exc, AppError) else classify_unexpected_error(exc, operation="worker.loop") + error = exc if isinstance(exc, AppError) else classify_unexpected_error(exc, operation=operation) logger.exception( "Worker loop exception error_id=%s category=%s", error.error_id, @@ -70,42 +70,6 @@ async def _get_queue_item(queue: asyncio.Queue[UUID]) -> AsyncGenerator[UUID]: queue.task_done() -async def process_job( - job: Job, - services: ServiceBundle, - settings: Settings | None = None, - session: AsyncSession | None = None, -) -> Job | None: - """Process a single job using the service bundle.""" - settings = settings or get_settings() - match job.status: - case JobStatus.QUEUED: - job.status = JobStatus.PROCESSING - try: - # await services.transcriptions.transcribe_image(job.document.file_path) - job.status = JobStatus.TRANSCRIBED - except TranscriptionError as exc: - job.status = JobStatus.FAILED - job.error_message = str(exc) - logger.error( - "Job failed operation=worker.process_job job_id=%s document_id=%s error=%s", - job.id, - job.document.id, - exc, - ) - case JobStatus.FAILED: - if job.retry_count < settings.worker_max_retries: - job.retry_count += 1 - job.status = JobStatus.QUEUED - logger.info(f"Queuing job {job.id} for retry {job.retry_count}/{settings.worker_max_retries}") - else: - logger.error(f"Job {job.id} has failed and reached max retries.") - return - case _: - return - return await services.jobs.update_job(job, session=session) - - async def run_worker_loop( *, session_factory: async_sessionmaker[AsyncSession] | None = None,