diff --git a/litellm/images/main.py b/litellm/images/main.py index 248e09d88f0..14b85638610 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -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, diff --git a/litellm/llms/wavespeed/common_utils.py b/litellm/llms/wavespeed/common_utils.py index e39f78b89ac..8206c1a245e 100644 --- a/litellm/llms/wavespeed/common_utils.py +++ b/litellm/llms/wavespeed/common_utils.py @@ -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]: diff --git a/litellm/llms/wavespeed/image_generation/handler.py b/litellm/llms/wavespeed/image_generation/handler.py index de79ea2ae6f..1d5cf60171c 100644 --- a/litellm/llms/wavespeed/image_generation/handler.py +++ b/litellm/llms/wavespeed/image_generation/handler.py @@ -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) diff --git a/litellm/llms/wavespeed/videos/transformation.py b/litellm/llms/wavespeed/videos/transformation.py index 4362cda1991..a5878fc7fcd 100644 --- a/litellm/llms/wavespeed/videos/transformation.py +++ b/litellm/llms/wavespeed/videos/transformation.py @@ -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 diff --git a/litellm/provider_endpoints_support_backup.json b/litellm/provider_endpoints_support_backup.json index 96f3f3ea9fd..a5823524eb0 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -2258,7 +2258,8 @@ "moderations": false, "batches": false, "rerank": false, - "a2a": false + "a2a": false, + "video_generations": true } }, "watsonx_text": { diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 9f38b150a4e..8f1db5e5bcc 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2569,7 +2569,8 @@ "moderations": false, "batches": false, "rerank": false, - "a2a": false + "a2a": false, + "video_generations": true } }, "watsonx_text": { diff --git a/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py b/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py index b17451d109b..08f618fdb16 100644 --- a/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py +++ b/tests/test_litellm/llms/wavespeed/image_generation/test_wavespeed_image_generation.py @@ -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, + } diff --git a/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py b/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py index f62efa0f0b8..83c22fc64e3 100644 --- a/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py +++ b/tests/test_litellm/llms/wavespeed/test_wavespeed_common_utils.py @@ -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()) diff --git a/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py b/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py index 105eeeac384..3526d39b82e 100644 --- a/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py +++ b/tests/test_litellm/llms/wavespeed/videos/test_wavespeed_video_transformation.py @@ -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"