diff --git a/tests/image_gen_tests/test_image_generation.py b/tests/image_gen_tests/test_image_generation.py index 0567f60ecfc..34ff080c507 100644 --- a/tests/image_gen_tests/test_image_generation.py +++ b/tests/image_gen_tests/test_image_generation.py @@ -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"}