test: use mocked HTTP for AIML image generation test (no API key/credits needed)

This commit is contained in:
Alexsander Hamir 2026-02-10 11:34:02 -08:00
parent e3b070ea83
commit ca7c74e358

View file

@ -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"}