diff --git a/src/transcription/app.py b/src/transcription/app.py index f437f2d..ecd4198 100644 --- a/src/transcription/app.py +++ b/src/transcription/app.py @@ -2,9 +2,8 @@ from __future__ import annotations -import asyncio +from contextlib import AsyncExitStack from contextlib import asynccontextmanager -from contextlib import suppress from fastapi import FastAPI @@ -17,42 +16,14 @@ from .db import dispose_database_runtime from .db import initialize_database_runtime from .services import ServiceBundle from .ui import register_pages -from .worker import run_worker_loop +from .worker import WorkerNotifier +from .worker import resolve_worker_notifier +from .worker import worker_consumer_lifespan -def _start_worker(app: FastAPI) -> None: - stop_event = asyncio.Event() - wake_queue: asyncio.Queue[None] = asyncio.Queue() - worker_task = asyncio.create_task( - run_worker_loop( - session_factory=app.state.runtime.session_factory, - stop_event=stop_event, - wake_queue=wake_queue, - poll_interval_seconds=1.0, - ) - ) - wake_queue.put_nowait(None) - app.state.worker_stop_event = stop_event - app.state.worker_wake_queue = wake_queue - app.state.worker_task = worker_task - - -async def _stop_worker(app: FastAPI) -> None: - stop_event = getattr(app.state, "worker_stop_event", None) - wake_queue = getattr(app.state, "worker_wake_queue", None) - worker_task = getattr(app.state, "worker_task", None) - - if stop_event is not None: - stop_event.set() - if wake_queue is not None: - wake_queue.put_nowait(None) - if worker_task is not None: - try: - await asyncio.wait_for(worker_task, timeout=2.0) - except TimeoutError: - worker_task.cancel() - with suppress(asyncio.CancelledError): - await worker_task +def get_worker_notifier(app: FastAPI) -> WorkerNotifier: + """Return app worker notifier, or a no-op fallback when unavailable.""" + return resolve_worker_notifier(app.state) @asynccontextmanager @@ -70,12 +41,17 @@ async def _lifespan(app: FastAPI): settings.upload_dir.mkdir(parents=True, exist_ok=True) settings.prompt_dir.mkdir(parents=True, exist_ok=True) - _start_worker(app) - try: + async with AsyncExitStack() as stack: + stack.push_async_callback(dispose_database_runtime) + stop_event, worker_notifier = await stack.enter_async_context( + worker_consumer_lifespan( + session_factory=app.state.runtime.session_factory, + poll_interval_seconds=1.0, + ) + ) + app.state.worker_stop_event = stop_event + app.state.worker_notifier = worker_notifier yield - finally: - await _stop_worker(app) - await dispose_database_runtime() def create_app() -> FastAPI: diff --git a/src/transcription/ui/components/upload.py b/src/transcription/ui/components/upload.py index 500554f..e8a0eb3 100644 --- a/src/transcription/ui/components/upload.py +++ b/src/transcription/ui/components/upload.py @@ -9,11 +9,11 @@ from nicegui import ui from nicegui.binding import bindable_dataclass from nicegui.events import UploadEventArguments -from transcription.services.documents import DocumentService -from transcription.services.documents import UploadError +from transcription.errors import AppError from transcription.services.documents import UploadJobResult from transcription.ui.components.error_presenter import show_error from transcription.ui.components.error_presenter import summarize_error +from transcription.worker import WorkerNotifier type UploadSubmitter = Callable[[str, bytes], Awaitable[UploadJobResult]] @@ -26,7 +26,7 @@ class UploadWidgetState: message: str = "" -def render_upload_widget(*, service: DocumentService) -> None: +def render_upload_widget(*, submitter: UploadSubmitter, notifier: WorkerNotifier | None = None) -> None: """Render upload controls and common status/error handling.""" state = UploadWidgetState() status_label = ui.label("Upload a document to start transcription.") @@ -41,12 +41,14 @@ def render_upload_widget(*, service: DocumentService) -> None: status_label.text = "Uploading..." try: payload = await event.file.read() - result = await service.upload_file(filename=event.file.name, file_bytes=payload) + result = await submitter(event.file.name, payload) job_id = result.job_id state.message = f"Created job {job_id}" if job_id is not None else "Upload complete" status_label.text = state.message + if notifier is not None: + notifier.notify() ui.notify(state.message, type="positive") - except UploadError as exc: + except AppError as exc: state.message = summarize_error(exc, operation="upload.submit") status_label.text = f"Upload failed: {state.message}" show_error(exc, title="Upload failed", operation="upload.submit") diff --git a/src/transcription/ui/pages/upload_page.py b/src/transcription/ui/pages/upload_page.py index 70d19ff..1b94622 100644 --- a/src/transcription/ui/pages/upload_page.py +++ b/src/transcription/ui/pages/upload_page.py @@ -2,20 +2,30 @@ from __future__ import annotations +from fastapi import Request from nicegui import ui -from transcription.services.documents import DocumentService +from transcription.db import get_session +from transcription.services.store import create_upload_job from transcription.ui.components.upload import render_upload_widget +from transcription.worker import resolve_worker_notifier def register_page() -> None: """Register the upload page route.""" - service = DocumentService() - @ui.page("/upload", title="Upload Document") - def upload_page() -> None: - render_upload_widget(service=service) + def upload_page(request: Request) -> None: + async def submit_upload(filename: str, file_bytes: bytes): + async with get_session(session_factory=request.app.state.runtime.session_factory) as session: + return await create_upload_job( + filename=filename, + file_bytes=file_bytes, + session=session, + ) + + notify_worker = resolve_worker_notifier(request.app.state) + render_upload_widget(submitter=submit_upload, notifier=notify_worker) with ui.row(): ui.link("View jobs", "/jobs") diff --git a/src/transcription/worker.py b/src/transcription/worker.py index e4ef785..c5a39de 100644 --- a/src/transcription/worker.py +++ b/src/transcription/worker.py @@ -8,6 +8,7 @@ from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from contextlib import contextmanager from contextlib import suppress +from typing import Protocol from uuid import UUID from sqlalchemy.ext.asyncio import async_sessionmaker @@ -27,6 +28,73 @@ from .services.workflows import process_next_queued_job as process_next_queued_j logger = logging.getLogger(__name__) +class WorkerNotifier(Protocol): + """Abstraction for signaling the worker loop about new work.""" + + def notify(self) -> None: + """Signal the worker loop that work may be available.""" + + +class EventWorkerNotifier: + """Worker notifier backed by an asyncio.Event.""" + + def __init__(self, wake_event: asyncio.Event): + self._wake_event = wake_event + + def notify(self) -> None: + self._wake_event.set() + + +class NoopWorkerNotifier: + """Fallback notifier used when worker signaling is unavailable.""" + + def notify(self) -> None: + return + + +def resolve_worker_notifier(state: object) -> WorkerNotifier: + """Resolve notifier from app-like state objects with no-op fallback.""" + notifier = getattr(state, "worker_notifier", None) + if isinstance(notifier, NoopWorkerNotifier): + return notifier + if notifier is None: + return NoopWorkerNotifier() + return notifier + + +@asynccontextmanager +async def worker_consumer_lifespan( + *, + session_factory: async_sessionmaker[AsyncSession] | None = None, + poll_interval_seconds: float = 1.0, +) -> AsyncGenerator[tuple[asyncio.Event, WorkerNotifier]]: + """Start and stop the worker consumer loop for app lifespan.""" + stop_event = asyncio.Event() + wake_event = asyncio.Event() + worker_notifier: WorkerNotifier = EventWorkerNotifier(wake_event) + worker_task = asyncio.create_task( + run_worker_loop( + session_factory=session_factory, + stop_event=stop_event, + wake_event=wake_event, + poll_interval_seconds=poll_interval_seconds, + ) + ) + worker_notifier.notify() + + try: + yield stop_event, worker_notifier + finally: + stop_event.set() + worker_notifier.notify() + try: + await asyncio.wait_for(worker_task, timeout=2.0) + except TimeoutError: + worker_task.cancel() + with suppress(asyncio.CancelledError): + await worker_task + + async def queue_consumer_loop(queue: asyncio.Queue[UUID], stop_event: asyncio.Event): """Main worker loop that consumes jobs from the queue and processes them. @@ -65,12 +133,12 @@ async def run_worker_loop( *, session_factory: async_sessionmaker[AsyncSession] | None = None, stop_event: asyncio.Event | None = None, - wake_queue: asyncio.Queue[None] | None = None, + wake_event: asyncio.Event | None = None, poll_interval_seconds: float = 1.0, ) -> None: """Run worker loop until stop_event is set. - If wake_queue is provided, queue activity wakes the loop immediately while + If wake_event is provided, signal activity wakes the loop immediately while timeout-based wakeups preserve current polling behavior. """ while True: @@ -78,15 +146,16 @@ async def run_worker_loop( logger.info("Worker stop event received") return - if wake_queue is not None: + if wake_event is not None: with suppress(TimeoutError): - await asyncio.wait_for(wake_queue.get(), timeout=poll_interval_seconds) + await asyncio.wait_for(wake_event.wait(), timeout=poll_interval_seconds) + wake_event.clear() processed_any = False while await process_next_queued_job(session_factory=session_factory): processed_any = True - if wake_queue is None and not processed_any: + if wake_event is None and not processed_any: await asyncio.sleep(poll_interval_seconds)