mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
test(runwayml): cover safe_get execution path in download tests
This commit is contained in:
parent
35d45a9992
commit
16b964b473
2 changed files with 32 additions and 16 deletions
|
|
@ -72,6 +72,7 @@ def test_transform_text_to_speech_response():
|
|||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
config = RunwayMLTextToSpeechConfig()
|
||||
|
|
@ -93,10 +94,12 @@ def test_transform_text_to_speech_response():
|
|||
mock_audio_response = Mock(spec=httpx.Response)
|
||||
mock_audio_response.raise_for_status = Mock()
|
||||
|
||||
with patch.object(config, "_poll_task_sync", return_value=mock_polled):
|
||||
with patch("litellm.llms.custom_httpx.http_handler._get_httpx_client"):
|
||||
with patch("litellm.llms.runwayml.text_to_speech.transformation.safe_get") as mock_safe_get:
|
||||
mock_safe_get.return_value = mock_audio_response
|
||||
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,
|
||||
|
|
@ -104,7 +107,7 @@ def test_transform_text_to_speech_response():
|
|||
)
|
||||
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
mock_safe_get.assert_called_once()
|
||||
mock_client.get.assert_called_once()
|
||||
|
||||
|
||||
def test_async_transform_text_to_speech_response():
|
||||
|
|
@ -114,6 +117,7 @@ def test_async_transform_text_to_speech_response():
|
|||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
config = RunwayMLTextToSpeechConfig()
|
||||
|
|
@ -135,11 +139,13 @@ def test_async_transform_text_to_speech_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(config, "_poll_task_async", new_callable=AsyncMock, return_value=mock_polled):
|
||||
with patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"):
|
||||
with patch("litellm.llms.runwayml.text_to_speech.transformation.async_safe_get", new_callable=AsyncMock) as mock_safe_get:
|
||||
mock_safe_get.return_value = mock_audio_response
|
||||
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,
|
||||
|
|
@ -149,4 +155,5 @@ def test_async_transform_text_to_speech_response():
|
|||
|
||||
result = asyncio.run(run_test())
|
||||
assert isinstance(result, HttpxBinaryResponseContent)
|
||||
mock_client.get.assert_called_once()
|
||||
|
||||
|
|
|
|||
|
|
@ -218,6 +218,8 @@ class TestRunwayMLVideoTransformation:
|
|||
"""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",
|
||||
|
|
@ -229,22 +231,26 @@ class TestRunwayMLVideoTransformation:
|
|||
mock_video_response.content = b"fake-video-bytes"
|
||||
mock_video_response.raise_for_status = Mock()
|
||||
|
||||
with patch("litellm.llms.runwayml.videos.transformation._get_httpx_client") as mock_client:
|
||||
with patch("litellm.llms.runwayml.videos.transformation.safe_get") as mock_safe_get:
|
||||
mock_safe_get.return_value = mock_video_response
|
||||
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_safe_get.assert_called_once()
|
||||
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",
|
||||
|
|
@ -256,10 +262,12 @@ class TestRunwayMLVideoTransformation:
|
|||
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("litellm.llms.runwayml.videos.transformation.get_async_httpx_client") as mock_client:
|
||||
with patch("litellm.llms.runwayml.videos.transformation.async_safe_get", new_callable=AsyncMock) as mock_safe_get:
|
||||
mock_safe_get.return_value = mock_video_response
|
||||
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,
|
||||
|
|
@ -268,6 +276,7 @@ class TestRunwayMLVideoTransformation:
|
|||
|
||||
result = asyncio.run(run_test())
|
||||
assert result == b"fake-video-bytes"
|
||||
mock_client.get.assert_called_once()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue