mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
test(test_gemini.py): add additional testing for additionalproperties case
This commit is contained in:
parent
cb9b65d6ed
commit
fdb9dd58bf
2 changed files with 113 additions and 60 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue