feat(wavespeed): add WaveSpeed AI image and video generation

WaveSpeed AI serves image and video models behind one asynchronous prediction
API: POST /api/v3/{model} submits a task and GET /api/v3/predictions/{id}/result
reports status and output URLs, both wrapped in a {code, message, data} envelope.

Image generation follows the Black Forest Labs pattern: the config transforms
data and the handler owns the submit-then-poll HTTP flow. The submit POST is
issued exactly once and is never retried, since every submission is a billable
task; poll GETs tolerate up to 5 consecutive transport failures.

Video generation maps onto the OpenAI video contract through BaseVideoConfig, so
create, status retrieve, and content download each hit the prediction API and the
client drives the polling.

Chat is registered as a JSON OpenAI-compatible provider against
https://llm.wavespeed.ai/v1. WaveSpeed model ids already carry an upstream
provider prefix (anthropic/claude-opus-4.8), so only the leading wavespeed/ is
stripped, which a test pins.
This commit is contained in:
chengzeyi 2026-08-20 11:23:37 +00:00
parent 6d47468dae
commit e6c01e49cb
21 changed files with 1483 additions and 0 deletions

View file

@ -757,6 +757,7 @@ openai_compatible_endpoints: Final[list] = [
"https://api.libertai.io/v1",
"https://pinstripes.io/v1",
"https://api.meta.ai/v1",
"https://llm.wavespeed.ai/v1",
]
@ -824,6 +825,7 @@ openai_compatible_providers: Final[list] = [
"pinstripes", # Pinstripes - JSON-configured provider
"darkbloom",
"meta", # Meta Model API (Muse Spark) - JSON-configured provider
"wavespeed", # WaveSpeed AI - JSON-configured provider
]
openai_text_completion_compatible_providers: Final[list] = [ # providers that support `/v1/completions`
"together_ai",

View file

@ -33,6 +33,7 @@ from openai.types.audio.transcription_create_params import FileTypes
# BFL handlers
from litellm.llms.black_forest_labs.image_edit.handler import bfl_image_edit
from litellm.llms.black_forest_labs.image_generation.handler import bfl_image_generation
from litellm.llms.wavespeed.image_generation.handler import wavespeed_image_generation
from litellm.main import (
azure_chat_completions,
base_llm_aiohttp_handler,
@ -405,6 +406,19 @@ def image_generation(
timeout=timeout,
client=client,
)
elif custom_llm_provider == "wavespeed":
return wavespeed_image_generation.image_generation(
model=model,
prompt=prompt,
model_response=model_response,
optional_params=optional_params,
litellm_params=litellm_params_dict,
logging_obj=litellm_logging_obj,
timeout=timeout,
extra_headers=extra_headers,
client=client,
aimg_generation=aimg_generation,
)
elif custom_llm_provider == "black_forest_labs":
# Route to BFL-specific handler (polling required)
if model is None:

View file

@ -349,6 +349,9 @@ def get_llm_provider(
elif endpoint == "https://api.meta.ai/v1":
custom_llm_provider = "meta"
dynamic_api_key = get_secret_str("META_API_KEY")
elif endpoint == "https://llm.wavespeed.ai/v1":
custom_llm_provider = "wavespeed"
dynamic_api_key = get_secret_str("WAVESPEED_API_KEY")
if api_base is not None and not isinstance(api_base, str):
raise Exception(f"api base needs to be a string. api_base={api_base}")

View file

@ -183,5 +183,12 @@
"max_completion_tokens": "max_tokens"
},
"supported_endpoints": ["/v1/chat/completions", "/v1/responses", "/v1/embeddings"]
},
"wavespeed": {
"base_url": "https://llm.wavespeed.ai/v1",
"api_key_env": "WAVESPEED_API_KEY",
"api_base_env": "WAVESPEED_API_BASE",
"base_class": "openai_gpt",
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"]
}
}

View file

View file

