Spaces:
Sleeping
Sleeping
| """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, | |
| }, | |
| ) | |
| 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, | |
| ) | |
| 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, | |
| ) | |
| 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) | |