basyx's picture
Upload 340 files
3493993 verified
Raw
History Blame Contribute Delete
5.65 kB
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