@ -0,0 +1,187 @@
"""
WaveSpeed AI common utilities.
WaveSpeed exposes every media model behind one asynchronous prediction API:
- ``POST {api_base}/api/v3/{model}`` submits a task and returns its id
- ``GET {api_base}/api/v3/predictions/{id}/result`` returns the task status and outputs
Both responses are wrapped in the platform envelope ``{"code": ..., "message": ..., "data": ...}``.
API Reference: https://wavespeed.ai/docs
"""
from collections.abc import Iterable, Mapping, Sequence
from types import MappingProxyType
from typing import Final, Literal, TypedDict
import httpx
from typing_extensions import ReadOnly
from litellm._version import version as litellm_version
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.secret_managers.main import get_secret_str
class WaveSpeedError(BaseLLMException):
"""Exception class for WaveSpeed AI API errors."""
DEFAULT_API_BASE: Final = "https://api.wavespeed.ai"
DEFAULT_POLLING_INTERVAL: Final = 1.0
DEFAULT_MAX_POLLING_TIME: Final = 600
MAX_CONSECUTIVE_POLL_FAILURES: Final = 5
SUCCESS_STATUS: Final = "completed"
FAILURE_STATUSES: Final = frozenset({"failed", "cancelled", "timeout"})
PENDING_STATUSES: Final = frozenset({"created", "processing"})
OPENAI_STATUS_BY_WAVESPEED_STATUS: Final = MappingProxyType(
{
"created": "queued",
"processing": "in_progress",
"completed": "completed",
"failed": "failed",
"cancelled": "failed",
"timeout": "failed",
}
)
class WaveSpeedPrediction(TypedDict, total=False):
id: ReadOnly[str]
model: ReadOnly[str]
status: ReadOnly[str]
outputs: ReadOnly[Sequence[str]]
error: ReadOnly[str]
created_at: ReadOnly[str]
has_nsfw_contents: ReadOnly[Sequence[bool]]
def to_request_payload(
payload: Mapping[str, object] | Iterable[tuple[str, object]],
) -> dict: # mutable-ok: base config contracts return bare `dict`
"""Materialize a read-only payload into the mutable ``dict`` the base config contracts declare."""
return dict(payload) # mutable-ok: base config contracts return bare `dict`
def optional_pair(key: str, value: object) -> tuple[tuple[str, object], ...]:
"""One key/value pair when the value is set, nothing otherwise, for splatting into a payload."""
return ((key, value),) if value is not None else ()
def optional_entry(key: str, value: object) -> Mapping[str, object]:
"""One-entry mapping when the value is set, empty otherwise, for splatting into a payload."""
return MappingProxyType({key: value}) if value is not None else MappingProxyType({})
def get_api_key(api_key: str | None) -> str:
resolved: Final = api_key or get_secret_str("WAVESPEED_API_KEY")
if not resolved:
raise WaveSpeedError(
status_code=401,
message="WaveSpeed API key is required. Set the WAVESPEED_API_KEY environment variable or pass api_key.",
)
return resolved
def get_api_base(api_base: str | None) -> str:
return (api_base or get_secret_str("WAVESPEED_API_BASE") or DEFAULT_API_BASE).rstrip("/")
def build_headers(api_key: str | None) -> Mapping[str, str]:
"""Auth plus the channel-attribution headers every WaveSpeed client sends."""
return MappingProxyType(
{
"Authorization": f"Bearer {get_api_key(api_key)}",
"Content-Type": "application/json",
"X-Client-Name": "litellm",
"X-Client-Version": litellm_version,
}
)
def build_submit_url(api_base: str | None, model: str) -> str:
encoded_model: Final = "/".join(
encode_url_path_segment(segment, field_name="model") for segment in model.split("/") if segment
)
if not encoded_model:
raise WaveSpeedError(status_code=400, message="model is required for WaveSpeed predictions")
return f"{get_api_base(api_base)}/api/v3/{encoded_model}"
def build_result_url(api_base: str | None, prediction_id: str) -> str:
encoded_id: Final = encode_url_path_segment(prediction_id, field_name="prediction_id")
return f"{get_api_base(api_base)}/api/v3/predictions/{encoded_id}/result"
def unwrap_envelope(raw_response: httpx.Response) -> WaveSpeedPrediction:
"""Return ``data`` from a WaveSpeed envelope, raising on transport or platform-level failure."""
if raw_response.status_code >= 400:
raise WaveSpeedError(
status_code=raw_response.status_code,
message=f"WaveSpeed request failed: {raw_response.text}",
headers=raw_response.headers,
)
try:
envelope: Final[object] = raw_response.json()
except ValueError as e:
raise WaveSpeedError(
status_code=raw_response.status_code,
message=f"Could not parse WaveSpeed response: {e}",
headers=raw_response.headers,
)
if not isinstance(envelope, Mapping):
raise WaveSpeedError(
status_code=raw_response.status_code,
message=f"Unexpected WaveSpeed response body: {raw_response.text}",
headers=raw_response.headers,
)
code: Final = envelope.get("code")
if code != 200:
raise WaveSpeedError(
status_code=raw_response.status_code,
message=str(envelope.get("message") or f"WaveSpeed returned code {code}"),
headers=raw_response.headers,
)
data: Final = envelope.get("data")
if not isinstance(data, Mapping):
raise WaveSpeedError(
status_code=raw_response.status_code,
message=f"WaveSpeed response is missing `data`: {raw_response.text}",
headers=raw_response.headers,
)
return WaveSpeedPrediction(**data)
def get_prediction_id(prediction: WaveSpeedPrediction) -> str:
prediction_id: Final = prediction.get("id")
if not prediction_id:
raise WaveSpeedError(status_code=500, message="WaveSpeed submit response is missing a prediction id")
return prediction_id
def get_outputs(prediction: WaveSpeedPrediction) -> Sequence[str]:
return prediction.get("outputs") or ()
def poll_outcome(prediction: WaveSpeedPrediction) -> Literal["done", "pending"]:
"""Classify a polled prediction, raising ``WaveSpeedError`` on a terminal failure."""
status: Final = prediction.get("status", "")
if status == SUCCESS_STATUS:
return "done"
if status in FAILURE_STATUSES:
raise WaveSpeedError(
status_code=400,
message=f"WaveSpeed prediction {status}: {prediction.get('error') or 'no error detail returned'}",
)
return "pending"
def map_status_to_openai(status: str) -> str:
return OPENAI_STATUS_BY_WAVESPEED_STATUS.get(status, "queued")

View file

@ -0,0 +1,269 @@
"""
WaveSpeed AI image generation handler.
WaveSpeed predictions are asynchronous: one submit POST returns a prediction id, then the
result endpoint is polled until the prediction reaches a terminal status.
The submit POST is issued exactly once and is never retried, because every submission is a
billable task and a retry would create a duplicate one. Poll GETs are read-only, so a short
run of connection failures is tolerated before giving up.
"""
import asyncio
import time
from collections.abc import Coroutine, Mapping
from types import MappingProxyType
from typing import Final, NamedTuple
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
_get_httpx_client, # pyright: ignore[reportPrivateUsage] # the shared sync client factory litellm providers use
get_async_httpx_client,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import ImageResponse
from ..common_utils import (
DEFAULT_MAX_POLLING_TIME,
DEFAULT_POLLING_INTERVAL,
MAX_CONSECUTIVE_POLL_FAILURES,
WaveSpeedError,
build_result_url,
get_prediction_id,
poll_outcome,
to_request_payload,
unwrap_envelope,
)
from .transformation import WaveSpeedImageGenerationConfig
class _PreparedRequest(NamedTuple):
headers: Mapping[str, str]
submit_url: str
body: Mapping[str, object]
class _ResolvedParams(NamedTuple):
api_key: str | None
api_base: str | None
litellm_params: Mapping[str, object]
class WaveSpeedImageGeneration:
def __init__(self, config: WaveSpeedImageGenerationConfig | None = None) -> None:
self.config: Final = config or WaveSpeedImageGenerationConfig()
def image_generation(
self,
model: str,
prompt: str,
model_response: ImageResponse,
optional_params: Mapping[str, object],
litellm_params: GenericLiteLLMParams | Mapping[str, object],
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout | None,
extra_headers: Mapping[str, str] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
aimg_generation: bool = False,
) -> "ImageResponse | Coroutine[object, object, ImageResponse]":
if aimg_generation:
return self.async_image_generation(
model=model,
prompt=prompt,
model_response=model_response,
optional_params=optional_params,
litellm_params=litellm_params,
logging_obj=logging_obj,
timeout=timeout,
extra_headers=extra_headers,
client=client if isinstance(client, AsyncHTTPHandler) else None,
)
resolved: Final = _resolve_params(litellm_params)
sync_client: Final = client if isinstance(client, HTTPHandler) else _get_httpx_client()
prepared: Final = self._prepare(model, prompt, resolved, optional_params, extra_headers, logging_obj)
submit_response: Final = sync_client.post(
url=prepared.submit_url,
headers=to_request_payload(prepared.headers),
json=to_request_payload(prepared.body),
timeout=timeout,
)
prediction_id: Final = get_prediction_id(unwrap_envelope(submit_response))
result_url: Final = build_result_url(resolved.api_base, prediction_id)
deadline: Final = time.time() + DEFAULT_MAX_POLLING_TIME
poll_headers: Final = to_request_payload(prepared.headers)
consecutive_failures = 0 # rebind-ok: counts consecutive poll transport failures
while time.time() < deadline:
try:
poll_response = sync_client.get(url=result_url, headers=poll_headers)
except Exception as e: # noqa: BLE001 # any transport failure is retried, never the billable submit
consecutive_failures = _record_poll_failure(consecutive_failures, prediction_id, e)
time.sleep(DEFAULT_POLLING_INTERVAL)
continue
consecutive_failures = 0
if poll_outcome(unwrap_envelope(poll_response)) == "done":
return self._transform(
model, poll_response, model_response, logging_obj, prepared, optional_params, resolved
)
time.sleep(DEFAULT_POLLING_INTERVAL)
raise _timeout_error(prediction_id)
async def async_image_generation(
self,
model: str,
prompt: str,
model_response: ImageResponse,
optional_params: Mapping[str, object],
litellm_params: GenericLiteLLMParams | Mapping[str, object],
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout | None,
extra_headers: Mapping[str, str] | None = None,
client: AsyncHTTPHandler | None = None,
) -> ImageResponse:
resolved: Final = _resolve_params(litellm_params)
async_client: Final = client or get_async_httpx_client(llm_provider=litellm.LlmProviders.WAVESPEED)
prepared: Final = self._prepare(model, prompt, resolved, optional_params, extra_headers, logging_obj)
submit_response: Final = await async_client.post(
url=prepared.submit_url,
headers=to_request_payload(prepared.headers),
json=to_request_payload(prepared.body),
timeout=timeout,
)
prediction_id: Final = get_prediction_id(unwrap_envelope(submit_response))
result_url: Final = build_result_url(resolved.api_base, prediction_id)
deadline: Final = time.time() + DEFAULT_MAX_POLLING_TIME
poll_headers: Final = to_request_payload(prepared.headers)
consecutive_failures = 0 # rebind-ok: counts consecutive poll transport failures
while time.time() < deadline:
try:
poll_response = await async_client.get(url=result_url, headers=poll_headers)
except Exception as e: # noqa: BLE001 # any transport failure is retried, never the billable submit
consecutive_failures = _record_poll_failure(consecutive_failures, prediction_id, e)
await asyncio.sleep(DEFAULT_POLLING_INTERVAL)
continue
consecutive_failures = 0
if poll_outcome(unwrap_envelope(poll_response)) == "done":
return self._transform(
model, poll_response, model_response, logging_obj, prepared, optional_params, resolved
)
await asyncio.sleep(DEFAULT_POLLING_INTERVAL)
raise _timeout_error(prediction_id)
def _prepare(
self,
model: str,
prompt: str,
resolved: _ResolvedParams,
optional_params: Mapping[str, object],
extra_headers: Mapping[str, str] | None,
logging_obj: LiteLLMLoggingObj,
) -> _PreparedRequest:
headers: Final = self.config.validate_environment(
headers=MappingProxyType({**(extra_headers or MappingProxyType({}))}),
model=model,
messages=(),
optional_params=optional_params,
litellm_params=resolved.litellm_params,
api_key=resolved.api_key,
api_base=resolved.api_base,
)
submit_url: Final = self.config.get_complete_url(
api_base=resolved.api_base,
api_key=resolved.api_key,
model=model,
optional_params=optional_params,
litellm_params=resolved.litellm_params,
)
body: Final = self.config.transform_image_generation_request(
model=model,
prompt=prompt,
optional_params=optional_params,
litellm_params=resolved.litellm_params,
headers=headers,
)
logging_obj.pre_call(
input=prompt,
api_key="",
additional_args=to_request_payload(
MappingProxyType({"complete_input_dict": body, "api_base": submit_url, "headers": headers})
),
)
return _PreparedRequest(headers=headers, submit_url=submit_url, body=body)
def _transform(
self,
model: str,
raw_response: httpx.Response,
model_response: ImageResponse,
logging_obj: LiteLLMLoggingObj,
prepared: _PreparedRequest,
optional_params: Mapping[str, object],
resolved: _ResolvedParams,
) -> ImageResponse:
return self.config.transform_image_generation_response(
model=model,
raw_response=raw_response,
model_response=model_response,
logging_obj=logging_obj,
request_data=prepared.body,
optional_params=optional_params,
litellm_params=resolved.litellm_params,
encoding=None,
)
def _resolve_params(litellm_params: GenericLiteLLMParams | Mapping[str, object]) -> _ResolvedParams:
if isinstance(litellm_params, Mapping):
api_key: Final = litellm_params.get("api_key")
api_base: Final = litellm_params.get("api_base")
return _ResolvedParams(
api_key=api_key if isinstance(api_key, str) else None,
api_base=api_base if isinstance(api_base, str) else None,
litellm_params=MappingProxyType(dict(litellm_params)),
)
return _ResolvedParams(
api_key=litellm_params.api_key,
api_base=litellm_params.api_base,
litellm_params=MappingProxyType(dict(litellm_params)),
)
def _record_poll_failure(consecutive_failures: int, prediction_id: str, error: Exception) -> int:
next_count: Final = consecutive_failures + 1
if next_count >= MAX_CONSECUTIVE_POLL_FAILURES:
raise WaveSpeedError(
status_code=500,
message=(
f"WaveSpeed result polling for prediction {prediction_id} failed {next_count} times in a row: {error}"
),
)
verbose_logger.debug("WaveSpeed poll attempt failed (%s/%s): %s", next_count, MAX_CONSECUTIVE_POLL_FAILURES, error)
return next_count
def _timeout_error(prediction_id: str) -> WaveSpeedError:
return WaveSpeedError(
status_code=408,
message=f"WaveSpeed prediction {prediction_id} did not finish within {DEFAULT_MAX_POLLING_TIME} seconds",
)
wavespeed_image_generation: Final = WaveSpeedImageGeneration()

