From e6c01e49cbfaade77cc005f2021cecf57f18f5eb Mon Sep 17 00:00:00 2001 From: chengzeyi Date: Thu, 20 Aug 2026 11:23:37 +0000 Subject: [PATCH] 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. --- litellm/constants.py | 2 + litellm/images/main.py | 14 + .../get_llm_provider_logic.py | 3 + litellm/llms/openai_like/providers.json | 7 + litellm/llms/wavespeed/__init__.py | 0 litellm/llms/wavespeed/common_utils.py | 187 +++++++++ .../wavespeed/image_generation/__init__.py | 0 .../wavespeed/image_generation/handler.py | 269 ++++++++++++ .../image_generation/transformation.py | 176 ++++++++ litellm/llms/wavespeed/videos/__init__.py | 0 .../llms/wavespeed/videos/transformation.py | 395 ++++++++++++++++++ .../provider_endpoints_support_backup.json | 17 + litellm/types/utils.py | 1 + litellm/utils.py | 10 + provider_endpoints_support.json | 17 + tests/test_litellm/llms/wavespeed/__init__.py | 0 .../wavespeed/image_generation/__init__.py | 0 .../test_wavespeed_image_generation.py | 207 +++++++++ .../llms/wavespeed/test_wavespeed_provider.py | 65 +++ .../llms/wavespeed/videos/__init__.py | 0 .../test_wavespeed_video_transformation.py | 113 +++++ 21 files changed, 1483 insertions(+) create mode 100644 litellm/llms/wavespeed/__init__.py create mode 100644 litellm/llms/wavespeed/common_utils.py create mode 100644 litellm/llms/wavespeed/image_generation/__init__.py create mode 100644 litellm/llms/wavespeed/image_generation/handler.py create mode 100644 litellm/llms/wavespeed/image_generation/transformation.py create mode 100644 litellm/llms/wavespeed/videos/__init__.py create mode 100644 litellm/llms/wavespeed/videos/transformation.py create mode 100644 tests/test_litellm/llms/wavespeed/__init__.py create mode 100644 tests/test_litellm/llms/wavespeed/image_generation/__init__.py create mode 100644 tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py create mode 100644 tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py create mode 100644 tests/test_litellm/llms/wavespeed/videos/__init__.py create mode 100644 tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py diff --git a/litellm/constants.py b/litellm/constants.py index facfc6f7c19..4087fd0045f 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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", diff --git a/litellm/images/main.py b/litellm/images/main.py index ae4818b1967..248e09d88f0 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -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: diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index dbb40913e14..d62c3cce280 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -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}") diff --git a/litellm/llms/openai_like/providers.json b/litellm/llms/openai_like/providers.json index 164100d4194..7ab244e3d3e 100644 --- a/litellm/llms/openai_like/providers.json +++ b/litellm/llms/openai_like/providers.json @@ -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"] } } diff --git a/litellm/llms/wavespeed/__init__.py b/litellm/llms/wavespeed/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/wavespeed/common_utils.py b/litellm/llms/wavespeed/common_utils.py new file mode 100644 index 00000000000..e39f78b89ac --- /dev/null +++ b/litellm/llms/wavespeed/common_utils.py @@ -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") diff --git a/litellm/llms/wavespeed/image_generation/__init__.py b/litellm/llms/wavespeed/image_generation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/wavespeed/image_generation/handler.py b/litellm/llms/wavespeed/image_generation/handler.py new file mode 100644 index 00000000000..de79ea2ae6f --- /dev/null +++ b/litellm/llms/wavespeed/image_generation/handler.py @@ -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() diff --git a/litellm/llms/wavespeed/image_generation/transformation.py b/litellm/llms/wavespeed/image_generation/transformation.py new file mode 100644 index 00000000000..ff168e6576c --- /dev/null +++ b/litellm/llms/wavespeed/image_generation/transformation.py @@ -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) diff --git a/litellm/llms/wavespeed/videos/__init__.py b/litellm/llms/wavespeed/videos/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/wavespeed/videos/transformation.py b/litellm/llms/wavespeed/videos/transformation.py new file mode 100644 index 00000000000..4362cda1991 --- /dev/null +++ b/litellm/llms/wavespeed/videos/transformation.py @@ -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) diff --git a/litellm/provider_endpoints_support_backup.json b/litellm/provider_endpoints_support_backup.json index dd7712aabca..96f3f3ea9fd 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -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", diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 96b9343353d..cb1fde61a09 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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" diff --git a/litellm/utils.py b/litellm/utils.py index d1b0cb882ac..52e55aecae2 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index ec0b1c27344..9f38b150a4e 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -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", diff --git a/tests/test_litellm/llms/wavespeed/__init__.py b/tests/test_litellm/llms/wavespeed/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/wavespeed/image_generation/__init__.py b/tests/test_litellm/llms/wavespeed/image_generation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py b/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py new file mode 100644 index 00000000000..6764999d408 --- /dev/null +++ b/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py @@ -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, + ) diff --git a/tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py b/tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py new file mode 100644 index 00000000000..01edaf38001 --- /dev/null +++ b/tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py @@ -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) diff --git a/tests/test_litellm/llms/wavespeed/videos/__init__.py b/tests/test_litellm/llms/wavespeed/videos/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py b/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py new file mode 100644 index 00000000000..0a82182bdda --- /dev/null +++ b/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py @@ -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(), {})