"""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"