mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge 16b964b473 into e6c4580a31
This commit is contained in:
commit
d7c3dce700
4 changed files with 168 additions and 9 deletions
|
|
@ -17,6 +17,7 @@ from litellm.constants import (
|
|||
RUNWAYML_DEFAULT_API_VERSION,
|
||||
RUNWAYML_POLLING_TIMEOUT,
|
||||
)
|
||||
from litellm.litellm_core_utils.url_utils import async_safe_get, safe_get
|
||||
from litellm.llms.base_llm.text_to_speech.transformation import (
|
||||
BaseTextToSpeechConfig,
|
||||
TextToSpeechRequestData,
|
||||
|
|
@ -496,11 +497,11 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
|
|||
if not isinstance(audio_url, str):
|
||||
raise ValueError(f"RunwayML TTS audio URL is not a string: {audio_url}")
|
||||
|
||||
# Download the audio file
|
||||
# Download the audio file with SSRF guards (provider output URLs are untrusted).
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
|
||||
client: Final = _get_httpx_client()
|
||||
audio_response: Final = client.get(url=audio_url)
|
||||
audio_response: Final = safe_get(client, audio_url)
|
||||
audio_response.raise_for_status()
|
||||
|
||||
verbose_logger.debug("RunwayML TTS audio downloaded successfully")
|
||||
|
|
@ -565,11 +566,11 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig):
|
|||
if not isinstance(audio_url, str):
|
||||
raise ValueError(f"RunwayML TTS audio URL is not a string: {audio_url}")
|
||||
|
||||
# Download the audio file (async)
|
||||
# Download the audio file (async) with SSRF guards.
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
|
||||
client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.RUNWAYML)
|
||||
audio_response: Final = await client.get(url=audio_url)
|
||||
audio_response: Final = await async_safe_get(client, audio_url)
|
||||
audio_response.raise_for_status()
|
||||
|
||||
verbose_logger.debug("RunwayML TTS audio downloaded successfully (async)")
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from typing_extensions import ReadOnly, TypedDict
|
|||
|
||||
import litellm
|
||||
from litellm.constants import RUNWAYML_DEFAULT_API_VERSION
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.litellm_core_utils.url_utils import async_safe_get, encode_url_path_segment, safe_get
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -456,9 +456,9 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
response_data: Final = raw_response.json()
|
||||
video_url: Final = self._extract_video_url_from_response(response_data)
|
||||
|
||||
# Download the video from the CloudFront URL synchronously
|
||||
# Download the video from the provider URL with SSRF guards (same as BFL image_edit).
|
||||
httpx_client: Final[HTTPHandler] = _get_httpx_client()
|
||||
video_response: Final = httpx_client.get(video_url)
|
||||
video_response: Final = safe_get(httpx_client, video_url)
|
||||
video_response.raise_for_status()
|
||||
|
||||
return video_response.content
|
||||
|
|
@ -485,11 +485,11 @@ class RunwayMLVideoConfig(BaseVideoConfig):
|
|||
response_data: Final = raw_response.json()
|
||||
video_url: Final = self._extract_video_url_from_response(response_data)
|
||||
|
||||
# Download the video from the CloudFront URL asynchronously
|
||||
# Download the video from the provider URL with SSRF guards (same as BFL image_edit).
|
||||
async_httpx_client: Final[AsyncHTTPHandler] = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders.RUNWAYML,
|
||||
)
|
||||
video_response: Final = await async_httpx_client.get(video_url)
|
||||
video_response: Final = await async_safe_get(async_httpx_client, video_url)
|
||||
video_response.raise_for_status()
|
||||
|
||||
return video_response.content
|
||||
|
|
|
|||
|
|
@ -62,3 +62,95 @@ def test_runwayml_native_voice_passthrough():
|
|||
assert "runwayml_voice" in mapped_params
|
||||
assert mapped_params["runwayml_voice"]["type"] == "runway-preset"
|
||||
assert mapped_params["runwayml_voice"]["presetId"] == runway_voice
|
||||
|
||||
def test_transform_text_to_speech_response():
|
||||
"""Test TTS audio download with SSRF-protected fetch."""
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
config = RunwayMLTextToSpeechConfig()
|
||||
|
||||
# Mock the initial response (task created)
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {"id": "task-123", "status": "PENDING"}
|
||||
mock_response.request.headers = {"Authorization": "Bearer test"}
|
||||
|
||||
# Mock the polled response (task completed)
|
||||
mock_polled = Mock(spec=httpx.Response)
|
||||
mock_polled.json.return_value = {
|
||||
"id": "task-123",
|
||||
"status": "SUCCEEDED",
|
||||
"output": ["https://example.com/audio.mp3"],
|
||||
}
|
||||
|
||||
# Mock the audio download response
|
||||
mock_audio_response = Mock(spec=httpx.Response)
|
||||
mock_audio_response.raise_for_status = Mock()
|
||||
|
||||
mock_client = Mock()
|
||||
mock_client.get.return_value = mock_audio_response
|
||||
|
||||
with patch.object(litellm, "user_url_validation", False):
|
||||
with patch.object(config, "_poll_task_sync", return_value=mock_polled):
|
||||
with patch("litellm.llms.custom_httpx.http_handler._get_httpx_client", return_value=mock_client):
|
||||
result = config.transform_text_to_speech_response(
|
||||
model="eleven_multilingual_v2",
|
||||
raw_response=mock_response,
|
||||
logging_obj=Mock(),
|
||||
)
|
||||
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
mock_client.get.assert_called_once()
|
||||
|
||||
|
||||
def test_async_transform_text_to_speech_response():
|
||||
"""Test async TTS audio download with SSRF-protected fetch."""
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
config = RunwayMLTextToSpeechConfig()
|
||||
|
||||
# Mock the initial response (task created)
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {"id": "task-123", "status": "PENDING"}
|
||||
mock_response.request.headers = {"Authorization": "Bearer test"}
|
||||
|
||||
# Mock the polled response (task completed)
|
||||
mock_polled = Mock(spec=httpx.Response)
|
||||
mock_polled.json.return_value = {
|
||||
"id": "task-123",
|
||||
"status": "SUCCEEDED",
|
||||
"output": ["https://example.com/audio.mp3"],
|
||||
}
|
||||
|
||||
# Mock the audio download response
|
||||
mock_audio_response = Mock(spec=httpx.Response)
|
||||
mock_audio_response.raise_for_status = Mock()
|
||||
|
||||
mock_client = Mock()
|
||||
mock_client.get = AsyncMock(return_value=mock_audio_response)
|
||||
|
||||
async def run_test():
|
||||
with patch.object(litellm, "user_url_validation", False):
|
||||
with patch.object(config, "_poll_task_async", new_callable=AsyncMock, return_value=mock_polled):
|
||||
with patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=mock_client):
|
||||
result = await config.async_transform_text_to_speech_response(
|
||||
model="eleven_multilingual_v2",
|
||||
raw_response=mock_response,
|
||||
logging_obj=Mock(),
|
||||
)
|
||||
return result
|
||||
|
||||
result = asyncio.run(run_test())
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
mock_client.get.assert_called_once()
|
||||
|
||||
|
|
|
|||
|
|
@ -370,5 +370,71 @@ class TestRunwayMLVideoTransformation:
|
|||
assert isinstance(status_obj.completed_at, int)
|
||||
|
||||
|
||||
|
||||
def test_transform_video_content_response(self):
|
||||
"""Test video content download with SSRF-protected fetch."""
|
||||
from unittest.mock import patch
|
||||
|
||||
import litellm
|
||||
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {
|
||||
"id": "test-id",
|
||||
"status": "SUCCEEDED",
|
||||
"output": ["https://example.com/video.mp4"],
|
||||
}
|
||||
|
||||
mock_video_response = Mock(spec=httpx.Response)
|
||||
mock_video_response.content = b"fake-video-bytes"
|
||||
mock_video_response.raise_for_status = Mock()
|
||||
|
||||
mock_client = Mock()
|
||||
mock_client.get.return_value = mock_video_response
|
||||
|
||||
with patch.object(litellm, "user_url_validation", False):
|
||||
with patch("litellm.llms.runwayml.videos.transformation._get_httpx_client", return_value=mock_client):
|
||||
result = self.config.transform_video_content_response(
|
||||
raw_response=mock_response,
|
||||
logging_obj=self.mock_logging_obj,
|
||||
)
|
||||
|
||||
assert result == b"fake-video-bytes"
|
||||
mock_client.get.assert_called_once()
|
||||
|
||||
def test_async_transform_video_content_response(self):
|
||||
"""Test async video content download with SSRF-protected fetch."""
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import litellm
|
||||
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {
|
||||
"id": "test-id",
|
||||
"status": "SUCCEEDED",
|
||||
"output": ["https://example.com/video.mp4"],
|
||||
}
|
||||
|
||||
mock_video_response = Mock(spec=httpx.Response)
|
||||
mock_video_response.content = b"fake-video-bytes"
|
||||
mock_video_response.raise_for_status = Mock()
|
||||
|
||||
mock_client = Mock()
|
||||
mock_client.get = AsyncMock(return_value=mock_video_response)
|
||||
|
||||
async def run_test():
|
||||
with patch.object(litellm, "user_url_validation", False):
|
||||
with patch("litellm.llms.runwayml.videos.transformation.get_async_httpx_client", return_value=mock_client):
|
||||
result = await self.config.async_transform_video_content_response(
|
||||
raw_response=mock_response,
|
||||
logging_obj=self.mock_logging_obj,
|
||||
)
|
||||
return result
|
||||
|
||||
result = asyncio.run(run_test())
|
||||
assert result == b"fake-video-bytes"
|
||||
mock_client.get.assert_called_once()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue