from __future__ import annotations from abc import ABC from collections.abc import AsyncIterator from contextlib import AbstractAsyncContextManager from pathlib import Path from app.generation.domain.capabilities import GenerationProviderCapabilities from app.generation.domain.errors import ( GenerationCapabilityUnsupportedError, GenerationProviderUnavailableError, GenerationWorkerError, ) from app.generation.domain.runtime import ( WorkerCancellationResult, WorkerHealth, WorkerInfo, WorkerJob, WorkerOutput, WorkerReadiness, ) from app.generation.schemas.requests import GenerationRequestCreate from app.security.models import CanonicalMediaAsset class _UnavailableOutputStream(AbstractAsyncContextManager[AsyncIterator[bytes]]): """Context-manager shape for the default unavailable output stream.""" def __init__(self, provider: str) -> None: self.provider = provider async def __aenter__(self) -> AsyncIterator[bytes]: raise GenerationProviderUnavailableError( f"{self.provider} output streaming is unavailable." ) async def __aexit__(self, *_: object) -> None: return None class GenerationProviderAdapter(ABC): """Provider-neutral worker contract. No credential, endpoint, or model-specific logic belongs in REST, MCP, SDK, n8n, or the generation service. A concrete adapter will be added only alongside a verified WAN or FLUX worker integration. """ capabilities: GenerationProviderCapabilities @property def provider(self) -> str: return self.capabilities.provider @property def available(self) -> bool: """Whether this process can safely accept new work for this adapter.""" return False def model_capability(self, model_id: str): for model in self.capabilities.models: if model.id == model_id: return model raise GenerationCapabilityUnsupportedError( f"{self.provider} does not support model '{model_id}'." ) async def validate_request( self, payload: GenerationRequestCreate ) -> dict[str, object]: """Return a normalized, non-secret worker payload. Concrete adapters must validate their own strict Pydantic model and must not pass through unknown fields. """ raise GenerationCapabilityUnsupportedError( f"{self.provider} generation is not implemented." ) async def validate_input_asset( self, payload: GenerationRequestCreate, asset: CanonicalMediaAsset ) -> None: """Validate a resolved canonical input descriptor before a job is queued. This hook intentionally receives a record owned by the service, never a client path, URL, or arbitrary upload object. Adapters may impose MIME/type constraints but byte-level readability is checked again by the trusted dispatcher immediately before submission. """ del payload, asset return None async def info(self) -> WorkerInfo: raise GenerationProviderUnavailableError( f"{self.provider} worker metadata is unavailable." ) async def health(self) -> WorkerHealth: raise GenerationProviderUnavailableError( f"{self.provider} worker health is unavailable." ) async def ready(self) -> WorkerReadiness: raise GenerationProviderUnavailableError( f"{self.provider} worker readiness is unavailable." ) async def submit( self, *, payload: dict[str, object], idempotency_key: str, input_path: Path | None = None, input_mime_type: str | None = None, ) -> WorkerJob: del input_path, input_mime_type raise GenerationProviderUnavailableError( f"{self.provider} generation worker is unavailable." ) async def get_job(self, *, external_job_id: str) -> WorkerJob: raise GenerationProviderUnavailableError( f"{self.provider} status reconciliation is unavailable." ) async def get_status(self, *, external_job_id: str) -> WorkerJob: """Compatibility alias; future code should call ``get_job``.""" return await self.get_job(external_job_id=external_job_id) async def cancel(self, *, external_job_id: str) -> WorkerCancellationResult: if not self.capabilities.supports_cancellation: from app.generation.domain.enums import WorkerCancellationStatus return WorkerCancellationResult(status=WorkerCancellationStatus.UNSUPPORTED) raise GenerationProviderUnavailableError( f"{self.provider} cancellation is unavailable." ) async def retrieve_output(self, *, external_job_id: str) -> WorkerOutput: raise GenerationProviderUnavailableError( f"{self.provider} output retrieval is unavailable." ) def stream_output( self, output: WorkerOutput ) -> AbstractAsyncContextManager[AsyncIterator[bytes]]: """Return a scoped byte stream for a validated worker output descriptor.""" del output return _UnavailableOutputStream(self.provider) def normalize_error(self, error: Exception) -> GenerationWorkerError | Exception: """Keep provider error mapping in the adapter, never in transports.""" if isinstance(error, GenerationWorkerError): return error return error async def close(self) -> None: """Close provider-owned worker clients when the container stops.""" return None