test(test_gemini.py): add additional testing for additionalproperties case

This commit is contained in:
Krrish Dholakia 2025-09-11 15:12:21 -07:00
parent cb9b65d6ed
commit fdb9dd58bf
2 changed files with 113 additions and 60 deletions

View file

@ -302,20 +302,16 @@ def test_gemini_2_5_flash_image_preview():
mock_response = ImageResponse()
mock_response.data = [ImageObject(b64_json="test_base64_data", url=None)]
with patch("litellm.llms.custom_httpx.llm_http_handler.HTTPHandler.post") as mock_post:
with patch(
"litellm.llms.custom_httpx.llm_http_handler.HTTPHandler.post"
) as mock_post:
# Mock successful HTTP response
mock_http_response = MagicMock()
mock_http_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"data": "test_base64_image_data"
}
}
]
"parts": [{"inlineData": {"data": "test_base64_image_data"}}]
}
}
]
@ -327,33 +323,38 @@ def test_gemini_2_5_flash_image_preview():
response = litellm.image_generation(
model="gemini/gemini-2.5-flash-image-preview",
prompt="Generate a simple test image",
api_key="test_api_key"
api_key="test_api_key",
)
# Validate response structure
assert response is not None
assert hasattr(response, 'data')
assert hasattr(response, "data")
assert response.data is not None
assert len(response.data) > 0
# Validate the correct endpoint was called
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', '')
called_url = (
call_args[0][0] if call_args[0] else call_args.kwargs.get("url", "")
)
# 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
# Verify request format is Gemini format (not Imagen)
request_data = call_args.kwargs.get('json', {})
request_data = call_args.kwargs.get("json", {})
assert "contents" in request_data
assert "parts" in request_data["contents"][0]
# Verify response_modalities is set correctly for image generation
assert "generationConfig" in request_data
assert "response_modalities" in request_data["generationConfig"]
assert request_data["generationConfig"]["response_modalities"] == ["IMAGE", "TEXT"]
assert request_data["generationConfig"]["response_modalities"] == [
"IMAGE",
"TEXT",
]
def test_gemini_imagen_models_use_predict_endpoint():
@ -363,15 +364,13 @@ def test_gemini_imagen_models_use_predict_endpoint():
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:
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"
}
]
"predictions": [{"bytesBase64Encoded": "test_base64_image_data"}]
}
mock_http_response.status_code = 200
mock_post.return_value = mock_http_response
@ -380,17 +379,19 @@ def test_gemini_imagen_models_use_predict_endpoint():
response = litellm.image_generation(
model="gemini/imagen-3.0-generate-001",
prompt="Generate a simple test image",
api_key="test_api_key"
api_key="test_api_key",
)
# Validate response structure
assert response is not None
assert hasattr(response, 'data')
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', '')
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
@ -398,7 +399,7 @@ def test_gemini_imagen_models_use_predict_endpoint():
assert ":generateContent" not in called_url
# Verify request format is Imagen format (not Gemini)
request_data = call_args.kwargs.get('json', {})
request_data = call_args.kwargs.get("json", {})
assert "instances" in request_data
assert "parameters" in request_data
@ -997,9 +998,7 @@ def test_gemini_exception_message_format():
# Create a mock exception that simulates a Gemini API error
mock_exception = httpx.HTTPStatusError(
message="Bad Request",
request=Mock(),
response=mock_response
message="Bad Request", request=Mock(), response=mock_response
)
mock_exception.response = mock_response
mock_exception.status_code = 400
@ -1011,7 +1010,7 @@ def test_gemini_exception_message_format():
original_exception=mock_exception,
custom_llm_provider="gemini",
completion_kwargs={},
extra_kwargs={}
extra_kwargs={},
)
# Should not reach here - exception should be raised
assert False, "Expected BadRequestError to be raised"
@ -1026,22 +1025,25 @@ def test_gemini_exception_message_format():
f"Expected 'GeminiException' in error message, got: {error_message}. "
f"This test should fail before the fix is implemented."
)
assert "VertexAIException" not in error_message, (
f"Should not contain 'VertexAIException' in error message, got: {error_message}"
)
assert (
"VertexAIException" not in error_message
), f"Should not contain 'VertexAIException' in error message, got: {error_message}"
@pytest.mark.parametrize("status_code,expected_exception", [
(400, "BadRequestError"),
(401, "AuthenticationError"),
(403, "PermissionDeniedError"),
(404, "NotFoundError"),
(408, "Timeout"),
(429, "RateLimitError"),
(500, "InternalServerError"),
(502, "APIConnectionError"),
(503, "ServiceUnavailableError"),
])
@pytest.mark.parametrize(
"status_code,expected_exception",
[
(400, "BadRequestError"),
(401, "AuthenticationError"),
(403, "PermissionDeniedError"),
(404, "NotFoundError"),
(408, "Timeout"),
(429, "RateLimitError"),
(500, "InternalServerError"),
(502, "APIConnectionError"),
(503, "ServiceUnavailableError"),
],
)
def l(status_code, expected_exception):
"""
Test comprehensive Gemini error handling for all HTTP status codes.
@ -1053,8 +1055,15 @@ def l(status_code, expected_exception):
from unittest.mock import Mock
from litellm.litellm_core_utils.exception_mapping_utils import exception_type
from litellm.exceptions import (
BadRequestError, AuthenticationError, PermissionDeniedError, NotFoundError,
Timeout, RateLimitError, InternalServerError, APIConnectionError, ServiceUnavailableError
BadRequestError,
AuthenticationError,
PermissionDeniedError,
NotFoundError,
Timeout,
RateLimitError,
InternalServerError,
APIConnectionError,
ServiceUnavailableError,
)
# Mock the appropriate error response
@ -1065,9 +1074,7 @@ def l(status_code, expected_exception):
# Create a mock exception
mock_exception = httpx.HTTPStatusError(
message=f"HTTP {status_code}",
request=Mock(),
response=mock_response
message=f"HTTP {status_code}", request=Mock(), response=mock_response
)
mock_exception.response = mock_response
mock_exception.status_code = status_code
@ -1081,9 +1088,11 @@ def l(status_code, expected_exception):
original_exception=mock_exception,
custom_llm_provider="gemini",
completion_kwargs={},
extra_kwargs={}
extra_kwargs={},
)
assert False, f"Expected {expected_exception} to be raised for status {status_code}"
assert (
False
), f"Expected {expected_exception} to be raised for status {status_code}"
except Exception as e:
# Verify the correct exception type is raised
exception_classes = {
@ -1098,13 +1107,40 @@ def l(status_code, expected_exception):
"ServiceUnavailableError": ServiceUnavailableError,
}
expected_class = exception_classes[expected_exception]
assert isinstance(e, expected_class), f"Expected {expected_exception}, got {type(e).__name__}"
assert isinstance(
e, expected_class
), f"Expected {expected_exception}, got {type(e).__name__}"
# Verify the error message contains GeminiException
error_message = str(e)
assert "GeminiException" in error_message, (
f"Expected 'GeminiException' in error message for status {status_code}, got: {error_message}"
)
assert "VertexAIException" not in error_message, (
f"Should not contain 'VertexAIException' for status {status_code}, got: {error_message}"
)
assert (
"GeminiException" in error_message
), f"Expected 'GeminiException' in error message for status {status_code}, got: {error_message}"
assert (
"VertexAIException" not in error_message
), f"Should not contain 'VertexAIException' for status {status_code}, got: {error_message}"
def test_gemini_additional_properties_bug():
# Simple tool with additionalProperties (simulating the TypedDict issue)
tools = [
{
"type": "function",
"function": {
"name": "test_tool",
"description": "Test tool",
"parameters": {
"type": "object",
"properties": {"param1": {"type": "string"}},
# This causes the error - any non-False value
"additionalProperties": True, # Could also be None, {}, etc.
},
},
}
]
messages = [{"role": "user", "content": "Test message"}]
response = litellm.completion(
model="gemini/gemini-2.5-flash", messages=messages, tools=tools
)

View file

@ -764,7 +764,9 @@ def test_gemini_pro_grounding(value_in_dict):
# @pytest.mark.skip(reason="exhausted vertex quota. need to refactor to mock the call")
@pytest.mark.parametrize("model", ["vertex_ai_beta/gemini-2.5-flash-lite"]) # "vertex_ai",
@pytest.mark.parametrize(
"model", ["vertex_ai_beta/gemini-2.5-flash-lite"]
) # "vertex_ai",
@pytest.mark.parametrize("sync_mode", [True]) # "vertex_ai",
@pytest.mark.asyncio
@pytest.mark.flaky(retries=3, delay=1)
@ -914,6 +916,10 @@ async def test_partner_models_httpx(model, region, sync_mode):
"vertex_ai/mistral-large-2411",
"us-central1",
), # critical - we had this issue: https://github.com/BerriAI/litellm/issues/13888
(
"vertex_ai/mistral-large-2411",
"us-central1",
), # critical - we had this issue: https://github.com/BerriAI/litellm/issues/13888
("vertex_ai/openai/gpt-oss-20b-maas", "us-central1"),
],
)
@ -2329,8 +2335,6 @@ def test_prompt_factory_nested():
), "'text' value not a string."
@pytest.mark.asyncio
async def test_completion_fine_tuned_model():
load_vertex_ai_credentials()
@ -3777,6 +3781,7 @@ def test_vertex_ai_gemini_audio_ogg():
async def test_vertex_ai_deepseek():
"""Test that deepseek models use the correct v1 API endpoint instead of v1beta1."""
# load_vertex_ai_credentials()
# load_vertex_ai_credentials()
litellm._turn_on_debug()
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
@ -3790,13 +3795,17 @@ async def test_vertex_ai_deepseek():
"message": {
"role": "assistant",
"content": "Hello! How can I help you today?",
"content": "Hello! How can I help you today?",
},
"index": 0,
"finish_reason": "stop",
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
"model": "deepseek-ai/deepseek-r1-0528-maas",
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
"model": "deepseek-ai/deepseek-r1-0528-maas",
}
mock_response.status_code = 200
@ -3855,7 +3864,16 @@ def test_gemini_google_maps_tool_simple():
litellm._turn_on_debug()
tools = [{"googleMaps": {"enableWidget": True}}]
tools_with_location = [{"googleMaps": {"enableWidget": True, "latitude": 37.7749, "longitude": -122.4194, "languageCode": "en_US"}}]
tools_with_location = [
{
"googleMaps": {
"enableWidget": True,
"latitude": 37.7749,
"longitude": -122.4194,
"languageCode": "en_US",
}
}
]
try:
for tools in [tools, tools_with_location]:
response = completion(
@ -3874,4 +3892,3 @@ def test_gemini_google_maps_tool_simple():
pass
except Exception as e:
pytest.fail(f"Error occurred: {e}")