diff --git a/litellm/constants.py b/litellm/constants.py index d75f9cbd371..5672899cc03 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -792,6 +792,7 @@ openai_compatible_endpoints: Final[list] = [ "https://api.meta.ai/v1", "https://api.cognition.ai/v1", "https://api.scx.ai/v1", + "https://llm.wavespeed.ai/v1", ] @@ -861,6 +862,7 @@ openai_compatible_providers: Final[list] = [ "meta", # Meta Model API (Muse Spark) - JSON-configured provider "cognition", "scx-ai", + "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 1688087c2da..e66cb1d7904 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -34,6 +34,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, @@ -406,6 +407,21 @@ def image_generation( timeout=timeout, client=client, ) + elif custom_llm_provider == "wavespeed": + litellm_params_dict["api_key"] = api_key or dynamic_api_key + litellm_params_dict["api_base"] = api_base or litellm.api_base + 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/llms/openai_like/providers.json b/litellm/llms/openai_like/providers.json index a458a209ea9..481e1c2f460 100644 --- a/litellm/llms/openai_like/providers.json +++ b/litellm/llms/openai_like/providers.json @@ -200,5 +200,12 @@ "temperature_max": 1.99 }, "supported_endpoints": ["/v1/chat/completions"] + }, + "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..c202cd36b2e --- /dev/null +++ b/litellm/llms/wavespeed/common_utils.py @@ -0,0 +1,258 @@ +""" +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 +""" + +import base64 +import mimetypes +import os +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" +CHAT_API_BASE: Final = "https://llm.wavespeed.ai/v1" +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({}) + + +_MAGIC_BYTE_MIME_TYPES: Final = ( + (b"\x89PNG\r\n\x1a\n", "image/png"), + (b"\xff\xd8\xff", "image/jpeg"), + (b"GIF87a", "image/gif"), + (b"GIF89a", "image/gif"), +) + + +def _sniff_mime_type(payload: bytes) -> str: + if payload[:4] == b"RIFF" and payload[8:12] == b"WEBP": + return "image/webp" + for magic, mime_type in _MAGIC_BYTE_MIME_TYPES: + if payload.startswith(magic): + return mime_type + raise WaveSpeedError( + status_code=400, + message="Could not determine the media type of the reference. Pass a URL, a data URI, or a named file.", + ) + + +def _to_data_uri(payload: bytes, filename: str | None) -> str: + guessed: Final = mimetypes.guess_type(filename)[0] if filename else None + mime_type: Final = guessed or _sniff_mime_type(payload) + return f"data:{mime_type};base64,{base64.b64encode(payload).decode()}" + + +def _reference_to_uri(reference: object, filename: str | None) -> str: + if isinstance(reference, str): + return reference + if isinstance(reference, (bytes, bytearray)): + return _to_data_uri(bytes(reference), filename) + if isinstance(reference, os.PathLike): + path: Final = os.fspath(reference) + with open(path, "rb") as handle: + return _to_data_uri(handle.read(), filename or str(path)) + read: Final = getattr(reference, "read", None) + if callable(read): + payload: Final = read() + if not isinstance(payload, bytes): + raise WaveSpeedError(status_code=400, message="Reference file handle must be opened in binary mode") + return _to_data_uri(payload, filename or getattr(reference, "name", None)) + raise WaveSpeedError(status_code=400, message=f"Unsupported reference type: {type(reference).__name__}") + + +def to_reference_uri(reference: object) -> str: + """Normalize an OpenAI ``input_reference`` into something a JSON body can carry. + + The shared video contract accepts URLs, raw bytes, paths, file handles and + ``(filename, content)`` tuples, but WaveSpeed submits predictions as JSON, so + anything that is not already a URL or data URI has to be inlined as one. + """ + if not isinstance(reference, tuple): + return _reference_to_uri(reference, None) + if len(reference) < 2 or reference[1] is None: + raise WaveSpeedError(status_code=400, message="Reference tuple is missing its content") + supplied_name: Final = reference[0] + return _reference_to_uri(reference[1], str(supplied_name) if supplied_name else None) + + +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: + """Resolve the base URL for the prediction API. + + Chat and media live on different hosts but share the ``wavespeed`` provider slug, so + provider resolution and ``WAVESPEED_API_BASE`` can both hand this the chat base. That + value would build an unreachable prediction URL, so it falls back to the media default. + A self-hosted base is any other value and is honored as-is. + """ + resolved: Final = (api_base or get_secret_str("WAVESPEED_API_BASE") or DEFAULT_API_BASE).rstrip("/") + return DEFAULT_API_BASE if resolved == CHAT_API_BASE else resolved + + +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..1d5cf60171c --- /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, timeout=timeout) + 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, timeout=timeout) + 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..a5878fc7fcd --- /dev/null +++ b/litellm/llms/wavespeed/videos/transformation.py @@ -0,0 +1,389 @@ +""" +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.litellm_core_utils.url_utils import async_safe_get, safe_get +from litellm.llms.base_llm.videos.transformation import BaseVideoConfig +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_reference_uri, + 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", to_reference_uri(input_reference) if input_reference else None), + *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) + video_response: Final = safe_get(litellm.module_level_client, 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) + video_response: Final = await async_safe_get(litellm.module_level_aclient, 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 86c14fb4cd8..c77ce38d152 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -2278,6 +2278,24 @@ "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, + "video_generations": true + } + }, "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 a7629fb2488..35a902ff07e 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3824,6 +3824,7 @@ class LlmProviders(str, Enum): SCX_AI = "scx-ai" 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 9cab81e1ba7..3c730e95de2 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8981,6 +8981,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, @@ -9038,6 +9044,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 1d8d374c2c4..0cffcf0b78b 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2589,6 +2589,24 @@ "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, + "video_generations": true + } + }, "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..08f618fdb16 --- /dev/null +++ b/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py @@ -0,0 +1,368 @@ +"""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.router import GenericLiteLLMParams +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, + ) + + +@pytest.fixture +def zero_poll_budget(monkeypatch): + """Make the polling deadline expire immediately so the timeout path is reachable.""" + monkeypatch.setattr("litellm.llms.wavespeed.image_generation.handler.DEFAULT_MAX_POLLING_TIME", 0) + + +@respx.mock +def test_sync_poll_timeout(generate, zero_poll_budget): + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + poll = respx.get(RESULT_URL).mock(return_value=httpx.Response(200, json=prediction("processing"))) + + with pytest.raises(WaveSpeedError) as exc_info: + generate() + + assert exc_info.value.status_code == 408 + assert "did not finish within" in str(exc_info.value) + assert submit.call_count == 1 + assert poll.call_count == 0 + + +@pytest.mark.asyncio +@respx.mock +async def test_async_poll_timeout(zero_poll_budget): + respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + + with pytest.raises(WaveSpeedError) as exc_info: + await WaveSpeedImageGeneration().async_image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params={"api_key": "sk-test", "api_base": None}, + logging_obj=MagicMock(), + timeout=None, + ) + + assert exc_info.value.status_code == 408 + + +@pytest.mark.asyncio +@respx.mock +async def test_async_submit_is_issued_exactly_once_when_polling_fails(): + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + poll = respx.get(RESULT_URL).mock(side_effect=httpx.ConnectError("connection reset")) + + with pytest.raises(WaveSpeedError) as exc_info: + await WaveSpeedImageGeneration().async_image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params={"api_key": "sk-test", "api_base": None}, + logging_obj=MagicMock(), + timeout=None, + ) + + assert submit.call_count == 1 + assert poll.call_count == 5 + assert "5 times in a row" in str(exc_info.value) + + +@pytest.mark.asyncio +@respx.mock +async def test_aimg_generation_flag_dispatches_to_the_async_path(): + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + respx.get(RESULT_URL).mock(return_value=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL]))) + + pending = WaveSpeedImageGeneration().image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params={"api_key": "sk-test", "api_base": None}, + logging_obj=MagicMock(), + timeout=None, + aimg_generation=True, + ) + + response = await pending + assert [image.url for image in response.data] == [OUTPUT_URL] + assert submit.call_count == 1 + + +@respx.mock +def test_litellm_params_object_is_accepted(monkeypatch): + """images/main.py can hand the handler a GenericLiteLLMParams rather than a dict.""" + monkeypatch.delenv("WAVESPEED_API_BASE", raising=False) + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + respx.get(RESULT_URL).mock(return_value=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL]))) + + response = WaveSpeedImageGeneration().image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params=GenericLiteLLMParams(api_key="sk-test"), + logging_obj=MagicMock(), + timeout=None, + ) + + assert [image.url for image in response.data] == [OUTPUT_URL] + assert submit.calls[0].request.headers["authorization"] == "Bearer sk-test" + + +@respx.mock +def test_extra_headers_are_merged_and_cannot_be_dropped(generate): + submit = respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + respx.get(RESULT_URL).mock(return_value=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL]))) + + WaveSpeedImageGeneration().image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params={"api_key": "sk-test", "api_base": None}, + logging_obj=MagicMock(), + timeout=None, + extra_headers={"X-Trace-Id": "abc123"}, + ) + + assert submit.calls[0].request.headers["x-trace-id"] == "abc123" + assert submit.calls[0].request.headers["x-client-name"] == "litellm" + + +def test_supported_openai_params_and_error_class(): + config = WaveSpeedImageGenerationConfig() + assert config.get_supported_openai_params(MODEL) == ["n", "size", "response_format"] + + error = config.get_error_class("boom", 503, {}) + assert isinstance(error, WaveSpeedError) + assert error.status_code == 503 + + +@respx.mock +def test_poll_requests_honour_the_caller_timeout(generate): + """A short caller timeout must not be replaced by the client's multi-minute default.""" + respx.post(SUBMIT_URL).mock(return_value=httpx.Response(200, json=prediction("created"))) + poll = respx.get(RESULT_URL).mock( + return_value=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])) + ) + + WaveSpeedImageGeneration().image_generation( + model=MODEL, + prompt="a red panda", + model_response=ImageResponse(), + optional_params={}, + litellm_params={"api_key": "sk-test", "api_base": None}, + logging_obj=MagicMock(), + timeout=2.5, + ) + + assert poll.call_count == 1 + assert poll.calls[0].request.extensions.get("timeout") == { + "connect": 2.5, + "read": 2.5, + "write": 2.5, + "pool": 2.5, + } diff --git a/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py b/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py new file mode 100644 index 00000000000..42067a551fd --- /dev/null +++ b/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py @@ -0,0 +1,227 @@ +"""Unit tests for the WaveSpeed AI envelope parsing and URL helpers.""" + +import base64 + +import httpx +import pytest + +from litellm.llms.wavespeed.common_utils import ( + CHAT_API_BASE, + DEFAULT_API_BASE, + WaveSpeedError, + build_headers, + build_result_url, + build_submit_url, + get_api_base, + get_outputs, + get_prediction_id, + map_status_to_openai, + optional_entry, + optional_pair, + poll_outcome, + to_reference_uri, + to_request_payload, + unwrap_envelope, +) + + +class TestUrls: + def test_submit_url_defaults_to_the_public_api(self, monkeypatch): + monkeypatch.delenv("WAVESPEED_API_BASE", raising=False) + assert build_submit_url(None, "wavespeed-ai/z-image/turbo") == ( + f"{DEFAULT_API_BASE}/api/v3/wavespeed-ai/z-image/turbo" + ) + + def test_api_base_env_override(self, monkeypatch): + monkeypatch.setenv("WAVESPEED_API_BASE", "https://proxy.internal/") + assert get_api_base(None) == "https://proxy.internal" + assert build_result_url(None, "pred-1") == "https://proxy.internal/api/v3/predictions/pred-1/result" + + def test_explicit_api_base_beats_the_env(self, monkeypatch): + monkeypatch.setenv("WAVESPEED_API_BASE", "https://proxy.internal") + assert get_api_base("https://other.internal") == "https://other.internal" + + def test_empty_model_is_rejected(self): + with pytest.raises(WaveSpeedError, match="model is required"): + build_submit_url(None, "///") + + def test_path_traversal_in_the_model_id_is_rejected(self): + with pytest.raises(ValueError, match="model cannot be a dot path segment"): + build_submit_url(None, "wavespeed-ai/../../admin") + + def test_prediction_id_is_percent_encoded(self): + assert build_result_url("https://api.wavespeed.ai", "a b").endswith("/predictions/a%20b/result") + + +class TestHeaders: + def test_headers_carry_auth_and_channel_attribution(self): + headers = build_headers("sk-test") + assert headers["Authorization"] == "Bearer sk-test" + assert headers["X-Client-Name"] == "litellm" + assert headers["X-Client-Version"] + + def test_api_key_falls_back_to_the_env(self, monkeypatch): + monkeypatch.setenv("WAVESPEED_API_KEY", "sk-env") + assert build_headers(None)["Authorization"] == "Bearer sk-env" + + def test_missing_api_key_raises_401(self, monkeypatch): + monkeypatch.delenv("WAVESPEED_API_KEY", raising=False) + with pytest.raises(WaveSpeedError) as exc_info: + build_headers(None) + assert exc_info.value.status_code == 401 + + +class TestUnwrapEnvelope: + def test_happy_path(self): + raw = httpx.Response(200, json={"code": 200, "message": "ok", "data": {"id": "pred-1"}}) + assert unwrap_envelope(raw)["id"] == "pred-1" + + def test_http_error_surfaces_the_status_code(self): + raw = httpx.Response(503, text="upstream down") + with pytest.raises(WaveSpeedError) as exc_info: + unwrap_envelope(raw) + assert exc_info.value.status_code == 503 + assert "upstream down" in str(exc_info.value) + + def test_non_json_body(self): + raw = httpx.Response(200, text="gateway") + with pytest.raises(WaveSpeedError, match="Could not parse"): + unwrap_envelope(raw) + + def test_non_object_body(self): + raw = httpx.Response(200, json=["not", "an", "envelope"]) + with pytest.raises(WaveSpeedError, match="Unexpected WaveSpeed response body"): + unwrap_envelope(raw) + + def test_platform_error_code_uses_the_platform_message(self): + raw = httpx.Response(200, json={"code": 401, "message": "invalid api key", "data": None}) + with pytest.raises(WaveSpeedError, match="invalid api key"): + unwrap_envelope(raw) + + def test_platform_error_code_without_a_message(self): + raw = httpx.Response(200, json={"code": 500, "data": None}) + with pytest.raises(WaveSpeedError, match="WaveSpeed returned code 500"): + unwrap_envelope(raw) + + def test_missing_data(self): + raw = httpx.Response(200, json={"code": 200, "message": "ok"}) + with pytest.raises(WaveSpeedError, match="missing `data`"): + unwrap_envelope(raw) + + +class TestPredictionHelpers: + def test_missing_prediction_id_raises(self): + with pytest.raises(WaveSpeedError, match="missing a prediction id"): + get_prediction_id({"status": "created"}) + + def test_get_outputs_defaults_to_empty(self): + assert get_outputs({"status": "completed"}) == () + assert get_outputs({"status": "completed", "outputs": None}) == () + assert get_outputs({"status": "completed", "outputs": ["a"]}) == ["a"] + + @pytest.mark.parametrize( + "status, expected", [("completed", "done"), ("created", "pending"), ("processing", "pending")] + ) + def test_poll_outcome_non_terminal_and_success(self, status, expected): + assert poll_outcome({"status": status}) == expected + + @pytest.mark.parametrize("status", ["failed", "cancelled", "timeout"]) + def test_poll_outcome_terminal_failures(self, status): + with pytest.raises(WaveSpeedError, match=status): + poll_outcome({"status": status, "error": "boom"}) + + def test_poll_outcome_failure_without_an_error_detail(self): + with pytest.raises(WaveSpeedError, match="no error detail returned"): + poll_outcome({"status": "failed"}) + + def test_status_mapping(self): + assert map_status_to_openai("processing") == "in_progress" + assert map_status_to_openai("cancelled") == "failed" + assert map_status_to_openai("brand-new-status") == "queued" + + +class TestPayloadHelpers: + def test_to_request_payload_accepts_mappings_and_pairs(self): + assert to_request_payload({"a": 1}) == {"a": 1} + assert to_request_payload((("a", 1), ("b", 2))) == {"a": 1, "b": 2} + + def test_optional_helpers_drop_none(self): + assert optional_pair("a", 1) == (("a", 1),) + assert optional_pair("a", None) == () + assert dict(optional_entry("a", 1)) == {"a": 1} + assert dict(optional_entry("a", None)) == {} + + +class TestApiBaseIsolation: + """Chat and media share the provider slug but not the host.""" + + def test_the_chat_base_never_builds_a_prediction_url(self, monkeypatch): + monkeypatch.delenv("WAVESPEED_API_BASE", raising=False) + assert get_api_base(CHAT_API_BASE) == DEFAULT_API_BASE + assert get_api_base(CHAT_API_BASE + "/") == DEFAULT_API_BASE + + def test_the_chat_base_in_the_env_does_not_break_media(self, monkeypatch): + monkeypatch.setenv("WAVESPEED_API_BASE", CHAT_API_BASE) + assert build_submit_url(None, "wavespeed-ai/z-image/turbo") == ( + f"{DEFAULT_API_BASE}/api/v3/wavespeed-ai/z-image/turbo" + ) + + def test_a_self_hosted_base_is_still_honored(self, monkeypatch): + monkeypatch.delenv("WAVESPEED_API_BASE", raising=False) + assert get_api_base("https://wavespeed.internal.corp") == "https://wavespeed.internal.corp" + + +class TestReferenceNormalization: + """WaveSpeed submits JSON, so a reference has to be a URL or a data URI.""" + + PNG = b"\x89PNG\r\n\x1a\n" + b"rest-of-the-png" + + def test_urls_and_data_uris_pass_through(self): + assert to_reference_uri("https://example.com/a.png") == "https://example.com/a.png" + assert to_reference_uri("data:image/png;base64,AAAA") == "data:image/png;base64,AAAA" + + def test_bytes_are_inlined_with_a_sniffed_media_type(self): + assert to_reference_uri(self.PNG).startswith("data:image/png;base64,") + assert to_reference_uri(b"\xff\xd8\xffrest").startswith("data:image/jpeg;base64,") + assert to_reference_uri(b"GIF89arest").startswith("data:image/gif;base64,") + assert to_reference_uri(b"RIFF1234WEBPrest").startswith("data:image/webp;base64,") + + def test_bytes_round_trip(self): + encoded = to_reference_uri(self.PNG).split(",", 1)[1] + assert base64.b64decode(encoded) == self.PNG + + def test_a_path_uses_its_extension_for_the_media_type(self, tmp_path): + path = tmp_path / "frame.png" + path.write_bytes(self.PNG) + assert to_reference_uri(path).startswith("data:image/png;base64,") + + def test_a_binary_file_handle_is_read(self, tmp_path): + path = tmp_path / "frame.png" + path.write_bytes(self.PNG) + with open(path, "rb") as handle: + assert to_reference_uri(handle).startswith("data:image/png;base64,") + + def test_a_named_tuple_reference_uses_the_filename(self): + assert to_reference_uri(("frame.jpg", self.PNG)).startswith("data:image/jpeg;base64,") + + def test_unsniffable_bytes_are_rejected_with_an_actionable_message(self): + with pytest.raises(WaveSpeedError, match="Pass a URL, a data URI, or a named file"): + to_reference_uri(b"not-a-known-format") + + def test_a_text_mode_handle_is_rejected(self, tmp_path): + path = tmp_path / "frame.txt" + path.write_text("hello") + with open(path) as handle: + with pytest.raises(WaveSpeedError, match="binary mode"): + to_reference_uri(handle) + + def test_a_short_tuple_is_rejected(self): + with pytest.raises(WaveSpeedError, match="missing its content"): + to_reference_uri(("frame.png",)) + + def test_a_tuple_wrapping_a_url_keeps_the_url(self): + assert to_reference_uri(("frame.png", "https://example.com/a.png")) == "https://example.com/a.png" + + def test_an_unsupported_type_is_rejected(self): + with pytest.raises(WaveSpeedError, match="Unsupported reference type"): + to_reference_uri(object()) 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..8c81d7deee1 --- /dev/null +++ b/tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py @@ -0,0 +1,128 @@ +"""Tests for WaveSpeed AI provider registration across chat, image, and video surfaces.""" + +import httpx +import respx + +import litellm +from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider +from litellm.llms.openai_like.json_loader import JSONProviderRegistry +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) + + +def test_api_base_autodetects_the_provider(monkeypatch): + """Pointing api_base at the WaveSpeed chat host is enough to route there.""" + monkeypatch.setenv("WAVESPEED_API_KEY", "sk-env-key") + + _, provider, api_key, _ = get_llm_provider( + model="glm-5", + custom_llm_provider=None, + api_base="https://llm.wavespeed.ai/v1", + api_key=None, + ) + + assert provider == "wavespeed" + assert api_key == "sk-env-key" + + +def test_api_key_and_base_resolved_from_env(monkeypatch): + monkeypatch.setenv("WAVESPEED_API_KEY", "sk-env-key") + monkeypatch.setenv("WAVESPEED_API_BASE", "https://proxy.internal/v1") + + _, provider, api_key, api_base = get_llm_provider( + model="wavespeed/deepseek/deepseek-v4-flash", + custom_llm_provider=None, + api_base=None, + api_key=None, + ) + + assert provider == "wavespeed" + assert api_key == "sk-env-key" + assert api_base == "https://proxy.internal/v1" + + +@respx.mock +def test_image_generation_routes_through_the_public_sdk(monkeypatch): + """litellm.image_generation dispatches wavespeed models to the polling handler.""" + monkeypatch.setenv("WAVESPEED_API_KEY", "sk-test") + monkeypatch.delenv("WAVESPEED_API_BASE", raising=False) + monkeypatch.setattr("litellm.llms.wavespeed.image_generation.handler.DEFAULT_POLLING_INTERVAL", 0) + + model = "wavespeed-ai/z-image/turbo" + output_url = "https://cdn.wavespeed.ai/pred-123.png" + envelope = {"code": 200, "message": "ok", "data": {"id": "pred-123", "status": "created"}} + completed = { + "code": 200, + "message": "ok", + "data": {"id": "pred-123", "status": "completed", "outputs": [output_url]}, + } + + submit = respx.post(f"https://api.wavespeed.ai/api/v3/{model}").mock( + return_value=httpx.Response(200, json=envelope) + ) + respx.get("https://api.wavespeed.ai/api/v3/predictions/pred-123/result").mock( + return_value=httpx.Response(200, json=completed) + ) + + response = litellm.image_generation(model=f"wavespeed/{model}", prompt="a red panda") + + assert response.data[0].url == output_url + assert submit.call_count == 1 diff --git a/tests/test_litellm/llms/wavespeed/videos/__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..3526d39b82e --- /dev/null +++ b/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py @@ -0,0 +1,361 @@ +"""Tests for WaveSpeed AI video generation transformation.""" + +import json + +from unittest.mock import Mock + +import httpx +import pytest +import respx + +import litellm + +from litellm.litellm_core_utils.url_utils import SSRFError +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" + + +class TestWaveSpeedVideoContentDownload: + def setup_method(self): + self.config = WaveSpeedVideoConfig() + self.logging_obj = Mock() + + @respx.mock + def test_content_response_downloads_the_output(self): + download = respx.get(OUTPUT_URL).mock(return_value=httpx.Response(200, content=b"mp4-bytes")) + + content = self.config.transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])), + logging_obj=self.logging_obj, + ) + + assert content == b"mp4-bytes" + assert download.call_count == 1 + + @respx.mock + def test_content_response_raises_on_a_dead_output_url(self): + respx.get(OUTPUT_URL).mock(return_value=httpx.Response(404)) + + with pytest.raises(httpx.HTTPStatusError): + self.config.transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])), + logging_obj=self.logging_obj, + ) + + @pytest.mark.asyncio + @respx.mock + async def test_async_content_response_downloads_the_output(self, monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + download = respx.get(OUTPUT_URL).mock(return_value=httpx.Response(200, content=b"mp4-bytes")) + + content = await self.config.async_transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[OUTPUT_URL])), + logging_obj=self.logging_obj, + ) + + assert content == b"mp4-bytes" + assert download.call_count == 1 + + @pytest.mark.asyncio + @respx.mock + async def test_async_content_response_raises_while_still_processing(self, monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + with pytest.raises(WaveSpeedError, match="still created"): + await self.config.async_transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("created")), + logging_obj=self.logging_obj, + ) + + +class TestWaveSpeedVideoMisc: + def setup_method(self): + self.config = WaveSpeedVideoConfig() + + def test_get_complete_url_defaults_and_overrides(self): + assert self.config.get_complete_url(MODEL, None, {}) == API_BASE + assert self.config.get_complete_url(MODEL, "https://proxy.internal/", {}) == "https://proxy.internal" + + def test_status_retrieve_response_encodes_the_provider_into_the_id(self): + video = self.config.transform_video_status_retrieve_response( + raw_response=httpx.Response(200, json=prediction("processing")), + logging_obj=Mock(), + custom_llm_provider="wavespeed", + ) + + assert video.status == "in_progress" + assert extract_original_video_id(video.id) == "pred-123" + + def test_unknown_status_falls_back_to_queued(self): + video = self.config.transform_video_status_retrieve_response( + raw_response=httpx.Response(200, json=prediction("something-new")), + logging_obj=Mock(), + ) + assert video.status == "queued" + + @pytest.mark.parametrize( + "created_at, expected", + [("2026-08-20T10:00:00Z", 1787220000), (None, 0), ("", 0), ("not-a-date", 0)], + ) + def test_created_at_parsing(self, created_at, expected): + payload = envelope({"id": "pred-123", "status": "created", "created_at": created_at}) + video = self.config.transform_video_status_retrieve_response( + raw_response=httpx.Response(200, json=payload), logging_obj=Mock() + ) + assert video.created_at == expected + + @pytest.mark.parametrize( + "seconds, expected_duration", + [("5", 5), (5, 5), (5.9, 5), ("5.9", 5), (None, None), ("abc", None), (True, None), (object(), None)], + ) + def test_seconds_coercion(self, seconds, expected_duration): + mapped = self.config.map_openai_params({"seconds": seconds}, MODEL, False) + assert mapped.get("duration") == expected_duration + + def test_size_without_an_x_is_left_alone(self): + assert "size" not in self.config.map_openai_params({"size": "720p"}, MODEL, False) + + def test_get_error_class(self): + error = self.config.get_error_class("boom", 503, {}) + assert isinstance(error, WaveSpeedError) + assert error.status_code == 503 + + @pytest.mark.parametrize( + "call", + [ + lambda c: c.transform_video_remix_request("v", "p", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_remix_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_list_request(API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_list_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_delete_request("v", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_delete_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_create_character_request("n", object(), API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_create_character_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_get_character_request("c", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_get_character_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_edit_request("p", "v", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_edit_response(httpx.Response(200), Mock()), + lambda c: c.transform_video_extension_request("p", "v", "5", API_BASE, GenericLiteLLMParams(), {}), + lambda c: c.transform_video_extension_response(httpx.Response(200), Mock()), + ], + ) + def test_unsupported_surfaces_raise_not_implemented(self, call): + with pytest.raises(NotImplementedError): + call(self.config) + + +class TestWaveSpeedVideoContentSSRF: + """The output URL comes from the upstream response, so it is untrusted input. + + A deployment pointed at a WaveSpeed-compatible endpoint could have that endpoint + hand back an internal address, and /videos/{id}/content would relay the response + back to the caller. Every fetch goes through the repo's safe_get helpers, which + validate the resolved IP and re-validate each redirect hop. + """ + + def setup_method(self): + self.config = WaveSpeedVideoConfig() + self.logging_obj = Mock() + + @pytest.mark.parametrize( + "internal_url", + [ + "http://169.254.169.254/latest/meta-data/iam/security-credentials/", + "https://169.254.169.254/latest/meta-data/", + "http://127.0.0.1:8080/admin", + "http://10.0.0.5/internal", + "http://192.168.1.1/router", + "http://172.16.0.1/internal", + "https://[::1]/admin", + "file:///etc/passwd", + ], + ) + @respx.mock + def test_internal_output_url_is_rejected(self, internal_url): + leak = respx.get(internal_url).mock(return_value=httpx.Response(200, content=b"secret")) + + with pytest.raises(SSRFError): + self.config.transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[internal_url])), + logging_obj=self.logging_obj, + ) + + assert leak.call_count == 0 + + @respx.mock + def test_redirect_to_the_metadata_service_is_rejected(self): + """A public first hop that 302s to link-local must not be followed.""" + public_url = "https://93.184.216.34/video.mp4" + metadata_url = "http://169.254.169.254/latest/meta-data/" + + first_hop = respx.get(public_url).mock(return_value=httpx.Response(302, headers={"location": metadata_url})) + leak = respx.get(metadata_url).mock(return_value=httpx.Response(200, content=b"secret")) + + with pytest.raises(SSRFError): + self.config.transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[public_url])), + logging_obj=self.logging_obj, + ) + + assert first_hop.call_count == 1 + assert leak.call_count == 0 + + @pytest.mark.asyncio + @respx.mock + async def test_async_internal_output_url_is_rejected(self, monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + metadata_url = "http://169.254.169.254/latest/meta-data/" + leak = respx.get(metadata_url).mock(return_value=httpx.Response(200, content=b"secret")) + + with pytest.raises(SSRFError): + await self.config.async_transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[metadata_url])), + logging_obj=self.logging_obj, + ) + + assert leak.call_count == 0 + + @pytest.mark.asyncio + @respx.mock + async def test_async_redirect_to_the_metadata_service_is_rejected(self, monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + public_url = "https://93.184.216.34/video.mp4" + metadata_url = "http://169.254.169.254/latest/meta-data/" + + respx.get(public_url).mock(return_value=httpx.Response(302, headers={"location": metadata_url})) + leak = respx.get(metadata_url).mock(return_value=httpx.Response(200, content=b"secret")) + + with pytest.raises(SSRFError): + await self.config.async_transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[public_url])), + logging_obj=self.logging_obj, + ) + + assert leak.call_count == 0 + + @respx.mock + def test_a_public_output_url_still_downloads(self): + public_url = "https://93.184.216.34/video.mp4" + download = respx.get(public_url).mock(return_value=httpx.Response(200, content=b"mp4-bytes")) + + content = self.config.transform_video_content_response( + raw_response=httpx.Response(200, json=prediction("completed", outputs=[public_url])), + logging_obj=self.logging_obj, + ) + + assert content == b"mp4-bytes" + assert download.call_count == 1 + + +class TestWaveSpeedVideoReferenceInputs: + def setup_method(self): + self.config = WaveSpeedVideoConfig() + + def test_a_binary_reference_is_inlined_so_the_json_body_stays_serializable(self, tmp_path): + path = tmp_path / "frame.png" + path.write_bytes(b"\x89PNG\r\n\x1a\nrest") + + with open(path, "rb") as handle: + mapped = self.config.map_openai_params({"input_reference": handle}, MODEL, False) + + assert mapped["image"].startswith("data:image/png;base64,") + json.dumps(mapped) + + def test_a_url_reference_is_left_alone(self): + mapped = self.config.map_openai_params({"input_reference": "https://example.com/frame.png"}, MODEL, False) + assert mapped["image"] == "https://example.com/frame.png"