mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
93eed536d6
commit
0286337ce2
9 changed files with 315 additions and 16 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -2258,7 +2258,8 @@
|
|||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": false,
|
||||
"a2a": false
|
||||
"a2a": false,
|
||||
"video_generations": true
|
||||
}
|
||||
},
|
||||
"watsonx_text": {
|
||||
|
|
|
|||
|
|
@ -2569,7 +2569,8 @@
|
|||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": false,
|
||||
"a2a": false
|
||||
"a2a": false,
|
||||
"video_generations": true
|
||||
}
|
||||
},
|
||||
"watsonx_text": {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue