MediaRouter / app /generation /workers /generation_worker.py
basyx's picture
Upload 340 files
3493993 verified
Raw
History Blame Contribute Delete
25.9 kB
"""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)