mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
test: use mocked HTTP for AIML image generation test (no API key/credits needed)
This commit is contained in:
parent
e3b070ea83
commit
ca7c74e358
1 changed files with 78 additions and 2 deletions
|
|
@ -5,7 +5,7 @@ import logging
|
|||
import os
|
||||
import sys
|
||||
import traceback
|
||||
from unittest.mock import AsyncMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
||||
sys.path.insert(
|
||||
|
|
@ -25,7 +25,7 @@ import pytest
|
|||
import litellm
|
||||
import json
|
||||
import tempfile
|
||||
from base_image_generation_test import BaseImageGenTest
|
||||
from base_image_generation_test import BaseImageGenTest, TestCustomLogger
|
||||
import logging
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
|
|
@ -182,6 +182,82 @@ class TestAimlImageGeneration(BaseImageGenTest):
|
|||
def get_base_image_generation_call_args(self) -> dict:
|
||||
return {"model": "aiml/flux-pro/v1.1"}
|
||||
|
||||
@pytest.mark.asyncio(scope="module")
|
||||
async def test_basic_image_generation(self):
|
||||
"""Test basic image generation"""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
mock_aiml_response = {
|
||||
"created": 1703658209,
|
||||
"data": [{"url": "https://example.com/generated_image.png"}],
|
||||
}
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = mock_aiml_response
|
||||
mock_response.text = json.dumps(mock_aiml_response)
|
||||
mock_response.headers = {}
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_async_post, patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
|
||||
) as mock_sync_post:
|
||||
mock_async_post.return_value = mock_response
|
||||
mock_sync_post.return_value = mock_response
|
||||
|
||||
try:
|
||||
litellm._turn_on_debug()
|
||||
custom_logger = TestCustomLogger()
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
litellm.callbacks = [custom_logger]
|
||||
base_image_generation_call_args = self.get_base_image_generation_call_args()
|
||||
litellm.set_verbose = True
|
||||
response = await litellm.aimage_generation(
|
||||
**base_image_generation_call_args, prompt="A image of a otter"
|
||||
)
|
||||
print("FAL AI RESPONSE: ", response)
|
||||
|
||||
await asyncio.sleep(1)
|
||||
|
||||
# assert response._hidden_params["response_cost"] is not None
|
||||
# assert response._hidden_params["response_cost"] > 0
|
||||
# print("response_cost", response._hidden_params["response_cost"])
|
||||
|
||||
logged_standard_logging_payload = custom_logger.standard_logging_payload
|
||||
print("logged_standard_logging_payload", logged_standard_logging_payload)
|
||||
assert logged_standard_logging_payload is not None
|
||||
assert logged_standard_logging_payload["response_cost"] is not None
|
||||
assert logged_standard_logging_payload["response_cost"] > 0
|
||||
import openai
|
||||
from openai.types.images_response import ImagesResponse
|
||||
|
||||
# print openai version
|
||||
print("openai version=", openai.__version__)
|
||||
|
||||
response_dict = dict(response)
|
||||
if "usage" in response_dict:
|
||||
response_dict["usage"] = dict(response_dict["usage"])
|
||||
print("response usage=", response_dict.get("usage"))
|
||||
|
||||
assert response.data is not None # type guard for iteration (base fails here if None)
|
||||
for d in response.data:
|
||||
assert isinstance(d, Image)
|
||||
print("data in response.data", d)
|
||||
assert d.b64_json is not None or d.url is not None
|
||||
except litellm.RateLimitError as e:
|
||||
pass
|
||||
except litellm.ContentPolicyViolationError:
|
||||
pass # Azure randomly raises these errors - skip when they occur
|
||||
except litellm.InternalServerError:
|
||||
pass
|
||||
except Exception as e:
|
||||
if "Your task failed as a result of our safety system." in str(e):
|
||||
pass
|
||||
else:
|
||||
pytest.fail(f"An exception occurred - {str(e)}")
|
||||
|
||||
|
||||
class TestGoogleImageGen(BaseImageGenTest):
|
||||
def get_base_image_generation_call_args(self) -> dict:
|
||||
return {"model": "gemini/imagen-4.0-generate-001"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue