MediaRouter / tests /test_generation_wan.py
basyx's picture
Upload 437 files
7cc81cb verified
Raw
History Blame Contribute Delete
11.7 kB
"""Mocked protocol tests for the audited WAN 2.2 worker integration."""
from __future__ import annotations
from collections.abc import Callable
from pathlib import Path
import httpx
import pytest
from app.core.config import Settings
from app.generation.domain.enums import (
GenerationModality,
WorkerCancellationStatus,
WorkerErrorCategory,
WorkerJobStatus,
)
from app.generation.domain.errors import GenerationWorkerError
from app.generation.domain.retry import GenerationRetryPolicy
from app.generation.model_registry import GenerationModelRegistration, GenerationModelRegistry
from app.generation.providers.wan import (
WAN_MODEL_CAPABILITY,
WAN_MODEL_ID,
WAN_PROVIDER_ID,
WanProviderAdapter,
)
from app.generation.providers.worker_client import RemoteWorkerClient
from app.generation.schemas.requests import GenerationRequestCreate
def _client(
handler: Callable[[httpx.Request], httpx.Response], *, retries: int = 2
) -> RemoteWorkerClient:
return RemoteWorkerClient(
base_url="https://wan-worker.example",
bearer_token="x" * 32,
connect_timeout_seconds=1,
request_timeout_seconds=1,
read_timeout_seconds=1,
retry_policy=GenerationRetryPolicy(max_retries=retries, backoff_seconds=0),
http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)),
sleep=lambda _: _no_sleep(),
)
async def _no_sleep() -> None:
return None
def _payload(**overrides: object) -> GenerationRequestCreate:
value: dict[str, object] = {
"provider": WAN_PROVIDER_ID,
"model_id": WAN_MODEL_ID,
"modality": "video",
"input_asset_id": "11111111-1111-4111-8111-111111111111",
"prompt": "Slow cinematic cloud movement",
"wan": {
"duration_seconds": 0.5,
"steps": 4,
"guidance_scale": 1.0,
"guidance_scale_2": 1.0,
"seed": 42,
"randomize_seed": False,
},
}
value.update(overrides)
return GenerationRequestCreate.model_validate(value)
def _info() -> dict[str, object]:
return {
"id": "wan2.2",
"name": "WAN 2.2 FP8 AOTI Faster",
"type": "video",
"task": "image-to-video",
"status": "ready",
"model_id": "Wan-AI/Wan2.2-I2V-A14B-Diffusers",
"fps": 16,
}
@pytest.mark.asyncio
async def test_wan_exact_model_discovery_and_readiness() -> None:
def handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/health":
return httpx.Response(200, json={"status": "ok", "service": "mediarouter-wan-worker"})
if request.url.path == "/ready":
return httpx.Response(
200,
json={
"status": "ready",
"model_loaded": True,
"model": "wan2.2",
"accepting_jobs": True,
},
)
return httpx.Response(200, json=_info())
adapter = WanProviderAdapter(client=_client(handler))
registry = GenerationModelRegistry(
[
GenerationModelRegistration(
provider_id=WAN_PROVIDER_ID,
model=WAN_MODEL_CAPABILITY,
configuration_reference="wan-space",
)
]
)
models = registry.verify_readiness(
provider_id=WAN_PROVIDER_ID,
worker_info=await adapter.info(),
readiness=await adapter.ready(),
provider_configured=adapter.available,
)
assert models[0].model.id == WAN_MODEL_ID
assert models[0].model.modality is GenerationModality.VIDEO
assert models[0].available
@pytest.mark.asyncio
async def test_wan_not_ready_and_model_mismatch_are_not_advertised() -> None:
def handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/ready":
return httpx.Response(
503,
json={
"status": "not_ready",
"model_loaded": False,
"model": "other-model",
"accepting_jobs": False,
},
)
return httpx.Response(
200, json={"status": "ok"} if request.url.path == "/health" else _info()
)
adapter = WanProviderAdapter(client=_client(handler))
with pytest.raises(GenerationWorkerError) as raised:
await adapter.ready()
assert raised.value.category is WorkerErrorCategory.WORKER_NOT_READY
@pytest.mark.asyncio
async def test_wan_model_identity_mismatch_remains_unavailable() -> None:
wrong_info = {
**_info(),
"id": "different-wan-model",
"name": "Different model",
}
def handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/ready":
return httpx.Response(
200,
json={
"status": "ready",
"model_loaded": True,
"model": WAN_MODEL_ID,
"accepting_jobs": True,
},
)
return httpx.Response(200, json=wrong_info)
adapter = WanProviderAdapter(client=_client(handler))
with pytest.raises(GenerationWorkerError) as raised:
await adapter.info()
assert raised.value.category is WorkerErrorCategory.PROVIDER_ERROR
@pytest.mark.asyncio
async def test_wan_submission_is_multipart_and_has_no_automatic_retry(tmp_path: Path) -> None:
seen: list[httpx.Request] = []
def handler(request: httpx.Request) -> httpx.Response:
seen.append(request)
return httpx.Response(202, json={"job_id": "wan_" + "a" * 32, "status": "queued"})
source = tmp_path / "input.png"
source.write_bytes(b"not-decoded-in-adapter-test")
adapter = WanProviderAdapter(client=_client(handler))
job = await adapter.submit(
payload={"prompt": "slow movement", "wan": {"duration_seconds": 0.5, "steps": 4}},
idempotency_key="generation-request-id",
input_path=source,
input_mime_type="image/png",
)
assert job.status is WorkerJobStatus.QUEUED
assert job.external_job_id.startswith("wan_")
assert seen[0].headers["authorization"] == "Bearer " + "x" * 32
body = seen[0].content.decode("latin-1")
assert 'name="image"' in body
assert 'name="duration_seconds"' in body
assert 'name="width"' not in body
@pytest.mark.asyncio
async def test_wan_submission_connection_ambiguity_is_not_retried(tmp_path: Path) -> None:
calls = 0
request = httpx.Request("POST", "https://wan-worker.example/v1/generate")
def handler(_: httpx.Request) -> httpx.Response:
nonlocal calls
calls += 1
raise httpx.ConnectError("Bearer " + "x" * 32, request=request)
source = tmp_path / "input.png"
source.write_bytes(b"input")
adapter = WanProviderAdapter(client=_client(handler, retries=3))
with pytest.raises(GenerationWorkerError) as raised:
await adapter.submit(
payload={"prompt": "slow movement"},
idempotency_key="generation-request-id",
input_path=source,
input_mime_type="image/png",
)
assert raised.value.category is WorkerErrorCategory.WORKER_UNAVAILABLE
assert calls == 1
assert "Bearer" not in str(raised.value)
@pytest.mark.asyncio
async def test_wan_completed_job_maps_a_safe_video_output() -> None:
job_id = "wan_" + "b" * 32
def handler(_: httpx.Request) -> httpx.Response:
return httpx.Response(
200,
json={
"job_id": job_id,
"status": "completed",
"output": {"type": "video", "filename": f"{job_id}.mp4"},
},
)
adapter = WanProviderAdapter(client=_client(handler))
job = await adapter.get_job(external_job_id=job_id)
assert job.output is not None
assert job.output.mime_type == "video/mp4"
assert job.output.provider_output_id == job_id
assert job.output.download_path == f"/v1/jobs/{job_id}/output"
@pytest.mark.asyncio
@pytest.mark.parametrize("status_code", [429, 502, 503, 504])
async def test_wan_polling_uses_shared_bounded_transient_retry(status_code: int) -> None:
calls = 0
job_id = "wan_" + "d" * 32
def handler(_: httpx.Request) -> httpx.Response:
nonlocal calls
calls += 1
if calls < 3:
return httpx.Response(status_code, json={"detail": {"token": "never-store"}})
return httpx.Response(200, json={"job_id": job_id, "status": "running"})
job = await WanProviderAdapter(client=_client(handler, retries=2)).get_job(
external_job_id=job_id
)
assert job.status is WorkerJobStatus.RUNNING
assert calls == 3
@pytest.mark.asyncio
async def test_wan_polling_does_not_retry_permanent_client_errors() -> None:
calls = 0
job_id = "wan_" + "e" * 32
def handler(_: httpx.Request) -> httpx.Response:
nonlocal calls
calls += 1
return httpx.Response(400, json={"detail": {"code": "WAN_PARAMETERS_INVALID"}})
with pytest.raises(GenerationWorkerError) as raised:
await WanProviderAdapter(client=_client(handler, retries=3)).get_job(
external_job_id=job_id
)
assert raised.value.category is WorkerErrorCategory.INVALID_REQUEST
assert calls == 1
@pytest.mark.asyncio
async def test_wan_cancellation_only_confirms_queued_worker_cancellation() -> None:
def queued_handler(_: httpx.Request) -> httpx.Response:
return httpx.Response(200, json={"job_id": "wan_" + "c" * 32, "status": "cancelled"})
def running_handler(_: httpx.Request) -> httpx.Response:
return httpx.Response(
409,
json={
"detail": {
"code": "WAN_JOB_NOT_CANCELLABLE",
"message": "A running job cannot be cancelled.",
"status": "running",
}
},
)
assert (
await WanProviderAdapter(client=_client(queued_handler)).cancel(
external_job_id="wan_" + "c" * 32
)
).status is WorkerCancellationStatus.CANCELLED
assert (
await WanProviderAdapter(client=_client(running_handler)).cancel(
external_job_id="wan_" + "c" * 32
)
).status is WorkerCancellationStatus.FAILED
@pytest.mark.parametrize(
"invalid",
[
{"prompt": " "},
{"modality": "image"},
{"input_asset_id": None},
],
)
@pytest.mark.asyncio
async def test_wan_request_validation_rejects_invalid_required_values(
invalid: dict[str, object]
) -> None:
adapter = WanProviderAdapter(client=None)
with pytest.raises(Exception):
payload = _payload(**invalid)
await adapter.validate_request(payload)
@pytest.mark.parametrize("field", ["width", "height", "num_frames", "provider_payload"])
def test_wan_schema_rejects_unsupported_parameters(field: str) -> None:
raw = _payload().model_dump()
wan = dict(raw["wan"] or {})
wan[field] = 1
raw["wan"] = wan
with pytest.raises(ValueError):
GenerationRequestCreate.model_validate(raw)
def test_wan_configuration_is_optional_and_never_enables_flux() -> None:
disabled = WanProviderAdapter.from_settings(Settings(_env_file=None))
invalid = WanProviderAdapter.from_settings(
Settings(_env_file=None, wan_space_url="https://wan-worker.example")
)
assert not disabled.available
assert not invalid.available
assert invalid.configuration_error is not None
assert WAN_PROVIDER_ID == "wan"