From 93eed536d688c455d7cf1f25ea3325b93f69a0cc Mon Sep 17 00:00:00 2001 From: chengzeyi Date: Thu, 20 Aug 2026 11:51:47 +0000 Subject: [PATCH] test(wavespeed): cover the polling, envelope and unsupported-surface paths Patch coverage was 85%, mostly the async polling branches, the envelope error paths and the not-supported video surfaces. This takes every wavespeed module to 100%, and covers the two registration lines through the public entry points rather than by calling them directly: image generation now routes through litellm.image_generation, and provider resolution is exercised through get_llm_provider with the chat host as api_base. Worth calling out: the async submit-once-on-poll-failure test mirrors the sync one, since a duplicate submit would bill a second task. --- .../test_wavespeed_image_generation.py | 134 ++++++++++++++++ .../wavespeed/test_wavespeed_common_utils.py | 148 ++++++++++++++++++ .../llms/wavespeed/test_wavespeed_provider.py | 63 ++++++++ .../test_wavespeed_video_transformation.py | 135 +++++++++++++++- 4 files changed, 474 insertions(+), 6 deletions(-) create mode 100644 tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py diff --git a/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py b/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py index 6764999d408..b17451d109b 100644 --- a/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py +++ b/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py @@ -13,6 +13,7 @@ from litellm.llms.wavespeed.image_generation.handler import WaveSpeedImageGenera from litellm.llms.wavespeed.image_generation.transformation import ( WaveSpeedImageGenerationConfig, ) +from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ImageResponse MODEL = "wavespeed-ai/z-image/turbo" @@ -205,3 +206,136 @@ class TestWaveSpeedImageGenerationConfig: litellm_params={}, encoding=None, ) + + +@pytest.fixture +def zero_poll_budget(monkeypatch): + """Make the polling deadline expire immediately so the timeout path is reachable.""" + monkeypatch.setattr("litellm.llms.wavespeed.image_generation.handler.DEFAULT_MAX_POLLING_TIME", 0) + + +@respx.mock +def test_sync_poll_timeout(generate, zero_poll_budget): + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + poll = respx.get(RESULT_URL).mock(return_value=httpx.Response(200, json=prediction("processing"))) + + with pytest.raises(WaveSpeedError) as exc_info: + generate() + + assert exc_info.value.status_code == 408 + assert "did not finish within" in str(exc_info.value) + assert submit.call_count == 1 + assert poll.call_count == 0 + + +@pytest.mark.asyncio +@respx.mock +async def test_async_poll_timeout(zero_poll_budget): + respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + + with pytest.raises(WaveSpeedError) as exc_info: + await WaveSpeedImageGeneration().async_image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params={"api_key": "sk-test", "api_base": None}, + logging_obj=MagicMock(), + timeout=None, + ) + + assert exc_info.value.status_code == 408 + + +@pytest.mark.asyncio +@respx.mock +async def test_async_submit_is_issued_exactly_once_when_polling_fails(): + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + poll = respx.get(RESULT_URL).mock(side_effect=httpx.ConnectError("connection reset")) + + with pytest.raises(WaveSpeedError) as exc_info: + await WaveSpeedImageGeneration().async_image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params={"api_key": "sk-test", "api_base": None}, + logging_obj=MagicMock(), + timeout=None, + ) + + assert submit.call_count == 1 + assert poll.call_count == 5 + assert "5 times in a row" in str(exc_info.value) + + +@pytest.mark.asyncio +@respx.mock +async def test_aimg_generation_flag_dispatches_to_the_async_path(): + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + respx.get(RESULT_URL).mock(return_value=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL]))) + + pending = WaveSpeedImageGeneration().image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params={"api_key": "sk-test", "api_base": None}, + logging_obj=MagicMock(), + timeout=None, + aimg_generation=True, + ) + + response = await pending + assert [image.url for image in response.data] == [OUTPUT_URL] + assert submit.call_count == 1 + + +@respx.mock +def test_litellm_params_object_is_accepted(monkeypatch): + """images/main.py can hand the handler a GenericLiteLLMParams rather than a dict.""" + monkeypatch.delenv("WAVESPEED_API_BASE", raising=False) + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + respx.get(RESULT_URL).mock(return_value=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL]))) + + response = WaveSpeedImageGeneration().image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params=GenericLiteLLMParams(api_key="sk-test"), + logging_obj=MagicMock(), + timeout=None, + ) + + assert [image.url for image in response.data] == [OUTPUT_URL] + assert submit.calls[0].request.headers["authorization"] == "Bearer sk-test" + + +@respx.mock +def test_extra_headers_are_merged_and_cannot_be_dropped(generate): + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + respx.get(RESULT_URL).mock(return_value=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL]))) + + WaveSpeedImageGeneration().image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params={"api_key": "sk-test", "api_base": None}, + logging_obj=MagicMock(), + timeout=None, + extra_headers={"X-Trace-Id": "abc123"}, + ) + + assert submit.calls[0].request.headers["x-trace-id"] == "abc123" + assert submit.calls[0].request.headers["x-client-name"] == "litellm" + + +def test_supported_openai_params_and_error_class(): + config = WaveSpeedImageGenerationConfig() + assert config.get_supported_openai_params(MODEL) == ["n", "size", "response_format"] + + error = config.get_error_class("boom", 503, {}) + assert isinstance(error, WaveSpeedError) + assert error.status_code == 503 diff --git a/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py b/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py new file mode 100644 index 00000000000..f62efa0f0b8 --- /dev/null +++ b/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py @@ -0,0 +1,148 @@ +"""Unit tests for the WaveSpeed AI envelope parsing and URL helpers.""" + +import httpx +import pytest + +from litellm.llms.wavespeed.common_utils import ( + DEFAULT_API_BASE, + WaveSpeedError, + build_headers, + build_result_url, + build_submit_url, + get_api_base, + get_outputs, + get_prediction_id, + map_status_to_openai, + optional_entry, + optional_pair, + poll_outcome, + to_request_payload, + unwrap_envelope, +) + + +class TestUrls: + def test_submit_url_defaults_to_the_public_api(self, monkeypatch): + monkeypatch.delenv("WAVESPEED_API_BASE", raising=False) + assert build_submit_url(None, "wavespeed-ai/z-image/turbo") == ( + f"{DEFAULT_API_BASE}/api/v3/wavespeed-ai/z-image/turbo" + ) + + def test_api_base_env_override(self, monkeypatch): + monkeypatch.setenv("WAVESPEED_API_BASE", "https://proxy.internal/") + assert get_api_base(None) == "https://proxy.internal" + assert build_result_url(None, "pred-1") == "https://proxy.internal/api/v3/predictions/pred-1/result" + + def test_explicit_api_base_beats_the_env(self, monkeypatch): + monkeypatch.setenv("WAVESPEED_API_BASE", "https://proxy.internal") + assert get_api_base("https://other.internal") == "https://other.internal" + + def test_empty_model_is_rejected(self): + with pytest.raises(WaveSpeedError, match="model is required"): + build_submit_url(None, "///") + + def test_path_traversal_in_the_model_id_is_rejected(self): + with pytest.raises(ValueError): + build_submit_url(None, "wavespeed-ai/../../admin") + + def test_prediction_id_is_percent_encoded(self): + assert build_result_url("https://api.wavespeed.ai", "a b").endswith("/predictions/a%20b/result") + + +class TestHeaders: + def test_headers_carry_auth_and_channel_attribution(self): + headers = build_headers("sk-test") + assert headers["Authorization"] == "Bearer sk-test" + assert headers["X-Client-Name"] == "litellm" + assert headers["X-Client-Version"] + + def test_api_key_falls_back_to_the_env(self, monkeypatch): + monkeypatch.setenv("WAVESPEED_API_KEY", "sk-env") + assert build_headers(None)["Authorization"] == "Bearer sk-env" + + def test_missing_api_key_raises_401(self, monkeypatch): + monkeypatch.delenv("WAVESPEED_API_KEY", raising=False) + with pytest.raises(WaveSpeedError) as exc_info: + build_headers(None) + assert exc_info.value.status_code == 401 + + +class TestUnwrapEnvelope: + def test_happy_path(self): + raw = httpx.Response(200, json={"code": 200, "message": "ok", "data": {"id": "pred-1"}}) + assert unwrap_envelope(raw)["id"] == "pred-1" + + def test_http_error_surfaces_the_status_code(self): + raw = httpx.Response(503, text="upstream down") + with pytest.raises(WaveSpeedError) as exc_info: + unwrap_envelope(raw) + assert exc_info.value.status_code == 503 + assert "upstream down" in str(exc_info.value) + + def test_non_json_body(self): + raw = httpx.Response(200, text="gateway") + with pytest.raises(WaveSpeedError, match="Could not parse"): + unwrap_envelope(raw) + + def test_non_object_body(self): + raw = httpx.Response(200, json=["not", "an", "envelope"]) + with pytest.raises(WaveSpeedError, match="Unexpected WaveSpeed response body"): + unwrap_envelope(raw) + + def test_platform_error_code_uses_the_platform_message(self): + raw = httpx.Response(200, json={"code": 401, "message": "invalid api key", "data": None}) + with pytest.raises(WaveSpeedError, match="invalid api key"): + unwrap_envelope(raw) + + def test_platform_error_code_without_a_message(self): + raw = httpx.Response(200, json={"code": 500, "data": None}) + with pytest.raises(WaveSpeedError, match="WaveSpeed returned code 500"): + unwrap_envelope(raw) + + def test_missing_data(self): + raw = httpx.Response(200, json={"code": 200, "message": "ok"}) + with pytest.raises(WaveSpeedError, match="missing `data`"): + unwrap_envelope(raw) + + +class TestPredictionHelpers: + def test_missing_prediction_id_raises(self): + with pytest.raises(WaveSpeedError, match="missing a prediction id"): + get_prediction_id({"status": "created"}) + + def test_get_outputs_defaults_to_empty(self): + assert get_outputs({"status": "completed"}) == () + assert get_outputs({"status": "completed", "outputs": None}) == () + assert get_outputs({"status": "completed", "outputs": ["a"]}) == ["a"] + + @pytest.mark.parametrize( + "status, expected", [("completed", "done"), ("created", "pending"), ("processing", "pending")] + ) + def test_poll_outcome_non_terminal_and_success(self, status, expected): + assert poll_outcome({"status": status}) == expected + + @pytest.mark.parametrize("status", ["failed", "cancelled", "timeout"]) + def test_poll_outcome_terminal_failures(self, status): + with pytest.raises(WaveSpeedError, match=status): + poll_outcome({"status": status, "error": "boom"}) + + def test_poll_outcome_failure_without_an_error_detail(self): + with pytest.raises(WaveSpeedError, match="no error detail returned"): + poll_outcome({"status": "failed"}) + + def test_status_mapping(self): + assert map_status_to_openai("processing") == "in_progress" + assert map_status_to_openai("cancelled") == "failed" + assert map_status_to_openai("brand-new-status") == "queued" + + +class TestPayloadHelpers: + def test_to_request_payload_accepts_mappings_and_pairs(self): + assert to_request_payload({"a": 1}) == {"a": 1} + assert to_request_payload((("a", 1), ("b", 2))) == {"a": 1, "b": 2} + + def test_optional_helpers_drop_none(self): + assert optional_pair("a", 1) == (("a", 1),) + assert optional_pair("a", None) == () + assert dict(optional_entry("a", 1)) == {"a": 1} + assert dict(optional_entry("a", None)) == {} diff --git a/tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py b/tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py index 01edaf38001..8c81d7deee1 100644 --- a/tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py +++ b/tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py @@ -1,5 +1,8 @@ """Tests for WaveSpeed AI provider registration across chat, image, and video surfaces.""" +import httpx +import respx + import litellm from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.llms.openai_like.json_loader import JSONProviderRegistry @@ -63,3 +66,63 @@ def test_image_and_video_configs_are_resolved(): assert isinstance(image_config, WaveSpeedImageGenerationConfig) assert isinstance(video_config, WaveSpeedVideoConfig) + + +def test_api_base_autodetects_the_provider(monkeypatch): + """Pointing api_base at the WaveSpeed chat host is enough to route there.""" + monkeypatch.setenv("WAVESPEED_API_KEY", "sk-env-key") + + _, provider, api_key, _ = get_llm_provider( + model="glm-5", + custom_llm_provider=None, + api_base="https://llm.wavespeed.ai/v1", + api_key=None, + ) + + assert provider == "wavespeed" + assert api_key == "sk-env-key" + + +def test_api_key_and_base_resolved_from_env(monkeypatch): + monkeypatch.setenv("WAVESPEED_API_KEY", "sk-env-key") + monkeypatch.setenv("WAVESPEED_API_BASE", "https://proxy.internal/v1") + + _, provider, api_key, api_base = get_llm_provider( + model="wavespeed/deepseek/deepseek-v4-flash", + custom_llm_provider=None, + api_base=None, + api_key=None, + ) + + assert provider == "wavespeed" + assert api_key == "sk-env-key" + assert api_base == "https://proxy.internal/v1" + + +@respx.mock +def test_image_generation_routes_through_the_public_sdk(monkeypatch): + """litellm.image_generation dispatches wavespeed models to the polling handler.""" + monkeypatch.setenv("WAVESPEED_API_KEY", "sk-test") + monkeypatch.delenv("WAVESPEED_API_BASE", raising=False) + monkeypatch.setattr("litellm.llms.wavespeed.image_generation.handler.DEFAULT_POLLING_INTERVAL", 0) + + model = "wavespeed-ai/z-image/turbo" + output_url = "https://cdn.wavespeed.ai/pred-123.png" + envelope = {"code": 200, "message": "ok", "data": {"id": "pred-123", "status": "created"}} + completed = { + "code": 200, + "message": "ok", + "data": {"id": "pred-123", "status": "completed", "outputs": [output_url]}, + } + + submit = respx.post(f"https://api.wavespeed.ai/api/v3/{model}").mock( + return_value=httpx.Response(200, json=envelope) + ) + respx.get("https://api.wavespeed.ai/api/v3/predictions/pred-123/result").mock( + return_value=httpx.Response(200, json=completed) + ) + + response = litellm.image_generation(model=f"wavespeed/{model}", prompt="a red panda") + + assert response.data[0].url == output_url + assert submit.call_count == 1 diff --git a/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py b/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py index 0a82182bdda..105eeeac384 100644 --- a/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py +++ b/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py @@ -4,6 +4,9 @@ from unittest.mock import Mock import httpx import pytest +import respx + +import litellm from litellm.llms.wavespeed.common_utils import WaveSpeedError from litellm.llms.wavespeed.videos.transformation import WaveSpeedVideoConfig @@ -104,10 +107,130 @@ class TestWaveSpeedVideoTransformation: assert headers["Authorization"] == "Bearer sk-test" assert headers["X-Client-Name"] == "litellm" - def test_unsupported_surfaces_raise_not_implemented(self): + +class TestWaveSpeedVideoContentDownload: + def setup_method(self): + self.config = WaveSpeedVideoConfig() + self.logging_obj = Mock() + + @respx.mock + def test_content_response_downloads_the_output(self): + download = respx.get(OUTPUT_URL).mock(return_value=httpx.Response(200, content=b"mp4-bytes")) + + content = self.config.transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])), + logging_obj=self.logging_obj, + ) + + assert content == b"mp4-bytes" + assert download.call_count == 1 + + @respx.mock + def test_content_response_raises_on_a_dead_output_url(self): + respx.get(OUTPUT_URL).mock(return_value=httpx.Response(404)) + + with pytest.raises(httpx.HTTPStatusError): + self.config.transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])), + logging_obj=self.logging_obj, + ) + + @pytest.mark.asyncio + @respx.mock + async def test_async_content_response_downloads_the_output(self, monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + download = respx.get(OUTPUT_URL).mock(return_value=httpx.Response(200, content=b"mp4-bytes")) + + content = await self.config.async_transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])), + logging_obj=self.logging_obj, + ) + + assert content == b"mp4-bytes" + assert download.call_count == 1 + + @pytest.mark.asyncio + @respx.mock + async def test_async_content_response_raises_while_still_processing(self, monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + with pytest.raises(WaveSpeedError, match="still created"): + await self.config.async_transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("created")), + logging_obj=self.logging_obj, + ) + + +class TestWaveSpeedVideoMisc: + def setup_method(self): + self.config = WaveSpeedVideoConfig() + + def test_get_complete_url_defaults_and_overrides(self): + assert self.config.get_complete_url(MODEL, None, {}) == API_BASE + assert self.config.get_complete_url(MODEL, "https://proxy.internal/", {}) == "https://proxy.internal" + + def test_status_retrieve_response_encodes_the_provider_into_the_id(self): + video = self.config.transform_video_status_retrieve_response( + raw_response=httpx.Response(200, json=prediction("processing")), + logging_obj=Mock(), + custom_llm_provider="wavespeed", + ) + + assert video.status == "in_progress" + assert extract_original_video_id(video.id) == "pred-123" + + def test_unknown_status_falls_back_to_queued(self): + video = self.config.transform_video_status_retrieve_response( + raw_response=httpx.Response(200, json=prediction("something-new")), + logging_obj=Mock(), + ) + assert video.status == "queued" + + @pytest.mark.parametrize( + "created_at, expected", + [("2026-08-20T10:00:00Z", 1787220000), (None, 0), ("", 0), ("not-a-date", 0)], + ) + def test_created_at_parsing(self, created_at, expected): + payload = envelope({"id": "pred-123", "status": "created", "created_at": created_at}) + video = self.config.transform_video_status_retrieve_response( + raw_response=httpx.Response(200, json=payload), logging_obj=Mock() + ) + assert video.created_at == expected + + @pytest.mark.parametrize( + "seconds, expected_duration", + [("5", 5), (5, 5), (5.9, 5), ("5.9", 5), (None, None), ("abc", None), (True, None), (object(), None)], + ) + def test_seconds_coercion(self, seconds, expected_duration): + mapped = self.config.map_openai_params({"seconds": seconds}, MODEL, False) + assert mapped.get("duration") == expected_duration + + def test_size_without_an_x_is_left_alone(self): + assert "size" not in self.config.map_openai_params({"size": "720p"}, MODEL, False) + + def test_get_error_class(self): + error = self.config.get_error_class("boom", 503, {}) + assert isinstance(error, WaveSpeedError) + assert error.status_code == 503 + + @pytest.mark.parametrize( + "call", + [ + lambda c: c.transform_video_remix_request("v", "p", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_remix_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_list_request(API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_list_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_delete_request("v", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_delete_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_create_character_request("n", object(), API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_create_character_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_get_character_request("c", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_get_character_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_edit_request("p", "v", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_edit_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_extension_request("p", "v", "5", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_extension_response(httpx.Response(200), Mock()), + ], + ) + def test_unsupported_surfaces_raise_not_implemented(self, call): with pytest.raises(NotImplementedError): - self.config.transform_video_list_request(API_BASE, GenericLiteLLMParams(), {}) - with pytest.raises(NotImplementedError): - self.config.transform_video_delete_request("pred-123", API_BASE, GenericLiteLLMParams(), {}) - with pytest.raises(NotImplementedError): - self.config.transform_video_remix_request("pred-123", "x", API_BASE, GenericLiteLLMParams(), {}) + call(self.config)