View file

@ -0,0 +1,176 @@
"""
WaveSpeed AI image generation configuration.
Transforms between the OpenAI image generation contract and WaveSpeed's prediction API.
The submit/poll HTTP flow lives in ``handler.py``; this class only transforms data.
API Reference: https://wavespeed.ai/docs
"""
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import (
TYPE_CHECKING,
Any, # noqa: TID251 # runtime stand-in for the TYPE_CHECKING-only logging type
Final,
TypeAlias,
)
import httpx
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIImageGenerationOptionalParams,
)
from litellm.types.utils import ImageObject, ImageResponse
from ..common_utils import (
WaveSpeedError,
build_headers,
build_submit_url,
get_outputs,
to_request_payload,
unwrap_envelope,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj: TypeAlias = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj: TypeAlias = Any
class WaveSpeedImageGenerationConfig(BaseImageGenerationConfig):
"""
Configuration for WaveSpeed AI image generation.
Any WaveSpeed image model id works as-is, e.g. ``wavespeed/bytedance/seedream-v5.0-pro``
or ``wavespeed/wavespeed-ai/z-image/turbo``. Model-specific fields that have no OpenAI
equivalent are passed straight through to the prediction body.
"""
def get_supported_openai_params(
self, model: str
) -> list[OpenAIImageGenerationOptionalParams]: # mutable-ok: base config contract returns `list`
return ["n", "size", "response_format"] # mutable-ok: base config contract returns `list`
def map_openai_params(
self,
non_default_params: Mapping[str, object],
optional_params: Mapping[str, object],
model: str,
drop_params: bool,
) -> dict: # mutable-ok: base config contract returns bare `dict`
return to_request_payload(
MappingProxyType(
{**optional_params, **self._mapped_params(non_default_params, optional_params, drop_params)}
)
)
def _mapped_params(
self,
non_default_params: Mapping[str, object],
optional_params: Mapping[str, object],
drop_params: bool,
) -> Mapping[str, object]:
return MappingProxyType(
{
key: value
for key, value in (
self._map_one(key, value, drop_params)
for key, value in non_default_params.items()
if key not in optional_params and value is not None
)
if key is not None
}
)
def _map_one(self, key: str, value: object, drop_params: bool) -> tuple[str | None, object]:
if key == "size":
return "size", self._map_size(value)
if key == "n":
return ("num_images", value) if isinstance(value, int) and value > 1 else (None, None)
if key == "response_format":
if value != "url" and not drop_params:
raise ValueError(
"WaveSpeed returns hosted image URLs, so only response_format='url' is supported. "
"Set drop_params=True to ignore this parameter."
)
return None, None
return key, value
def _map_size(self, size: object) -> str:
"""WaveSpeed takes ``{width}*{height}`` where OpenAI takes ``{width}x{height}``."""
width, separator, height = str(size).lower().partition("x")
if not separator or not width.isdigit() or not height.isdigit():
raise ValueError(f"Invalid size format: '{size}'. Expected 'WIDTHxHEIGHT' (e.g. '1024x1024').")
return f"{width}*{height}"
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
messages: Sequence[AllMessageValues],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
api_key: str | None = None,
api_base: str | None = None,
) -> dict: # mutable-ok: base config contract returns bare `dict`
return to_request_payload(MappingProxyType({**build_headers(api_key), **headers}))
def get_complete_url(
self,
api_base: str | None,
api_key: str | None,
model: str,
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
stream: bool | None = None,
) -> str:
return build_submit_url(api_base, model)
def transform_image_generation_request(
self,
model: str,
prompt: str,
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
headers: Mapping[str, str],
) -> dict: # mutable-ok: base config contract returns bare `dict`
return to_request_payload(MappingProxyType({"prompt": prompt, **optional_params}))
def transform_image_generation_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ImageResponse,
logging_obj: LiteLLMLoggingObj,
request_data: Mapping[str, object],
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
encoding: object,
api_key: str | None = None,
json_mode: bool | None = None,
) -> ImageResponse:
"""Transform the final polled prediction into an OpenAI image response."""
prediction: Final = unwrap_envelope(raw_response)
outputs: Final = get_outputs(prediction)
if not outputs:
raise WaveSpeedError(
status_code=500,
message=f"WaveSpeed prediction {prediction.get('id', '')} completed without any outputs",
)
images: Final = [ImageObject(url=url, b64_json=None) for url in outputs] # mutable-ok: pydantic field is `list`
model_response.data = images # rebind-ok: the base contract fills in and returns the caller's ImageResponse
return model_response
def get_error_class(
self, error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers
) -> WaveSpeedError:
return WaveSpeedError(status_code=status_code, message=error_message, headers=headers)

View file

