fix(wavespeed): validate the video output URL before fetching it

The output URL comes from the upstream prediction response, so a deployment
pointed at a WaveSpeed-compatible endpoint could have that endpoint return an
internal address and read the response back through /videos/{id}/content.
Both download paths now go through the repo's safe_get and async_safe_get,
which resolve and validate the IP actually connected to and re-validate every
redirect hop instead of following them blindly.

Also addresses the other review findings:

- poll GETs now carry the caller's timeout rather than falling back to the
  client's 600 second default
- api_key and api_base reach the handler, so a Router supplying either per
  model no longer needs matching environment variables
- get_api_base ignores the chat base, which provider resolution and
  WAVESPEED_API_BASE can both hand to the media path and which would build an
  unreachable prediction URL; any other override is still honored
- input_reference accepts bytes, paths, file handles and (filename, content)
  tuples by inlining them as data URIs, since the prediction body is JSON
- provider_endpoints_support.json advertises video_generations
This commit is contained in:
chengzeyi 2026-08-20 12:04:16 +00:00
parent 93eed536d6
commit 0286337ce2
9 changed files with 315 additions and 16 deletions

View file

@ -407,6 +407,8 @@ def image_generation(
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,

View file

@ -11,6 +11,9 @@ Both responses are wrapped in the platform envelope ``{"code": ..., "message": .
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
@ -29,6 +32,7 @@ class WaveSpeedError(BaseLLMException):
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
@ -76,6 +80,64 @@ def optional_entry(key: str, value: object) -> Mapping[str, object]:
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 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 isinstance(reference, str):
return reference
if isinstance(reference, (bytes, bytearray)):
return _to_data_uri(bytes(reference), None)
if isinstance(reference, os.PathLike):
path: Final = os.fspath(reference)
with open(path, "rb") as handle:
return _to_data_uri(handle.read(), str(path))
if isinstance(reference, tuple):
filename, content = (reference[0], reference[1]) if len(reference) >= 2 else (None, None)
if content is None:
raise WaveSpeedError(status_code=400, message="Reference tuple is missing its content")
inner: Final = to_reference_uri(content)
if inner.startswith("data:") and filename:
return _to_data_uri(base64.b64decode(inner.split(",", 1)[1]), str(filename))
return inner
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, getattr(reference, "name", None))
raise WaveSpeedError(status_code=400, message=f"Unsupported reference type: {type(reference).__name__}")
def get_api_key(api_key: str | None) -> str:
resolved: Final = api_key or get_secret_str("WAVESPEED_API_KEY")
if not resolved:
@ -87,7 +149,15 @@ def get_api_key(api_key: str | None) -> str:
def get_api_base(api_base: str | None) -> str:
return (api_base or get_secret_str("WAVESPEED_API_BASE") or DEFAULT_API_BASE).rstrip("/")
"""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]:

View file

@ -103,7 +103,7 @@ class WaveSpeedImageGeneration:
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)
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)
@ -149,7 +149,7 @@ class WaveSpeedImageGeneration:
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)
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)

View file

@ -29,13 +29,8 @@ 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.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
_get_httpx_client, # pyright: ignore[reportPrivateUsage] # the shared sync client factory litellm providers use
get_async_httpx_client,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.videos.main import VideoCreateOptionalRequestParams, VideoObject
from litellm.types.videos.utils import (
@ -55,6 +50,7 @@ from ..common_utils import (
map_status_to_openai,
optional_entry,
optional_pair,
to_reference_uri,
to_request_payload,
unwrap_envelope,
)
@ -143,7 +139,7 @@ class WaveSpeedVideoConfig(BaseVideoConfig):
return to_request_payload(
(
*((k, v) for k, v in video_create_optional_params.items() if k not in supported),
*optional_pair("image", input_reference),
*optional_pair("image", to_reference_uri(input_reference) if input_reference else None),
*optional_pair("size", mapped_size),
*optional_pair("duration", seconds),
)
@ -248,8 +244,7 @@ class WaveSpeedVideoConfig(BaseVideoConfig):
logging_obj: LiteLLMLoggingObj,
) -> bytes:
output_url: Final = self._extract_output_url(raw_response)
httpx_client: Final[HTTPHandler] = _get_httpx_client()
video_response: Final = httpx_client.get(output_url)
video_response: Final = safe_get(litellm.module_level_client, output_url)
video_response.raise_for_status()
return video_response.content
@ -259,8 +254,7 @@ class WaveSpeedVideoConfig(BaseVideoConfig):
logging_obj: LiteLLMLoggingObj,
) -> bytes:
output_url: Final = self._extract_output_url(raw_response)
async_client: Final[AsyncHTTPHandler] = get_async_httpx_client(llm_provider=litellm.LlmProviders.WAVESPEED)
video_response: Final = await async_client.get(output_url)
video_response: Final = await async_safe_get(litellm.module_level_aclient, output_url)
video_response.raise_for_status()
return video_response.content

View file

@ -2258,7 +2258,8 @@
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false
"a2a": false,
"video_generations": true
}
},
"watsonx_text": {

View file

@ -2569,7 +2569,8 @@
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false
"a2a": false,
"video_generations": true
}
},
"watsonx_text": {

View file

@ -339,3 +339,30 @@ def test_supported_openai_params_and_error_class():
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,
}

View file

@ -1,9 +1,12 @@
"""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,
@ -16,6 +19,7 @@ from litellm.llms.wavespeed.common_utils import (
optional_entry,
optional_pair,
poll_outcome,
to_reference_uri,
to_request_payload,
unwrap_envelope,
)
@ -146,3 +150,78 @@ class TestPayloadHelpers:
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())

View file

@ -1,5 +1,7 @@
"""Tests for WaveSpeed AI video generation transformation."""
import json
from unittest.mock import Mock
import httpx
@ -8,6 +10,7 @@ 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
@ -234,3 +237,125 @@ class TestWaveSpeedVideoMisc:
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"