"""Provider-neutral durable dispatcher for remote generation workers.""" from __future__ import annotations import asyncio from datetime import datetime, timedelta, timezone from pathlib import Path from app.core.config import Settings from app.core.logger import get_logger from app.generation.domain.enums import ( GenerationJobStatus, WorkerErrorCategory, WorkerJobStatus, ) from app.generation.domain.errors import ( GenerationError, GenerationOutputError, GenerationValidationError, GenerationWorkerError, ) from app.generation.domain.runtime import WorkerJob from app.generation.providers.base import GenerationProviderAdapter from app.generation.repositories.generation import GenerationDispatchRecord from app.generation.services.generation_service import GenerationService from app.security.assets import CanonicalAssetNotFoundError, CanonicalAssetService from app.security.audit import AuditService from app.security.models import CanonicalMediaAsset from app.services.cleanup import CleanupService logger = get_logger(__name__) class GenerationWorker: """Claim, dispatch, reconcile, and ingest provider-neutral generation jobs. This worker owns no provider URLs, credentials, schema, or storage. It calls the existing GenerationService internal boundary, which preserves authoritative workspace ownership and the durable state machine. """ def __init__( self, *, settings: Settings, generation: GenerationService, assets: CanonicalAssetService, cleanup: CleanupService, audit: AuditService, ) -> None: self.settings = settings self.generation = generation self.assets = assets self.cleanup = cleanup self.audit = audit self._task: asyncio.Task[None] | None = None self._stop = asyncio.Event() self._poll_not_before: dict[str, datetime] = {} self._poll_failures: dict[str, int] = {} async def start(self) -> None: if not self.settings.generation_worker_enabled: logger.info("generation worker disabled") return self._stop.clear() self._task = asyncio.create_task(self._run(), name="generation-dispatch-worker") async def stop(self) -> None: self._stop.set() if self._task is not None: self._task.cancel() try: await self._task except asyncio.CancelledError: logger.debug("generation worker task cancelled") self._task = None async def run_once(self) -> None: """Execute one bounded dispatch/reconciliation cycle for tests and runtime.""" if not self.settings.generation_enabled or not self.generation.ready: return await self._refresh_configured_providers() stale = await self.generation.fail_stale_unbound_submissions() if stale: logger.warning("generation submissions failed as ambiguous", extra={"count": stale}) await self._reconcile_bound_jobs() for item in await self.generation.claim_dispatchable_jobs( limit=self.settings.generation_worker_batch_size ): await self._dispatch(item) async def _run(self) -> None: while not self._stop.is_set(): try: await self.run_once() except asyncio.CancelledError: raise except Exception: # A single broken worker/provider cannot take down the ASGI # process or prevent the next durable jobs from being handled. logger.exception("generation worker iteration failed") try: await asyncio.wait_for( self._stop.wait(), timeout=self.settings.generation_worker_interval_seconds ) except asyncio.TimeoutError: continue async def _refresh_configured_providers(self) -> None: for adapter in self.generation.providers.list(): if not adapter.available: continue try: await self.generation.refresh_provider_runtime(adapter.provider) except Exception: # refresh_provider_runtime marks models unavailable. The # worker client already redacts transport errors. logger.warning( "generation worker provider unavailable", extra={"provider": adapter.provider}, ) async def _dispatch(self, item: GenerationDispatchRecord) -> None: adapter = self.generation.providers.get(item.provider) bound = False try: registered_model = self.generation.get_model(item.provider, item.model_id) if not registered_model.available: # No WAN request has been sent, so this is safe to retry as a # local availability delay rather than a duplicate-risk # submission. The durable attempt limit still bounds it. if item.attempt_number <= self.settings.generation_job_retry_limit: await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=GenerationJobStatus.RETRYING, error_code="GENERATION_MODEL_UNAVAILABLE", error_message="Generation model is not ready to accept work.", next_attempt_at=datetime.now(timezone.utc) + timedelta(seconds=self.settings.generation_worker_poll_backoff_seconds), ) else: await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=GenerationJobStatus.FAILED, error_code="GENERATION_MODEL_UNAVAILABLE", error_message=( "Generation model did not become ready within the retry limit." ), ) await self._audit_terminal( item, "ai.generation_failed", error_code="GENERATION_MODEL_UNAVAILABLE", ) return current = await self.generation.get_job( item.workspace_id, item.user_id, item.generation_job_id ) if current.status is GenerationJobStatus.CANCEL_REQUESTED: await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=GenerationJobStatus.CANCELLED, ) await self._audit_terminal(item, "ai.generation_cancelled") return if current.status is not GenerationJobStatus.SUBMITTING: return input_path: Path | None = None input_mime_type: str | None = None if item.input_asset_id is not None: asset, input_path = await self._resolve_input(item) input_mime_type = asset.mime_type remote_job = await adapter.submit( payload=item.spec, # This ID is only an internal correlation value for workers # that support idempotency. An adapter for a worker without # that protocol disables transport resubmission. idempotency_key=item.generation_request_id, input_path=input_path, input_mime_type=input_mime_type, ) await self.generation.bind_provider_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, worker_job_id=remote_job.external_job_id, provider_metadata={ "model_id": item.model_id, "worker_metadata": remote_job.metadata, }, ) bound = True logger.info( "generation provider submission accepted", extra={ "generation_job_id": item.generation_job_id, "provider": item.provider, "model_id": item.model_id, "attempt_number": item.attempt_number, }, ) await self._apply_remote_job(item, adapter, remote_job) except asyncio.CancelledError: raise except Exception as exc: if bound: await self._handle_bound_execution_failure(item, adapter, exc) else: await self._handle_submission_failure(item, adapter, exc) async def _resolve_input( self, item: GenerationDispatchRecord ) -> tuple[CanonicalMediaAsset, Path]: if item.input_asset_id is None: raise GenerationValidationError("Generation model requires a canonical input asset.") try: asset = await self.assets.get_owned_by_id( workspace_id=item.workspace_id, user_id=item.user_id, asset_id=item.input_asset_id, ) path = self.cleanup.resolve_download(asset.request_id, asset.filename) await self.assets.verify_file(asset, path) if path.is_symlink() or not path.is_file(): raise CanonicalAssetNotFoundError("Input asset is unsafe.") return asset, path except (CanonicalAssetNotFoundError, OSError) as exc: raise GenerationValidationError( "Generation input asset is no longer readable." ) from exc async def _handle_submission_failure( self, item: GenerationDispatchRecord, adapter: GenerationProviderAdapter, exc: Exception, ) -> None: error = self._normalise(adapter, exc) # A client-side timeout/connection failure, and a gateway 502/504, # may have reached a non-idempotent remote worker. A provider without # request lookup/idempotency support must never resubmit that ambiguity. if self._submission_is_ambiguous(error): code = "GENERATION_SUBMISSION_AMBIGUOUS" message = "Worker submission could not be reconciled safely; it was not retried." destination = GenerationJobStatus.FAILED else: retry_at = self._safe_submission_retry(item, error) if retry_at is not None: code = "GENERATION_SUBMISSION_RETRYING" message = "Generation worker did not accept the request; retry scheduled." destination = GenerationJobStatus.RETRYING else: code = "GENERATION_SUBMISSION_FAILED" message = "Generation worker rejected or failed the submission." destination = GenerationJobStatus.FAILED try: await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=destination, error_code=code, error_message=message, next_attempt_at=retry_at if destination is GenerationJobStatus.RETRYING else None, ) except Exception: logger.exception( "generation submission failure could not be persisted", extra={"generation_job_id": item.generation_job_id, "provider": item.provider}, ) return logger.warning( "generation provider submission failed", extra={ "generation_job_id": item.generation_job_id, "provider": item.provider, "model_id": item.model_id, "attempt_number": item.attempt_number, "category": error.category.value, "http_status": error.http_status, "retrying": destination is GenerationJobStatus.RETRYING, "ambiguous": destination is GenerationJobStatus.FAILED and self._submission_is_ambiguous(error), }, ) if destination is GenerationJobStatus.FAILED: await self._audit_terminal(item, "ai.generation_failed", error_code=code) async def _handle_bound_execution_failure( self, item: GenerationDispatchRecord, adapter: GenerationProviderAdapter, exc: Exception ) -> None: """Contain post-bind errors without ever attempting a second submit.""" error = self._normalise(adapter, exc) current = await self.generation.get_job( item.workspace_id, item.user_id, item.generation_job_id ) if current.status in { GenerationJobStatus.SUCCEEDED, GenerationJobStatus.FAILED, GenerationJobStatus.CANCELLED, }: return if self._is_transient_poll_error(error): self._defer_poll(item.generation_job_id) await self.generation.set_worker_reconciliation_due( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, due_at=self._poll_not_before[item.generation_job_id], ) return try: await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=GenerationJobStatus.FAILED, error_code="GENERATION_PROVIDER_EXECUTION_FAILED", error_message="Generation worker execution could not be reconciled.", ) await self._audit_terminal( item, "ai.generation_failed", error_code="GENERATION_PROVIDER_EXECUTION_FAILED", ) except Exception: logger.exception( "generation bound execution failure could not be persisted", extra={"generation_job_id": item.generation_job_id, "provider": item.provider}, ) async def _reconcile_bound_jobs(self) -> None: for item in await self.generation.list_reconcilable_jobs( limit=self.settings.generation_worker_batch_size ): if item.external_job_id is None: continue adapter = self.generation.providers.get(item.provider) try: remote_job = await adapter.get_job(external_job_id=item.external_job_id) if remote_job.external_job_id != item.external_job_id: raise GenerationOutputError("Worker returned an unexpected job identity.") self._clear_poll_backoff(item.generation_job_id) await self._apply_remote_job(item, adapter, remote_job) await self.generation.set_worker_reconciliation_due( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, due_at=None, ) except asyncio.CancelledError: raise except Exception as exc: error = self._normalise(adapter, exc) current = await self.generation.get_job( item.workspace_id, item.user_id, item.generation_job_id ) if current.status in { GenerationJobStatus.SUCCEEDED, GenerationJobStatus.FAILED, GenerationJobStatus.CANCELLED, }: self._clear_poll_backoff(item.generation_job_id) continue if self._is_transient_poll_error(error): self._defer_poll(item.generation_job_id) await self.generation.set_worker_reconciliation_due( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, due_at=self._poll_not_before[item.generation_job_id], ) logger.warning( "generation provider poll deferred", extra={ "generation_job_id": item.generation_job_id, "provider": item.provider, "category": error.category.value, "http_status": error.http_status, }, ) continue await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=GenerationJobStatus.FAILED, error_code="GENERATION_PROVIDER_STATUS_FAILED", error_message="Generation worker status reconciliation failed.", ) await self._audit_terminal( item, "ai.generation_failed", error_code="GENERATION_PROVIDER_STATUS_FAILED", ) logger.warning( "generation provider reconciliation failed", extra={ "generation_job_id": item.generation_job_id, "provider": item.provider, "category": error.category.value, }, ) async def _apply_remote_job( self, item: GenerationDispatchRecord, adapter: GenerationProviderAdapter, remote_job: WorkerJob, ) -> None: current = await self.generation.get_job( item.workspace_id, item.user_id, item.generation_job_id ) current_status = current.status if current_status is GenerationJobStatus.CANCEL_REQUESTED and remote_job.status in { WorkerJobStatus.QUEUED, WorkerJobStatus.RUNNING, }: # A cancellation may have raced the original worker submission. # Reuse the service cancellation boundary; it will never claim a # running WAN inference stopped unless the worker confirms it. try: await self.generation.cancel( item.workspace_id, item.user_id, item.generation_job_id ) except GenerationError: pass return if remote_job.status is WorkerJobStatus.QUEUED: if current_status is GenerationJobStatus.SUBMITTING: await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=GenerationJobStatus.QUEUED, ) return if remote_job.status is WorkerJobStatus.RUNNING: if current_status in {GenerationJobStatus.QUEUED, GenerationJobStatus.SUBMITTING}: await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=GenerationJobStatus.RUNNING, ) logger.info( "generation provider job running", extra={ "generation_job_id": item.generation_job_id, "provider": item.provider, "model_id": item.model_id, "attempt_number": item.attempt_number, }, ) return if remote_job.status is WorkerJobStatus.COMPLETED: await self.generation.ingest_completed_provider_output( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, worker_job=remote_job, ) await self._audit_terminal(item, "ai.generation_completed") logger.info( "generation provider job completed", extra={ "generation_job_id": item.generation_job_id, "provider": item.provider, "model_id": item.model_id, "attempt_number": item.attempt_number, }, ) return if remote_job.status is WorkerJobStatus.CANCELLED: await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=GenerationJobStatus.CANCELLED, ) await self._audit_terminal(item, "ai.generation_cancelled") return if remote_job.status is WorkerJobStatus.FAILED: await self.generation.transition_worker_job( workspace_id=item.workspace_id, user_id=item.user_id, job_id=item.generation_job_id, status=GenerationJobStatus.FAILED, error_code=remote_job.error_code or "GENERATION_INFERENCE_FAILED", error_message="Generation worker reported an inference failure.", ) await self._audit_terminal( item, "ai.generation_failed", error_code=remote_job.error_code or "GENERATION_INFERENCE_FAILED", ) logger.warning( "generation provider job failed", extra={ "generation_job_id": item.generation_job_id, "provider": item.provider, "model_id": item.model_id, "attempt_number": item.attempt_number, }, ) return raise GenerationOutputError("Generation worker returned an unsupported job status.") async def _audit_terminal( self, item: GenerationDispatchRecord, event_type: str, *, error_code: str | None = None, ) -> None: if item.product_surface != "ai_studio": return await self.audit.record_event( workspace_id=item.workspace_id, user_id=item.user_id, request_id=item.generation_request_id, event_type=event_type, entity_type="generation_job", entity_id=item.generation_job_id, metadata={ "operation": ("generate_image" if item.modality == "image" else "generate_video"), "model": item.model_id, "project_id": item.project_id, "error_code": error_code, }, ) @staticmethod def _normalise(adapter: GenerationProviderAdapter, exc: Exception) -> GenerationWorkerError: try: normalized = adapter.normalize_error(exc) except Exception: normalized = None if isinstance(normalized, GenerationWorkerError): return normalized return GenerationWorkerError( category=WorkerErrorCategory.UNKNOWN_ERROR, message="Generation worker operation failed unexpectedly.", retryable=False, ) @staticmethod def _submission_is_ambiguous(error: GenerationWorkerError) -> bool: return error.category is WorkerErrorCategory.TIMEOUT or ( error.category is WorkerErrorCategory.WORKER_UNAVAILABLE and error.http_status in {None, 502, 504} ) def _safe_submission_retry( self, item: GenerationDispatchRecord, error: GenerationWorkerError ) -> datetime | None: # A received 429/503 means the worker did not accept a job. Other # potentially transient submit failures are treated as ambiguous above # unless a future worker offers explicit idempotency/reconciliation. if error.http_status not in {429, 503}: return None if item.attempt_number > self.settings.generation_job_retry_limit: return None return self.generation.submission_retry_at( attempt_number=item.attempt_number, category=error.category, http_status=error.http_status, ) @staticmethod def _is_transient_poll_error(error: GenerationWorkerError) -> bool: return error.category in { WorkerErrorCategory.WORKER_UNAVAILABLE, WorkerErrorCategory.WORKER_NOT_READY, WorkerErrorCategory.TIMEOUT, WorkerErrorCategory.RATE_LIMITED, } or error.http_status in {429, 502, 503, 504} def _defer_poll(self, job_id: str) -> None: failures = self._poll_failures.get(job_id, 0) + 1 self._poll_failures[job_id] = failures delay = min( self.settings.generation_worker_poll_backoff_seconds * (2 ** (failures - 1)), 60.0, ) self._poll_not_before[job_id] = datetime.now(timezone.utc) + timedelta(seconds=delay) def _clear_poll_backoff(self, job_id: str) -> None: self._poll_not_before.pop(job_id, None) self._poll_failures.pop(job_id, None)