@ -0,0 +1,395 @@
"""
WaveSpeed AI video generation configuration.
WaveSpeed video models use the same prediction API as image models:
- ``POST {api_base}/api/v3/{model}`` creates the task
- ``GET {api_base}/api/v3/predictions/{id}/result`` reports status and, once complete, the output URLs
That maps onto the OpenAI video contract as create, status retrieve, and content download,
so the client polls the status endpoint instead of the provider blocking on a poll loop.
API Reference: https://wavespeed.ai/docs
"""
from collections.abc import Mapping
from datetime import datetime
from types import MappingProxyType
from typing import (
TYPE_CHECKING,
Any, # noqa: TID251 # runtime stand-in for the TYPE_CHECKING-only logging type
Final,
Literal,
Never,
TypeAlias,
)
import httpx
from httpx._types import RequestFiles
from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
_get_httpx_client, # pyright: ignore[reportPrivateUsage] # the shared sync client factory litellm providers use
get_async_httpx_client,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.videos.main import VideoCreateOptionalRequestParams, VideoObject
from litellm.types.videos.utils import (
encode_video_id_with_provider,
extract_original_video_id,
)
from ..common_utils import (
PENDING_STATUSES,
WaveSpeedError,
WaveSpeedPrediction,
build_headers,
build_result_url,
build_submit_url,
get_api_base,
get_outputs,
map_status_to_openai,
optional_entry,
optional_pair,
to_request_payload,
unwrap_envelope,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj: TypeAlias = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj: TypeAlias = Any
class _VideoObjectData(TypedDict, extra_items=object):
id: ReadOnly[str]
object: ReadOnly[Literal["video"]]
status: ReadOnly[str]
created_at: ReadOnly[int]
def _parse_created_at(created_at: str | None) -> int:
if not created_at:
return 0
try:
return int(datetime.fromisoformat(created_at.replace("Z", "+00:00")).timestamp())
except ValueError:
return 0
def _to_error(prediction: WaveSpeedPrediction, error: str | None) -> Mapping[str, str] | None:
if not error:
return None
return MappingProxyType({"code": prediction.get("status", "failed"), "message": error})
def _to_video_object(prediction: WaveSpeedPrediction, model: str | None) -> VideoObject:
error: Final = prediction.get("error")
video_data: Final[_VideoObjectData] = {
"id": prediction.get("id", ""),
"object": "video",
"status": map_status_to_openai(prediction.get("status", "")),
"created_at": _parse_created_at(prediction.get("created_at")),
**optional_entry("model", model),
**optional_entry("error", _to_error(prediction, error)),
}
return VideoObject(**video_data)
def _to_int_seconds(seconds: object) -> int | None:
if seconds is None or isinstance(seconds, bool):
return None
if isinstance(seconds, (int, float)):
return int(seconds)
if isinstance(seconds, str):
try:
return int(float(seconds))
except ValueError:
return None
return None
class WaveSpeedVideoConfig(BaseVideoConfig):
"""
Configuration for WaveSpeed AI video generation.
Any WaveSpeed video model id works as-is, e.g.
``wavespeed/bytedance/seedance-2.5/text-to-video``. Model-specific fields that have no
OpenAI equivalent are passed straight through to the prediction body.
"""
def get_supported_openai_params(self, model: str) -> list: # mutable-ok: base contract returns bare `list`
return ["model", "prompt", "input_reference", "seconds", "size", "user", "extra_headers"] # mutable-ok: ditto
def map_openai_params(
self,
video_create_optional_params: VideoCreateOptionalRequestParams,
model: str,
drop_params: bool,
) -> dict: # mutable-ok: base config contract returns bare `dict`
supported: Final = frozenset(self.get_supported_openai_params(model))
size: Final = video_create_optional_params.get("size")
seconds: Final = _to_int_seconds(video_create_optional_params.get("seconds"))
input_reference: Final = video_create_optional_params.get("input_reference")
mapped_size: Final = size.lower().replace("x", "*") if isinstance(size, str) and "x" in size.lower() else None
return to_request_payload(
(
*((k, v) for k, v in video_create_optional_params.items() if k not in supported),
*optional_pair("image", input_reference),
*optional_pair("size", mapped_size),
*optional_pair("duration", seconds),
)
)
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
api_key: str | None = None,
litellm_params: GenericLiteLLMParams | None = None,
) -> dict: # mutable-ok: base config contract returns bare `dict`
resolved_key: Final = api_key or (litellm_params.api_key if litellm_params else None) or litellm.api_key
return to_request_payload(MappingProxyType({**build_headers(resolved_key), **headers}))
def get_complete_url(
self,
model: str,
api_base: str | None,
litellm_params: Mapping[str, object],
) -> str:
return get_api_base(api_base)
def transform_video_create_request(
self,
model: str,
prompt: str,
api_base: str,
video_create_optional_request_params: Mapping[str, object],
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
) -> tuple[dict, RequestFiles, str]: # mutable-ok: base config contract returns bare `dict`
body: Final = to_request_payload(MappingProxyType({"prompt": prompt, **video_create_optional_request_params}))
return body, [], build_submit_url(api_base, model) # mutable-ok: httpx RequestFiles is a list
def transform_video_create_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str | None = None,
request_data: Mapping[str, object] | None = None,
) -> VideoObject:
video_obj: Final = _to_video_object(unwrap_envelope(raw_response), model)
if custom_llm_provider and video_obj.id:
video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, model)
return video_obj
def transform_video_status_retrieve_request(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
) -> tuple[str, dict]: # mutable-ok: base config contract returns bare `dict`
return build_result_url(api_base, extract_original_video_id(video_id)), to_request_payload(MappingProxyType({}))
def transform_video_status_retrieve_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str | None = None,
) -> VideoObject:
video_obj: Final = _to_video_object(unwrap_envelope(raw_response), None)
if custom_llm_provider and video_obj.id:
video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, None)
return video_obj
def transform_video_content_request(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
variant: str | None = None,
) -> tuple[str, dict]: # mutable-ok: base config contract returns bare `dict`
return build_result_url(api_base, extract_original_video_id(video_id)), to_request_payload(MappingProxyType({}))
def _extract_output_url(self, raw_response: httpx.Response) -> str:
prediction: Final = unwrap_envelope(raw_response)
outputs: Final = get_outputs(prediction)
if outputs:
return outputs[0]
status: Final = prediction.get("status", "")
if status in PENDING_STATUSES:
raise WaveSpeedError(
status_code=409,
message=f"WaveSpeed prediction {prediction.get('id', '')} is still {status}. Retry once it completes.",
)
raise WaveSpeedError(
status_code=400,
message=(
f"WaveSpeed prediction {prediction.get('id', '')} has no video output "
f"(status {status or 'unknown'}): {prediction.get('error') or 'no error detail returned'}"
),
)
def transform_video_content_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> bytes:
output_url: Final = self._extract_output_url(raw_response)
httpx_client: Final[HTTPHandler] = _get_httpx_client()
video_response: Final = httpx_client.get(output_url)
video_response.raise_for_status()
return video_response.content
async def async_transform_video_content_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> bytes:
output_url: Final = self._extract_output_url(raw_response)
async_client: Final[AsyncHTTPHandler] = get_async_httpx_client(llm_provider=litellm.LlmProviders.WAVESPEED)
video_response: Final = await async_client.get(output_url)
video_response.raise_for_status()
return video_response.content
def transform_video_remix_request(
self,
video_id: str,
prompt: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
extra_body: Mapping[str, object] | None = None,
) -> Never:
raise NotImplementedError("video remix is not supported for WaveSpeed")
def transform_video_remix_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str | None = None,
) -> Never:
raise NotImplementedError("video remix is not supported for WaveSpeed")
def transform_video_list_request(
self,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
after: str | None = None,
limit: int | None = None,
order: str | None = None,
extra_query: Mapping[str, object] | None = None,
) -> Never:
raise NotImplementedError("video listing is not supported for WaveSpeed")
def transform_video_list_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str | None = None,
) -> Never:
raise NotImplementedError("video listing is not supported for WaveSpeed")
def transform_video_delete_request(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
) -> Never:
raise NotImplementedError("video delete is not supported for WaveSpeed")
def transform_video_delete_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> Never:
raise NotImplementedError("video delete is not supported for WaveSpeed")
def transform_video_create_character_request(
self,
name: str,
video: object,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
) -> Never:
raise NotImplementedError("video create character is not supported for WaveSpeed")
def transform_video_create_character_response(
self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj
) -> Never:
raise NotImplementedError("video create character is not supported for WaveSpeed")
def transform_video_get_character_request(
self,
character_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
) -> Never:
raise NotImplementedError("video get character is not supported for WaveSpeed")
def transform_video_get_character_response(
self, raw_response: httpx.Response, logging_obj: LiteLLMLoggingObj
) -> Never:
raise NotImplementedError("video get character is not supported for WaveSpeed")
def transform_video_edit_request(
self,
prompt: str,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
extra_body: Mapping[str, object] | None = None,
prefetched_source_data: object | None = None,
) -> Never:
raise NotImplementedError("video edit is not supported for WaveSpeed")
def transform_video_edit_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str | None = None,
request_data: Mapping[str, object] | None = None,
) -> Never:
raise NotImplementedError("video edit is not supported for WaveSpeed")
def transform_video_extension_request(
self,
prompt: str,
video_id: str,
seconds: str | None,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: Mapping[str, str],
extra_body: Mapping[str, object] | None = None,
) -> Never:
raise NotImplementedError("video extension is not supported for WaveSpeed")
def transform_video_extension_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str | None = None,
) -> Never:
raise NotImplementedError("video extension is not supported for WaveSpeed")
def get_error_class(
self, error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers
) -> WaveSpeedError:
return WaveSpeedError(status_code=status_code, message=error_message, headers=headers)

