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.
This commit is contained in:
chengzeyi 2026-08-20 11:51:47 +00:00
parent e6c01e49cb
commit 93eed536d6
4 changed files with 474 additions and 6 deletions

View file

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

View file

@ -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="<html>gateway</html>")
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)) == {}

View file

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

View file

@ -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)