generated from john/python-template
async provider
This commit is contained in:
@@ -41,6 +41,6 @@ class TranscriptionResult:
|
|||||||
class TranscriptionProvider(Protocol):
|
class TranscriptionProvider(Protocol):
|
||||||
"""Contract every transcription provider adapter must satisfy."""
|
"""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."""
|
"""Transcribe the provided image according to the prompt text."""
|
||||||
...
|
...
|
||||||
|
|||||||
@@ -44,11 +44,11 @@ class OpenRouterTranscriptionProvider:
|
|||||||
"""Return the resolved OpenRouter model slug."""
|
"""Return the resolved OpenRouter model slug."""
|
||||||
return self._model
|
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."""
|
"""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)
|
request = self._build_request(prompt_text=prompt_text, image_bytes=image_bytes, mime_type=mime_type)
|
||||||
try:
|
try:
|
||||||
response = self._client.chat.send(
|
response = await self._client.chat.send_async(
|
||||||
messages=request.messages,
|
messages=request.messages,
|
||||||
model=request.model,
|
model=request.model,
|
||||||
http_referer=request.http_referer,
|
http_referer=request.http_referer,
|
||||||
@@ -63,7 +63,7 @@ class OpenRouterTranscriptionProvider:
|
|||||||
text = self._extract_text(response)
|
text = self._extract_text(response)
|
||||||
model = self._get_optional_attr(response, "model") or self.model
|
model = self._get_optional_attr(response, "model") or self.model
|
||||||
logger.info("OpenRouter transcription completed using model=%s", 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:
|
def _build_request(self, *, prompt_text: str, image_bytes: bytes, mime_type: str) -> OpenRouterRequest:
|
||||||
image_b64 = base64.b64encode(image_bytes).decode("ascii")
|
image_b64 = base64.b64encode(image_bytes).decode("ascii")
|
||||||
|
|||||||
@@ -101,7 +101,7 @@ class TranscriptionService(ServiceBase):
|
|||||||
session: AsyncSession | None = None,
|
session: AsyncSession | None = None,
|
||||||
):
|
):
|
||||||
"""Transcribe a local image using the configured prompt and provider."""
|
"""Transcribe a local image using the configured prompt and provider."""
|
||||||
result = transcribe_document_image(
|
result = await transcribe_document_image(
|
||||||
image_path=image_path,
|
image_path=image_path,
|
||||||
prompt_name=prompt_name,
|
prompt_name=prompt_name,
|
||||||
settings=self.settings,
|
settings=self.settings,
|
||||||
@@ -110,7 +110,7 @@ class TranscriptionService(ServiceBase):
|
|||||||
await self.create_transcript(transcript=result.to_transcript(job_id=job_id), session=session)
|
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,
|
image_path: str | Path,
|
||||||
*,
|
*,
|
||||||
prompt_name: str = DEFAULT_PROMPT_FILE,
|
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)
|
logger.info("Starting transcription for image=%s mime_type=%s", image_path, mime_type)
|
||||||
|
|
||||||
with handle_transcription_errors():
|
with handle_transcription_errors():
|
||||||
result = adapter.transcribe(
|
result = await adapter.transcribe(
|
||||||
prompt_text=prompt_text,
|
prompt_text=prompt_text,
|
||||||
image_bytes=image_bytes,
|
image_bytes=image_bytes,
|
||||||
mime_type=mime_type,
|
mime_type=mime_type,
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -27,11 +27,11 @@ from transcription.errors import format_error_detail
|
|||||||
from transcription.models import Job
|
from transcription.models import Job
|
||||||
from transcription.models import JobStatus
|
from transcription.models import JobStatus
|
||||||
from transcription.models import Transcript
|
from transcription.models import Transcript
|
||||||
from transcription.services.transcription import TranscriptionError
|
|
||||||
from transcription.services.transcription import transcribe_document_image
|
from transcription.services.transcription import transcribe_document_image
|
||||||
|
|
||||||
from .services import ServiceBundle
|
from .services import ServiceBundle
|
||||||
from .services.jobs import JobService
|
from .services.jobs import JobService
|
||||||
|
from .services.workflows import process_job
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -50,12 +50,12 @@ async def queue_consumer_loop(queue: asyncio.Queue[UUID], stop_event: asyncio.Ev
|
|||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def handle_worker_exceptions():
|
def handle_worker_exceptions(operation: str = "worker.loop"):
|
||||||
"""Context manager to log and suppress exceptions in the worker loop."""
|
"""Context manager to log and suppress exceptions in the worker loop."""
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
except Exception as exc:
|
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(
|
logger.exception(
|
||||||
"Worker loop exception error_id=%s category=%s",
|
"Worker loop exception error_id=%s category=%s",
|
||||||
error.error_id,
|
error.error_id,
|
||||||
@@ -70,42 +70,6 @@ async def _get_queue_item(queue: asyncio.Queue[UUID]) -> AsyncGenerator[UUID]:
|
|||||||
queue.task_done()
|
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(
|
async def run_worker_loop(
|
||||||
*,
|
*,
|
||||||
session_factory: async_sessionmaker[AsyncSession] | None = None,
|
session_factory: async_sessionmaker[AsyncSession] | None = None,
|
||||||
|
|||||||
Reference in New Issue
Block a user