View file

@ -2244,6 +2244,23 @@
"interactions": true
}
},
"wavespeed": {
"display_name": "WaveSpeed AI (`wavespeed`)",
"url": "https://docs.litellm.ai/docs/providers/wavespeed",
"endpoints": {
"chat_completions": true,
"messages": false,
"responses": true,
"embeddings": false,
"image_generations": true,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false
}
},
"watsonx_text": {
"display_name": "Watsonx Text (`watsonx_text`)",
"url": "https://docs.litellm.ai/docs/providers/watsonx",

View file

@ -3783,6 +3783,7 @@ class LlmProviders(str, Enum):
PINSTRIPES = "pinstripes"
DARKBLOOM = "darkbloom"
META = "meta"
WAVESPEED = "wavespeed"
LITELLM_AGENT = "litellm_agent"
CURSOR = "cursor"
BEDROCK_MANTLE = "bedrock_mantle"

View file

@ -8847,6 +8847,12 @@ class ProviderConfigManager:
)
return get_runwayml_image_generation_config(model)
elif LlmProviders.WAVESPEED == provider:
from litellm.llms.wavespeed.image_generation.transformation import (
WaveSpeedImageGenerationConfig,
)
return WaveSpeedImageGenerationConfig()
elif LlmProviders.BLACK_FOREST_LABS == provider:
from litellm.llms.black_forest_labs.image_generation import (
get_black_forest_labs_image_generation_config,
@ -8904,6 +8910,10 @@ class ProviderConfigManager:
from litellm.llms.runwayml.videos.transformation import RunwayMLVideoConfig
return RunwayMLVideoConfig()
elif LlmProviders.WAVESPEED == provider:
from litellm.llms.wavespeed.videos.transformation import WaveSpeedVideoConfig
return WaveSpeedVideoConfig()
return None
@staticmethod

View file

@ -2555,6 +2555,23 @@
"interactions": true
}
},
"wavespeed": {
"display_name": "WaveSpeed AI (`wavespeed`)",
"url": "https://docs.litellm.ai/docs/providers/wavespeed",
"endpoints": {
"chat_completions": true,
"messages": false,
"responses": true,
"embeddings": false,
"image_generations": true,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false
}
},
"watsonx_text": {
"display_name": "Watsonx Text (`watsonx_text`)",
"url": "https://docs.litellm.ai/docs/providers/watsonx",

View file

@ -0,0 +1,207 @@
"""Unit tests for the WaveSpeed AI image generation submit/poll flow."""
from unittest.mock import MagicMock
import httpx
import pytest
import respx
import litellm
from litellm.llms.wavespeed.common_utils import WaveSpeedError
from litellm.llms.wavespeed.image_generation.handler import WaveSpeedImageGeneration
from litellm.llms.wavespeed.image_generation.transformation import (
WaveSpeedImageGenerationConfig,
)
from litellm.types.utils import ImageResponse
MODEL = "wavespeed-ai/z-image/turbo"
SUBMIT_URL = f"https://api.wavespeed.ai/api/v3/{MODEL}"
RESULT_URL = "https://api.wavespeed.ai/api/v3/predictions/pred-123/result"
OUTPUT_URL = "https://cdn.wavespeed.ai/pred-123.png"
def envelope(data):
return {"code": 200, "message": "success", "data": data}
def prediction(status, **extra):
return envelope({"id": "pred-123", "model": MODEL, "status": status, **extra})
@pytest.fixture(autouse=True)
def mocked_transport(monkeypatch):
"""No test here may reach the network: respx only intercepts httpx, so pin httpx transport."""
monkeypatch.setattr("litellm.llms.wavespeed.image_generation.handler.DEFAULT_POLLING_INTERVAL", 0)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
@pytest.fixture
def generate():
handler = WaveSpeedImageGeneration()
def run():
return handler.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,
)
return run
@respx.mock
def test_submit_then_poll_until_completed(generate):
submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created")))
poll = respx.get(RESULT_URL).mock(
side_effect=[
httpx.Response(200, json=prediction("processing")),
httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])),
]
)
response = generate()
assert [image.url for image in response.data] == [OUTPUT_URL]
assert submit.call_count == 1
assert poll.call_count == 2
assert submit.calls[0].request.headers["authorization"] == "Bearer sk-test"
assert submit.calls[0].request.headers["x-client-name"] == "litellm"
@pytest.mark.parametrize("status", ["failed", "cancelled", "timeout"])
@respx.mock
def test_terminal_failure_status_raises(generate, status):
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(status, error="nsfw content")))
with pytest.raises(WaveSpeedError) as exc_info:
generate()
assert status in str(exc_info.value)
assert "nsfw content" in str(exc_info.value)
@respx.mock
def test_submit_is_issued_exactly_once_when_polling_fails(generate):
"""A submission is a billable task, so a poll failure must never re-submit it."""
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:
generate()
assert submit.call_count == 1
assert poll.call_count == 5
assert "5 times in a row" in str(exc_info.value)
@respx.mock
def test_transient_poll_failures_are_tolerated(generate):
submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created")))
respx.get(RESULT_URL).mock(
side_effect=[
httpx.ConnectError("connection reset"),
httpx.ConnectError("connection reset"),
httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])),
]
)
response = generate()
assert [image.url for image in response.data] == [OUTPUT_URL]
assert submit.call_count == 1
@respx.mock
def test_non_200_envelope_code_raises(generate):
respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json={"code": 401, "message": "invalid api key"}))
with pytest.raises(WaveSpeedError) as exc_info:
generate()
assert "invalid api key" in str(exc_info.value)
@pytest.mark.asyncio
@respx.mock
async def test_async_submit_then_poll_until_completed():
submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created")))
respx.get(RESULT_URL).mock(
side_effect=[
httpx.Response(200, json=prediction("processing")),
httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])),
]
)
response = 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 [image.url for image in response.data] == [OUTPUT_URL]
assert submit.call_count == 1
class TestWaveSpeedImageGenerationConfig:
def setup_method(self):
self.config = WaveSpeedImageGenerationConfig()
def test_size_is_mapped_to_wavespeed_format(self):
assert self.config.map_openai_params({"size": "1024x1536"}, {}, MODEL, False) == {"size": "1024*1536"}
def test_invalid_size_is_rejected(self):
with pytest.raises(ValueError, match="Invalid size format"):
self.config.map_openai_params({"size": "huge"}, {}, MODEL, False)
def test_n_greater_than_one_maps_to_num_images(self):
assert self.config.map_openai_params({"n": 4}, {}, MODEL, False) == {"num_images": 4}
assert self.config.map_openai_params({"n": 1}, {}, MODEL, False) == {}
def test_b64_response_format_is_rejected_unless_dropped(self):
with pytest.raises(ValueError, match="response_format"):
self.config.map_openai_params({"response_format": "b64_json"}, {}, MODEL, False)
assert self.config.map_openai_params({"response_format": "b64_json"}, {}, MODEL, True) == {}
def test_model_specific_params_pass_through(self):
assert self.config.map_openai_params({"guidance_scale": 3.5}, {}, MODEL, False) == {"guidance_scale": 3.5}
def test_submit_url_keeps_multi_segment_model_ids(self):
assert (
self.config.get_complete_url(None, "sk-test", "bytedance/seedance-2.5/text-to-video", {}, {})
== "https://api.wavespeed.ai/api/v3/bytedance/seedance-2.5/text-to-video"
)
def test_api_base_override(self):
assert (
self.config.get_complete_url("https://proxy.internal/", "sk-test", MODEL, {}, {})
== f"https://proxy.internal/api/v3/{MODEL}"
)
def test_missing_api_key_raises(self, monkeypatch):
monkeypatch.delenv("WAVESPEED_API_KEY", raising=False)
with pytest.raises(WaveSpeedError, match="WAVESPEED_API_KEY"):
self.config.validate_environment({}, MODEL, [], {}, {})
def test_completed_prediction_without_outputs_raises(self):
raw = httpx.Response(200, json=prediction("completed", outputs=[]))
with pytest.raises(WaveSpeedError, match="without any outputs"):
self.config.transform_image_generation_response(
model=MODEL,
raw_response=raw,
model_response=ImageResponse(),
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)

View file

@ -0,0 +1,65 @@
"""Tests for WaveSpeed AI provider registration across chat, image, and video surfaces."""
import litellm
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
from litellm.llms.wavespeed.image_generation.transformation import (
WaveSpeedImageGenerationConfig,
)
from litellm.llms.wavespeed.videos.transformation import WaveSpeedVideoConfig
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
def test_wavespeed_is_a_known_provider():
assert LlmProviders.WAVESPEED.value == "wavespeed"
assert "wavespeed" in litellm.provider_list
def test_chat_json_registry_entry():
from litellm.constants import openai_compatible_providers
config = JSONProviderRegistry.get("wavespeed")
assert config is not None
assert config.base_url == "https://llm.wavespeed.ai/v1"
assert config.api_key_env == "WAVESPEED_API_KEY"
assert config.api_base_env == "WAVESPEED_API_BASE"
assert "wavespeed" in openai_compatible_providers
def test_upstream_model_prefix_is_preserved(monkeypatch):
"""WaveSpeed chat model ids are themselves `{provider}/{model}`, so only the routing prefix is stripped."""
model, provider, api_key, api_base = get_llm_provider(
model="wavespeed/anthropic/claude-opus-4.8",
custom_llm_provider=None,
api_base=None,
api_key="sk-test",
)
assert model == "anthropic/claude-opus-4.8"
assert provider == "wavespeed"
assert api_base == "https://llm.wavespeed.ai/v1"
def test_media_model_routes_to_wavespeed():
model, provider, _, _ = get_llm_provider(
model="wavespeed/bytedance/seedance-2.5/text-to-video",
custom_llm_provider=None,
api_base=None,
api_key="sk-test",
)
assert provider == "wavespeed"
assert model == "bytedance/seedance-2.5/text-to-video"
def test_image_and_video_configs_are_resolved():
image_config = ProviderConfigManager.get_provider_image_generation_config(
model="bytedance/seedream-v5.0-pro", provider=LlmProviders.WAVESPEED
)
video_config = ProviderConfigManager.get_provider_video_config(
model="bytedance/seedance-2.5/text-to-video", provider=LlmProviders.WAVESPEED
)
assert isinstance(image_config, WaveSpeedImageGenerationConfig)
assert isinstance(video_config, WaveSpeedVideoConfig)

View file

@ -0,0 +1,113 @@
"""Tests for WaveSpeed AI video generation transformation."""
from unittest.mock import Mock
import httpx
import pytest
from litellm.llms.wavespeed.common_utils import WaveSpeedError
from litellm.llms.wavespeed.videos.transformation import WaveSpeedVideoConfig
from litellm.types.router import GenericLiteLLMParams
from litellm.types.videos.utils import extract_original_video_id
MODEL = "bytedance/seedance-2.5/text-to-video"
API_BASE = "https://api.wavespeed.ai"
OUTPUT_URL = "https://cdn.wavespeed.ai/pred-123.mp4"
def envelope(data):
return {"code": 200, "message": "success", "data": data}
def prediction(status, **extra):
return envelope({"id": "pred-123", "status": status, "created_at": "2026-08-20T10:00:00Z", **extra})
class TestWaveSpeedVideoTransformation:
def setup_method(self):
self.config = WaveSpeedVideoConfig()
self.logging_obj = Mock()
def test_transform_video_create_request(self):
data, files, url = self.config.transform_video_create_request(
model=MODEL,
prompt="a red panda skateboarding",
api_base=API_BASE,
video_create_optional_request_params={"size": "1280*720", "duration": 5},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert data == {"prompt": "a red panda skateboarding", "size": "1280*720", "duration": 5}
assert files == []
assert url == f"{API_BASE}/api/v3/{MODEL}"
def test_map_openai_params(self):
assert self.config.map_openai_params(
{"size": "1280x720", "seconds": "5", "input_reference": "https://example.com/a.png", "guidance": 3},
MODEL,
False,
) == {"size": "1280*720", "duration": 5, "image": "https://example.com/a.png", "guidance": 3}
def test_create_response_maps_to_queued_video_object(self):
raw = httpx.Response(200, json=prediction("created"))
video = self.config.transform_video_create_response(
model=MODEL, raw_response=raw, logging_obj=self.logging_obj, custom_llm_provider="wavespeed"
)
assert video.status == "queued"
assert video.object == "video"
assert extract_original_video_id(video.id) == "pred-123"
assert video.model == MODEL
def test_status_retrieve_response_maps_terminal_statuses(self):
completed = self.config.transform_video_status_retrieve_response(
raw_response=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])),
logging_obj=self.logging_obj,
)
assert completed.status == "completed"
failed = self.config.transform_video_status_retrieve_response(
raw_response=httpx.Response(200, json=prediction("failed", error="upstream rejected the prompt")),
logging_obj=self.logging_obj,
)
assert failed.status == "failed"
assert failed.error["message"] == "upstream rejected the prompt"
def test_status_retrieve_request_url(self):
url, params = self.config.transform_video_status_retrieve_request(
video_id="pred-123", api_base=API_BASE, litellm_params=GenericLiteLLMParams(), headers={}
)
assert url == f"{API_BASE}/api/v3/predictions/pred-123/result"
assert params == {}
def test_content_request_url(self):
url, params = self.config.transform_video_content_request(
video_id="pred-123", api_base=API_BASE, litellm_params=GenericLiteLLMParams(), headers={}
)
assert url == f"{API_BASE}/api/v3/predictions/pred-123/result"
assert params == {}
def test_content_download_raises_while_still_processing(self):
raw = httpx.Response(200, json=prediction("processing"))
with pytest.raises(WaveSpeedError, match="still processing"):
self.config.transform_video_content_response(raw_response=raw, logging_obj=self.logging_obj)
def test_content_download_raises_on_failed_prediction(self):
raw = httpx.Response(200, json=prediction("failed", error="upstream rejected the prompt"))
with pytest.raises(WaveSpeedError, match="upstream rejected the prompt"):
self.config.transform_video_content_response(raw_response=raw, logging_obj=self.logging_obj)
def test_validate_environment_sets_bearer_and_attribution_headers(self):
headers = self.config.validate_environment({}, MODEL, api_key="sk-test")
assert headers["Authorization"] == "Bearer sk-test"
assert headers["X-Client-Name"] == "litellm"
def test_unsupported_surfaces_raise_not_implemented(self):
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(), {})