mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge 26448beb67 into 77765fd302
This commit is contained in:
commit
d73c0ffc9a
21 changed files with 2248 additions and 0 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
}
|
||||
}
|
||||
|
|
|
|||
0
litellm/llms/wavespeed/__init__.py
Normal file
0
litellm/llms/wavespeed/__init__.py
Normal file
258
litellm/llms/wavespeed/common_utils.py
Normal file
258
litellm/llms/wavespeed/common_utils.py
Normal file
|
|
@ -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")
|
||||
0
litellm/llms/wavespeed/image_generation/__init__.py
Normal file
0
litellm/llms/wavespeed/image_generation/__init__.py
Normal file
269
litellm/llms/wavespeed/image_generation/handler.py
Normal file
269
litellm/llms/wavespeed/image_generation/handler.py
Normal file
|
|
@ -0,0 +1,269 @@
|
|||
"""
|
||||
WaveSpeed AI image generation handler.
|
||||
|
||||
WaveSpeed predictions are asynchronous: one submit POST returns a prediction id, then the
|
||||
result endpoint is polled until the prediction reaches a terminal status.
|
||||
|
||||
The submit POST is issued exactly once and is never retried, because every submission is a
|
||||
billable task and a retry would create a duplicate one. Poll GETs are read-only, so a short
|
||||
run of connection failures is tolerated before giving up.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from collections.abc import Coroutine, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final, NamedTuple
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
_get_httpx_client, # pyright: ignore[reportPrivateUsage] # the shared sync client factory litellm providers use
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
from ..common_utils import (
|
||||
DEFAULT_MAX_POLLING_TIME,
|
||||
DEFAULT_POLLING_INTERVAL,
|
||||
MAX_CONSECUTIVE_POLL_FAILURES,
|
||||
WaveSpeedError,
|
||||
build_result_url,
|
||||
get_prediction_id,
|
||||
poll_outcome,
|
||||
to_request_payload,
|
||||
unwrap_envelope,
|
||||
)
|
||||
from .transformation import WaveSpeedImageGenerationConfig
|
||||
|
||||
|
||||
class _PreparedRequest(NamedTuple):
|
||||
headers: Mapping[str, str]
|
||||
submit_url: str
|
||||
body: Mapping[str, object]
|
||||
|
||||
|
||||
class _ResolvedParams(NamedTuple):
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
litellm_params: Mapping[str, object]
|
||||
|
||||
|
||||
class WaveSpeedImageGeneration:
|
||||
def __init__(self, config: WaveSpeedImageGenerationConfig | None = None) -> None:
|
||||
self.config: Final = config or WaveSpeedImageGenerationConfig()
|
||||
|
||||
def image_generation(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
model_response: ImageResponse,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: GenericLiteLLMParams | Mapping[str, object],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
extra_headers: Mapping[str, str] | None = None,
|
||||
client: HTTPHandler | AsyncHTTPHandler | None = None,
|
||||
aimg_generation: bool = False,
|
||||
) -> "ImageResponse | Coroutine[object, object, ImageResponse]":
|
||||
if aimg_generation:
|
||||
return self.async_image_generation(
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
model_response=model_response,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
timeout=timeout,
|
||||
extra_headers=extra_headers,
|
||||
client=client if isinstance(client, AsyncHTTPHandler) else None,
|
||||
)
|
||||
|
||||
resolved: Final = _resolve_params(litellm_params)
|
||||
sync_client: Final = client if isinstance(client, HTTPHandler) else _get_httpx_client()
|
||||
prepared: Final = self._prepare(model, prompt, resolved, optional_params, extra_headers, logging_obj)
|
||||
|
||||
submit_response: Final = sync_client.post(
|
||||
url=prepared.submit_url,
|
||||
headers=to_request_payload(prepared.headers),
|
||||
json=to_request_payload(prepared.body),
|
||||
timeout=timeout,
|
||||
)
|
||||
prediction_id: Final = get_prediction_id(unwrap_envelope(submit_response))
|
||||
result_url: Final = build_result_url(resolved.api_base, prediction_id)
|
||||
deadline: Final = time.time() + DEFAULT_MAX_POLLING_TIME
|
||||
poll_headers: Final = to_request_payload(prepared.headers)
|
||||
|
||||
consecutive_failures = 0 # rebind-ok: counts consecutive poll transport failures
|
||||
while time.time() < deadline:
|
||||
try:
|
||||
poll_response = sync_client.get(url=result_url, headers=poll_headers, 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()
|
||||
176
litellm/llms/wavespeed/image_generation/transformation.py
Normal file
176
litellm/llms/wavespeed/image_generation/transformation.py
Normal file
|
|
@ -0,0 +1,176 @@
|
|||
"""
|
||||
WaveSpeed AI image generation configuration.
|
||||
|
||||
Transforms between the OpenAI image generation contract and WaveSpeed's prediction API.
|
||||
The submit/poll HTTP flow lives in ``handler.py``; this class only transforms data.
|
||||
|
||||
API Reference: https://wavespeed.ai/docs
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any, # noqa: TID251 # runtime stand-in for the TYPE_CHECKING-only logging type
|
||||
Final,
|
||||
TypeAlias,
|
||||
)
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
OpenAIImageGenerationOptionalParams,
|
||||
)
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
from ..common_utils import (
|
||||
WaveSpeedError,
|
||||
build_headers,
|
||||
build_submit_url,
|
||||
get_outputs,
|
||||
to_request_payload,
|
||||
unwrap_envelope,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj: TypeAlias = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj: TypeAlias = Any
|
||||
|
||||
|
||||
class WaveSpeedImageGenerationConfig(BaseImageGenerationConfig):
|
||||
"""
|
||||
Configuration for WaveSpeed AI image generation.
|
||||
|
||||
Any WaveSpeed image model id works as-is, e.g. ``wavespeed/bytedance/seedream-v5.0-pro``
|
||||
or ``wavespeed/wavespeed-ai/z-image/turbo``. Model-specific fields that have no OpenAI
|
||||
equivalent are passed straight through to the prediction body.
|
||||
"""
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> list[OpenAIImageGenerationOptionalParams]: # mutable-ok: base config contract returns `list`
|
||||
return ["n", "size", "response_format"] # mutable-ok: base config contract returns `list`
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
optional_params: Mapping[str, object],
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict: # mutable-ok: base config contract returns bare `dict`
|
||||
return to_request_payload(
|
||||
MappingProxyType(
|
||||
{**optional_params, **self._mapped_params(non_default_params, optional_params, drop_params)}
|
||||
)
|
||||
)
|
||||
|
||||
def _mapped_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
optional_params: Mapping[str, object],
|
||||
drop_params: bool,
|
||||
) -> Mapping[str, object]:
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in (
|
||||
self._map_one(key, value, drop_params)
|
||||
for key, value in non_default_params.items()
|
||||
if key not in optional_params and value is not None
|
||||
)
|
||||
if key is not None
|
||||
}
|
||||
)
|
||||
|
||||
def _map_one(self, key: str, value: object, drop_params: bool) -> tuple[str | None, object]:
|
||||
if key == "size":
|
||||
return "size", self._map_size(value)
|
||||
if key == "n":
|
||||
return ("num_images", value) if isinstance(value, int) and value > 1 else (None, None)
|
||||
if key == "response_format":
|
||||
if value != "url" and not drop_params:
|
||||
raise ValueError(
|
||||
"WaveSpeed returns hosted image URLs, so only response_format='url' is supported. "
|
||||
"Set drop_params=True to ignore this parameter."
|
||||
)
|
||||
return None, None
|
||||
return key, value
|
||||
|
||||
def _map_size(self, size: object) -> str:
|
||||
"""WaveSpeed takes ``{width}*{height}`` where OpenAI takes ``{width}x{height}``."""
|
||||
width, separator, height = str(size).lower().partition("x")
|
||||
if not separator or not width.isdigit() or not height.isdigit():
|
||||
raise ValueError(f"Invalid size format: '{size}'. Expected 'WIDTHxHEIGHT' (e.g. '1024x1024').")
|
||||
return f"{width}*{height}"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict: # mutable-ok: base config contract returns bare `dict`
|
||||
return to_request_payload(MappingProxyType({**build_headers(api_key), **headers}))
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
return build_submit_url(api_base, model)
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
headers: Mapping[str, str],
|
||||
) -> dict: # mutable-ok: base config contract returns bare `dict`
|
||||
return to_request_payload(MappingProxyType({"prompt": prompt, **optional_params}))
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: Mapping[str, object],
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
encoding: object,
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
"""Transform the final polled prediction into an OpenAI image response."""
|
||||
prediction: Final = unwrap_envelope(raw_response)
|
||||
outputs: Final = get_outputs(prediction)
|
||||
|
||||
if not outputs:
|
||||
raise WaveSpeedError(
|
||||
status_code=500,
|
||||
message=f"WaveSpeed prediction {prediction.get('id', '')} completed without any outputs",
|
||||
)
|
||||
|
||||
images: Final = [ImageObject(url=url, b64_json=None) for url in outputs] # mutable-ok: pydantic field is `list`
|
||||
model_response.data = images # rebind-ok: the base contract fills in and returns the caller's ImageResponse
|
||||
return model_response
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers
|
||||
) -> WaveSpeedError:
|
||||
return WaveSpeedError(status_code=status_code, message=error_message, headers=headers)
|
||||
0
litellm/llms/wavespeed/videos/__init__.py
Normal file
0
litellm/llms/wavespeed/videos/__init__.py
Normal file
389
litellm/llms/wavespeed/videos/transformation.py
Normal file
389
litellm/llms/wavespeed/videos/transformation.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
0
tests/test_litellm/llms/wavespeed/__init__.py
Normal file
0
tests/test_litellm/llms/wavespeed/__init__.py
Normal file
|
|
@ -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,
|
||||
}
|
||||
227
tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py
Normal file
227
tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py
Normal file
|
|
@ -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="<html>gateway</html>")
|
||||
with pytest.raises(WaveSpeedError, match="Could not parse"):
|
||||
unwrap_envelope(raw)
|
||||
|
||||
def test_non_object_body(self):
|
||||
raw = httpx.Response(200, json=["not", "an", "envelope"])
|
||||
with pytest.raises(WaveSpeedError, match="Unexpected WaveSpeed response body"):
|
||||
unwrap_envelope(raw)
|
||||
|
||||
def test_platform_error_code_uses_the_platform_message(self):
|
||||
raw = httpx.Response(200, json={"code": 401, "message": "invalid api key", "data": None})
|
||||
with pytest.raises(WaveSpeedError, match="invalid api key"):
|
||||
unwrap_envelope(raw)
|
||||
|
||||
def test_platform_error_code_without_a_message(self):
|
||||
raw = httpx.Response(200, json={"code": 500, "data": None})
|
||||
with pytest.raises(WaveSpeedError, match="WaveSpeed returned code 500"):
|
||||
unwrap_envelope(raw)
|
||||
|
||||
def test_missing_data(self):
|
||||
raw = httpx.Response(200, json={"code": 200, "message": "ok"})
|
||||
with pytest.raises(WaveSpeedError, match="missing `data`"):
|
||||
unwrap_envelope(raw)
|
||||
|
||||
|
||||
class TestPredictionHelpers:
|
||||
def test_missing_prediction_id_raises(self):
|
||||
with pytest.raises(WaveSpeedError, match="missing a prediction id"):
|
||||
get_prediction_id({"status": "created"})
|
||||
|
||||
def test_get_outputs_defaults_to_empty(self):
|
||||
assert get_outputs({"status": "completed"}) == ()
|
||||
assert get_outputs({"status": "completed", "outputs": None}) == ()
|
||||
assert get_outputs({"status": "completed", "outputs": ["a"]}) == ["a"]
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"status, expected", [("completed", "done"), ("created", "pending"), ("processing", "pending")]
|
||||
)
|
||||
def test_poll_outcome_non_terminal_and_success(self, status, expected):
|
||||
assert poll_outcome({"status": status}) == expected
|
||||
|
||||
@pytest.mark.parametrize("status", ["failed", "cancelled", "timeout"])
|
||||
def test_poll_outcome_terminal_failures(self, status):
|
||||
with pytest.raises(WaveSpeedError, match=status):
|
||||
poll_outcome({"status": status, "error": "boom"})
|
||||
|
||||
def test_poll_outcome_failure_without_an_error_detail(self):
|
||||
with pytest.raises(WaveSpeedError, match="no error detail returned"):
|
||||
poll_outcome({"status": "failed"})
|
||||
|
||||
def test_status_mapping(self):
|
||||
assert map_status_to_openai("processing") == "in_progress"
|
||||
assert map_status_to_openai("cancelled") == "failed"
|
||||
assert map_status_to_openai("brand-new-status") == "queued"
|
||||
|
||||
|
||||
class TestPayloadHelpers:
|
||||
def test_to_request_payload_accepts_mappings_and_pairs(self):
|
||||
assert to_request_payload({"a": 1}) == {"a": 1}
|
||||
assert to_request_payload((("a", 1), ("b", 2))) == {"a": 1, "b": 2}
|
||||
|
||||
def test_optional_helpers_drop_none(self):
|
||||
assert optional_pair("a", 1) == (("a", 1),)
|
||||
assert optional_pair("a", None) == ()
|
||||
assert dict(optional_entry("a", 1)) == {"a": 1}
|
||||
assert dict(optional_entry("a", None)) == {}
|
||||
|
||||
|
||||
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())
|
||||
128
tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py
Normal file
128
tests/test_litellm/llms/wavespeed/test_wavespeed_provider.py
Normal file
|
|
@ -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
|
||||
0
tests/test_litellm/llms/wavespeed/videos/__init__.py
Normal file
0
tests/test_litellm/llms/wavespeed/videos/__init__.py
Normal file
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue