Spaces:
Running
Running
| """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, | |
| } | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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) | |
| 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" | |
| 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 | |
| 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 | |
| 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 | |
| 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) | |
| 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" | |