mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
e6c01e49cb
commit
93eed536d6
4 changed files with 474 additions and 6 deletions
|
|
@ -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
|
||||
|
|
|
|||
148
tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py
Normal file
148
tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py
Normal 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)) == {}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue