From 3a98fd609677ce39fd85c204d1f290bfa245d7d7 Mon Sep 17 00:00:00 2001 From: Tim Elfrink Date: Fri, 19 Sep 2025 09:29:14 +0200 Subject: [PATCH] Remove hardcoded model name and fix breaking change - Reverted GEMINI_2_5_FLASH_IMAGE_PREVIEW_MODEL constant usage - Made endpoint selection conditional for gemini-2.5-flash-image-preview only - Preserved existing Imagen models functionality with :predict endpoint - Fixed potential breaking change that would affect 6 other Gemini image models --- .../gemini/image_generation/transformation.py | 13 +++-- tests/llm_translation/test_gemini.py | 49 ++++++++++++++++++- 2 files changed, 58 insertions(+), 4 deletions(-) diff --git a/litellm/llms/gemini/image_generation/transformation.py b/litellm/llms/gemini/image_generation/transformation.py index 01ed665290..734bced005 100644 --- a/litellm/llms/gemini/image_generation/transformation.py +++ b/litellm/llms/gemini/image_generation/transformation.py @@ -86,8 +86,8 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): """ Get the complete url for the request - Google AI API format: https://generativelanguage.googleapis.com/v1beta/models/{model}:generateContent - Note: Gemini image generation models use generateContent, not predict + Gemini 2.5 Flash Image Preview: :generateContent + Other Imagen models: :predict """ complete_url: str = ( api_base @@ -96,7 +96,14 @@ class GoogleImageGenConfig(BaseImageGenerationConfig): ) complete_url = complete_url.rstrip("/") - complete_url = f"{complete_url}/models/{model}:generateContent" + + # Gemini 2.5 Flash Image Preview uses generateContent endpoint + if "2.5-flash-image-preview" in model: + complete_url = f"{complete_url}/models/{model}:generateContent" + else: + # All other Imagen models use predict endpoint + complete_url = f"{complete_url}/models/{model}:predict" + return complete_url def validate_environment( diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index 44f921536c..779b0cb96d 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -323,7 +323,7 @@ def test_gemini_2_5_flash_image_preview(): call_args = mock_post.call_args called_url = call_args[0][0] if call_args[0] else call_args.kwargs.get('url', '') - # Verify it uses generateContent endpoint (not predict) + # Verify it uses generateContent endpoint for gemini-2.5-flash-image-preview (not predict) assert ":generateContent" in called_url assert "gemini-2.5-flash-image-preview" in called_url @@ -333,6 +333,53 @@ def test_gemini_2_5_flash_image_preview(): assert "parts" in request_data["contents"][0] +def test_gemini_imagen_models_use_predict_endpoint(): + """ + Test that Imagen models still use :predict endpoint (not broken by gemini-2.5-flash-image-preview fix) + """ + from unittest.mock import patch, MagicMock + from litellm.types.utils import ImageResponse, ImageObject + + with patch("litellm.llms.custom_httpx.llm_http_handler.HTTPHandler.post") as mock_post: + # Mock successful HTTP response for Imagen + mock_http_response = MagicMock() + mock_http_response.json.return_value = { + "predictions": [ + { + "bytesBase64Encoded": "test_base64_image_data" + } + ] + } + mock_http_response.status_code = 200 + mock_post.return_value = mock_http_response + + # Test an Imagen model + response = litellm.image_generation( + model="gemini/imagen-3.0-generate-001", + prompt="Generate a simple test image", + api_key="test_api_key" + ) + + # Validate response structure + assert response is not None + assert hasattr(response, 'data') + + # Validate the correct endpoint was called for Imagen models + mock_post.assert_called_once() + call_args = mock_post.call_args + called_url = call_args[0][0] if call_args[0] else call_args.kwargs.get('url', '') + + # Verify Imagen models use predict endpoint (not generateContent) + assert ":predict" in called_url + assert "imagen-3.0-generate-001" in called_url + assert ":generateContent" not in called_url + + # Verify request format is Imagen format (not Gemini) + request_data = call_args.kwargs.get('json', {}) + assert "instances" in request_data + assert "parameters" in request_data + + def test_gemini_thinking(): litellm._turn_on_debug() from litellm.types.utils import Message, CallTypes