basyx's picture
Upload 340 files
3493993 verified
Raw
History Blame Contribute Delete
17.9 kB
"""WAN 2.2 image-to-video adapter for the trusted MediaRouter worker API.
The worker protocol in this module was audited against the companion WAN
Space. This adapter deliberately knows only that protocol; orchestration,
tenancy, durable jobs, output storage, and public transport remain owned by
the provider-neutral generation services.
"""
from __future__ import annotations
from collections.abc import AsyncIterator
from contextlib import AbstractAsyncContextManager
from pathlib import Path
from app.core.config import Settings
from app.generation.domain.capabilities import (
GenerationModelCapability,
GenerationProviderCapabilities,
)
from app.generation.domain.enums import (
GenerationModality,
WorkerCancellationStatus,
WorkerErrorCategory,
WorkerJobStatus,
WorkerReadinessStatus,
)
from app.generation.domain.errors import (
GenerationCapabilityUnsupportedError,
GenerationOutputError,
GenerationProviderUnavailableError,
GenerationValidationError,
GenerationWorkerError,
)
from app.generation.domain.runtime import (
WorkerCancellationResult,
WorkerHealth,
WorkerInfo,
WorkerJob,
WorkerOutput,
WorkerReadiness,
safe_worker_metadata,
)
from app.generation.providers.base import GenerationProviderAdapter
from app.generation.providers.worker_client import RemoteWorkerClient
from app.generation.schemas.requests import GenerationRequestCreate, WanGenerationOptions
from app.security.models import CanonicalMediaAsset
WAN_PROVIDER_ID = "wan"
WAN_MODEL_ID = "wan2.2"
WAN_MODEL_NAME = "WAN 2.2 FP8 AOTI Faster"
WAN_UNDERLYING_MODEL_ID = "Wan-AI/Wan2.2-I2V-A14B-Diffusers"
WAN_INPUT_MAX_BYTES = 20 * 1024 * 1024
WAN_INPUT_MIME_PREFIX = "image/"
WAN_MODEL_CAPABILITY = GenerationModelCapability(
id=WAN_MODEL_ID,
name=WAN_MODEL_NAME,
modality=GenerationModality.VIDEO,
input_asset_supported=True,
input_schema={
"input_asset": {
"required": True,
"media_type": "image/*",
"max_bytes": WAN_INPUT_MAX_BYTES,
},
"prompt": {"required": True, "min_length": 1, "max_length": 4000},
"negative_prompt": {"required": False, "max_length": 4000},
"duration_seconds": {"required": False, "minimum": 0.5, "maximum": 5.0},
"steps": {"required": False, "minimum": 1, "maximum": 30},
"guidance_scale": {"required": False, "minimum": 0.0, "maximum": 10.0},
"guidance_scale_2": {"required": False, "minimum": 0.0, "maximum": 10.0},
"seed": {"required": False, "minimum": 0, "maximum": 2_147_483_647},
"randomize_seed": {"required": False, "type": "boolean"},
"derived": {
"dimensions": "source image, resized by worker to its supported bounds",
"frames": "duration_seconds at worker-reported 16 fps",
},
},
)
class WanProviderAdapter(GenerationProviderAdapter):
"""Strict adapter for the audited, authenticated WAN worker contract."""
capabilities = GenerationProviderCapabilities(
provider=WAN_PROVIDER_ID,
name="WAN 2.2",
models=[WAN_MODEL_CAPABILITY],
implementation_status="implemented",
supports_cancellation=True,
supports_status_reconciliation=True,
)
def __init__(
self,
*,
client: RemoteWorkerClient | None,
configuration_error: str | None = None,
) -> None:
self.client = client
self.configuration_error = configuration_error
@classmethod
def from_settings(cls, settings: Settings) -> "WanProviderAdapter":
"""Build an optional adapter without allowing a bad config to abort startup."""
url = settings.wan_space_url.strip()
token = settings.wan_space_token
if not url and token is None:
return cls(client=None)
if not url or token is None:
return cls(
client=None,
configuration_error="WAN worker URL and token must be configured together.",
)
secret = token.get_secret_value()
if not 32 <= len(secret) <= 4096:
return cls(
client=None,
configuration_error="WAN worker token has an invalid length.",
)
try:
from app.generation.domain.retry import GenerationRetryPolicy
return cls(
client=RemoteWorkerClient(
base_url=url,
bearer_token=token,
connect_timeout_seconds=settings.ai_worker_connect_timeout_seconds,
request_timeout_seconds=settings.ai_worker_request_timeout_seconds,
read_timeout_seconds=settings.ai_worker_read_timeout_seconds,
retry_policy=GenerationRetryPolicy(
max_retries=settings.ai_worker_max_retries,
backoff_seconds=settings.ai_worker_retry_backoff_seconds,
),
)
)
except (TypeError, ValueError):
# Never include the configured URL or token in a startup error or
# log. An operator can correct configuration without affecting
# the non-generation application surface.
return cls(client=None, configuration_error="WAN worker configuration is invalid.")
@property
def available(self) -> bool:
# Availability at the model level still requires a successful health,
# readiness and exact-info verification in GenerationModelRegistry.
return self.client is not None
async def validate_request(
self, payload: GenerationRequestCreate
) -> dict[str, object]:
if payload.provider != self.provider or payload.model_id != WAN_MODEL_ID:
raise GenerationCapabilityUnsupportedError("WAN request targets an unsupported model.")
if payload.modality is not GenerationModality.VIDEO:
raise GenerationCapabilityUnsupportedError("WAN 2.2 supports video generation only.")
if payload.input_asset_id is None:
raise GenerationValidationError("WAN 2.2 requires an image input asset.")
if not payload.prompt or not payload.prompt.strip() or len(payload.prompt) > 4_000:
raise GenerationValidationError("WAN prompt must contain 1 through 4000 characters.")
if payload.wan is None:
# Preserve the worker's audited defaults when no optional control
# was requested; do not persist an invented empty provider blob.
return {"prompt": payload.prompt}
# Validate again at the adapter boundary so internal callers cannot
# hand the worker an arbitrary model-dumped object.
options = WanGenerationOptions.model_validate(payload.wan)
return {"prompt": payload.prompt, "wan": options.model_dump(exclude_unset=True)}
async def validate_input_asset(
self, payload: GenerationRequestCreate, asset: CanonicalMediaAsset
) -> None:
del payload
mime_type = (asset.mime_type or "").split(";", 1)[0].strip().lower()
if not mime_type.startswith(WAN_INPUT_MIME_PREFIX):
raise GenerationValidationError("WAN 2.2 requires a canonical image input asset.")
if asset.file_size <= 0 or asset.file_size > WAN_INPUT_MAX_BYTES:
raise GenerationValidationError(
"WAN 2.2 input image exceeds the worker's 20 MiB limit."
)
async def info(self) -> WorkerInfo:
info = await self._client().info()
# The audited WAN worker is a single-model worker whose top-level
# identity is the public model identifier. Do not accept a different
# worker that merely happens to list ``wan2.2`` in a secondary model
# collection: it could expose different preprocessing, output, or
# cancellation semantics than this adapter has been reviewed for.
if (
info.id != WAN_MODEL_ID
or GenerationModality.VIDEO not in info.media_types
or info.metadata.get("task") != "image-to-video"
or info.metadata.get("model_id") != WAN_UNDERLYING_MODEL_ID
or info.metadata.get("fps") != 16
):
raise GenerationWorkerError(
category=WorkerErrorCategory.PROVIDER_ERROR,
message="Configured WAN worker identity does not match the registered model.",
retryable=False,
)
return info
async def health(self) -> WorkerHealth:
return await self._client().health()
async def ready(self) -> WorkerReadiness:
readiness = await self._client().ready()
accepting_jobs = readiness.metadata.get("accepting_jobs")
# The audited worker returns 503/"not_ready" while it is not ready,
# but retain this extra guard for malformed or future responses.
if accepting_jobs is not True:
return readiness.model_copy(
update={"status": WorkerReadinessStatus.UNAVAILABLE}
)
return readiness
async def submit(
self,
*,
payload: dict[str, object],
idempotency_key: str,
input_path: Path | None = None,
input_mime_type: str | None = None,
) -> WorkerJob:
if input_path is None or input_mime_type is None:
raise GenerationValidationError("WAN generation requires a verified image input asset.")
mime_type = input_mime_type.split(";", 1)[0].strip().lower()
if not mime_type.startswith(WAN_INPUT_MIME_PREFIX):
raise GenerationValidationError("WAN generation requires an image input asset.")
try:
size = input_path.stat().st_size
except OSError as exc:
raise GenerationValidationError("WAN input asset is no longer readable.") from exc
if size <= 0 or size > WAN_INPUT_MAX_BYTES:
raise GenerationValidationError("WAN input image exceeds the worker's 20 MiB limit.")
form = self._multipart_fields(payload)
# The audited worker has no idempotency-key protocol. Include the
# canonical request ID for correlation only and forbid transport
# retries: a lost response is intentionally handled as ambiguous by
# the durable dispatcher rather than creating a duplicate video.
response = await self._client().submit_multipart(
fields=form,
file_field="image",
file_path=input_path,
filename=input_path.name,
mime_type=mime_type,
idempotency_key=idempotency_key,
idempotent=False,
)
return self._job_from_payload(response, expected_job_id=None)
async def get_job(self, *, external_job_id: str) -> WorkerJob:
return self._job_from_payload(
await self._client().get_job_payload(external_job_id),
expected_job_id=external_job_id,
)
async def cancel(self, *, external_job_id: str) -> WorkerCancellationResult:
payload = await self._client().cancel_job_payload(
external_job_id, expected_statuses={200, 409}
)
if payload.get("status") == "cancelled":
return WorkerCancellationResult(status=WorkerCancellationStatus.CANCELLED)
detail = payload.get("detail")
if isinstance(detail, dict) and detail.get("code") == "WAN_JOB_NOT_CANCELLABLE":
# The worker explicitly says a running GPU operation was not
# stopped. This is a cancellation failure, not a success or a
# provider-wide lack of cancellation capability.
return WorkerCancellationResult(
status=WorkerCancellationStatus.FAILED,
metadata={"reason": "job_not_cancellable"},
)
return WorkerCancellationResult(
status=WorkerCancellationStatus.FAILED,
metadata={"reason": "unexpected_cancellation_response"},
)
async def retrieve_output(self, *, external_job_id: str) -> WorkerOutput:
job = await self.get_job(external_job_id=external_job_id)
if job.status is not WorkerJobStatus.COMPLETED or job.output is None:
raise GenerationOutputError("WAN output is not ready.")
return job.output
def stream_output(
self, output: WorkerOutput
) -> AbstractAsyncContextManager[AsyncIterator[bytes]]:
if output.output_type is not GenerationModality.VIDEO or output.mime_type != "video/mp4":
return super().stream_output(output)
return self._client().stream_output(output)
def normalize_error(self, error: Exception) -> GenerationWorkerError | Exception:
if isinstance(error, GenerationWorkerError):
return error
return GenerationWorkerError(
category=WorkerErrorCategory.UNKNOWN_ERROR,
message="WAN worker operation failed unexpectedly.",
retryable=False,
)
async def close(self) -> None:
if self.client is not None:
await self.client.aclose()
def _client(self) -> RemoteWorkerClient:
if self.client is None:
raise GenerationProviderUnavailableError("WAN worker is not configured.")
return self.client
@staticmethod
def _multipart_fields(payload: dict[str, object]) -> dict[str, str]:
prompt = payload.get("prompt")
raw_options = payload.get("wan", {})
if not isinstance(prompt, str) or not prompt.strip() or not isinstance(raw_options, dict):
raise GenerationValidationError("WAN generation request is invalid.")
try:
options = WanGenerationOptions.model_validate(raw_options)
except ValueError as exc:
raise GenerationValidationError("WAN generation options are invalid.") from exc
# Persisted specs use ``exclude_unset``. Sending only those fields
# preserves WAN's own documented defaults for omitted controls.
fields: dict[str, str] = {"prompt": prompt}
for name, value in options.model_dump(exclude_unset=True).items():
if value is None:
continue
if isinstance(value, bool):
fields[name] = "true" if value else "false"
else:
fields[name] = str(value)
return fields
@staticmethod
def _job_from_payload(
payload: dict[str, object], *, expected_job_id: str | None
) -> WorkerJob:
job_id = payload.get("job_id")
raw_status = payload.get("status")
if not isinstance(job_id, str) or not isinstance(raw_status, str):
raise GenerationWorkerError(
category=WorkerErrorCategory.PROVIDER_ERROR,
message="WAN worker returned an invalid job response.",
retryable=False,
)
if expected_job_id is not None and job_id != expected_job_id:
raise GenerationWorkerError(
category=WorkerErrorCategory.PROVIDER_ERROR,
message="WAN worker returned an unexpected job identity.",
retryable=False,
)
statuses = {
"queued": WorkerJobStatus.QUEUED,
"running": WorkerJobStatus.RUNNING,
"completed": WorkerJobStatus.COMPLETED,
"failed": WorkerJobStatus.FAILED,
"cancelled": WorkerJobStatus.CANCELLED,
}
state = statuses.get(raw_status.lower())
if state is None:
raise GenerationWorkerError(
category=WorkerErrorCategory.PROVIDER_ERROR,
message="WAN worker returned an unsupported job status.",
retryable=False,
)
output: WorkerOutput | None = None
if state is WorkerJobStatus.COMPLETED:
raw_output = payload.get("output")
if not isinstance(raw_output, dict) or raw_output.get("type") != "video":
raise GenerationOutputError("WAN worker returned an invalid video output.")
filename = raw_output.get("filename")
if not isinstance(filename, str):
raise GenerationOutputError("WAN worker returned an invalid video filename.")
output = WorkerOutput(
output_type=GenerationModality.VIDEO,
mime_type="video/mp4",
provider_output_id=job_id,
download_path=f"/v1/jobs/{job_id}/output",
filename=filename,
)
error_code: str | None = None
error_message: str | None = None
raw_error = payload.get("error")
if isinstance(raw_error, dict):
code = raw_error.get("code")
message = raw_error.get("message")
error_code = code if isinstance(code, str) else None
error_message = message if isinstance(message, str) else None
metadata = safe_worker_metadata(
{
key: value
for key, value in payload.items()
if key not in {"job_id", "status", "output", "error"}
}
)
try:
return WorkerJob(
external_job_id=job_id,
status=state,
output=output,
error_category=(
WorkerErrorCategory.INFERENCE_ERROR
if state is WorkerJobStatus.FAILED
else None
),
error_code=error_code,
error_message=error_message,
metadata=metadata if isinstance(metadata, dict) else {},
)
except ValueError as exc:
raise GenerationWorkerError(
category=WorkerErrorCategory.PROVIDER_ERROR,
message="WAN worker returned invalid job metadata.",
retryable=False,
) from exc