litellm/tests/unit/google_genai/test_google_genai_handler.py

277 lines
11 KiB
Python

#!/usr/bin/env python3
"""
Test to verify the Google GenAI generate_content handler functionality
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.google_genai.adapters.handler import GenerateContentToCompletionHandler
from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter
def test_stream_response_when_stream_requested_sync():
"""
Test that when a stream response is returned and streaming was requested,
the sync handler correctly transforms it to generate_content streaming format.
"""
# Mock a stream response
mock_stream = MagicMock()
mock_stream.__iter__ = MagicMock(return_value=iter([]))
# Mock the GoogleGenAIAdapter's translate_completion_output_params_streaming method
with patch.object(
GoogleGenAIAdapter,
"translate_completion_output_params_streaming",
return_value=mock_stream,
) as mock_translate:
with patch("litellm.completion", return_value=mock_stream):
# Call the handler with stream=True
result = GenerateContentToCompletionHandler.generate_content_handler(
model="gemini-pro",
contents=[{"role": "user", "parts": [{"text": "Hello"}]}],
litellm_params={}, # Empty dict for params
stream=True,
)
# Verify that translate_completion_output_params_streaming was called
mock_translate.assert_called_once_with(mock_stream)
# Verify the result is the transformed stream
assert result == mock_stream
@pytest.mark.asyncio
async def test_stream_response_when_stream_requested_async():
"""
Test that when a stream response is returned and streaming was requested,
the async handler correctly transforms it to generate_content streaming format.
"""
# Mock a stream response
mock_stream = MagicMock()
mock_stream.__aiter__ = AsyncMock(return_value=iter([])) # Return an empty async iterator
# Mock the GoogleGenAIAdapter's translate_completion_output_params_streaming method
with patch.object(
GoogleGenAIAdapter,
"translate_completion_output_params_streaming",
return_value=mock_stream,
) as mock_translate:
with patch("litellm.acompletion", return_value=mock_stream):
# Call the handler with stream=True
result = await GenerateContentToCompletionHandler.async_generate_content_handler(
model="gemini-pro",
contents=[{"role": "user", "parts": [{"text": "Hello"}]}],
litellm_params={}, # Empty dict for params
stream=True,
)
# Verify that translate_completion_output_params_streaming was called
mock_translate.assert_called_once_with(mock_stream)
# Verify the result is the transformed stream
assert result == mock_stream
def test_stream_transformation_error_sync():
"""
Test that when a stream transformation fails, the sync handler raises a ValueError.
"""
# Mock a stream response
mock_stream = MagicMock()
mock_stream.__iter__ = MagicMock(return_value=iter([]))
# Mock the GoogleGenAIAdapter's translate_completion_output_params_streaming method to return None
with patch.object(
GoogleGenAIAdapter,
"translate_completion_output_params_streaming",
return_value=None,
):
# Patch litellm.completion directly to prevent real API calls
with patch("litellm.completion", return_value=mock_stream):
# Call the handler with stream=True and expect a ValueError
with pytest.raises(ValueError, match="Failed to transform streaming response"):
GenerateContentToCompletionHandler.generate_content_handler(
model="gemini-pro",
contents=[{"role": "user", "parts": [{"text": "Hello"}]}],
litellm_params={}, # Empty dict for params
stream=True,
)
@pytest.mark.asyncio
async def test_stream_transformation_error_async():
"""
Test that when a stream transformation fails, the async handler raises a ValueError.
"""
# Mock a stream response
mock_stream = MagicMock()
mock_stream.__aiter__ = AsyncMock(return_value=mock_stream)
# Mock the GoogleGenAIAdapter's translate_completion_output_params_streaming method to return None
with patch.object(
GoogleGenAIAdapter,
"translate_completion_output_params_streaming",
return_value=None,
):
# Mock litellm.acompletion at the module level where it's imported
# We need to patch it in the handler module, not in litellm itself
with patch("litellm.google_genai.adapters.handler.litellm") as mock_litellm:
# Use AsyncMock for async function
mock_litellm.acompletion = AsyncMock(return_value=mock_stream)
# Call the handler with stream=True and expect a ValueError
with pytest.raises(ValueError, match="Failed to transform streaming response"):
await GenerateContentToCompletionHandler.async_generate_content_handler(
model="gemini-pro",
contents=[{"role": "user", "parts": [{"text": "Hello"}]}],
litellm_params={}, # Empty dict for params
stream=True,
)
def test_citation_metadata_transformation():
"""
Test that citationMetadata.citationSources is properly transformed to citationMetadata.citations
to avoid Pydantic validation errors.
"""
from unittest.mock import MagicMock
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.gemini.google_genai.transformation import GoogleGenAIConfig
# Create a mock response with citationMetadata.citationSources (the problematic format)
mock_response_data = {
"candidates": [
{
"content": {
"parts": [{"text": "This is a video analysis response with citation metadata."}],
"role": "model",
},
"finishReason": "STOP",
"index": 0,
"safetyRatings": [],
"citationMetadata": {
"citationSources": [
{
"startIndex": 5848,
"endIndex": 5900,
"uri": "https://example.com/video-source",
"license": "MIT",
"title": "Video Analysis Source",
"publicationDate": "2024-01-15",
},
{
"startIndex": 6200,
"endIndex": 6250,
"uri": "https://another-source.com/reference",
"license": "CC-BY",
"title": "Another Reference",
"publicationDate": "2024-02-01",
},
]
},
}
],
"usageMetadata": {
"promptTokenCount": 150,
"candidatesTokenCount": 200,
"totalTokenCount": 350,
},
"responseId": "test-response-123",
}
# Create mock httpx response
mock_httpx_response = MagicMock(spec=httpx.Response)
mock_httpx_response.json.return_value = mock_response_data
mock_httpx_response.status_code = 200
mock_httpx_response.headers = {}
# Create logging object
logging_obj = LiteLLMLoggingObj(
model="gemini-2.5-flash",
messages=[],
stream=False,
call_type="generate_content",
start_time=1234567890,
litellm_call_id="test-call-123",
function_id="test-function-123",
)
# Create GoogleGenAI config
config = GoogleGenAIConfig()
# Test the transformation
try:
result = config.transform_generate_content_response(
model="gemini-2.5-flash",
raw_response=mock_httpx_response,
logging_obj=logging_obj,
)
# Verify the transformation worked
assert result is not None
# Check that citationSources was transformed to citations
if hasattr(result, "candidates") and result.candidates:
candidate = result.candidates[0]
if hasattr(candidate, "citationMetadata") and candidate.citationMetadata:
# The citationMetadata should now have 'citations' instead of 'citationSources'
citation_metadata = candidate.citationMetadata
# Check that citations field exists
assert hasattr(citation_metadata, "citations"), "citations field should exist after transformation"
# Verify the citations data is preserved
if hasattr(citation_metadata, "citations") and citation_metadata.citations:
assert len(citation_metadata.citations) == 2, "Should have 2 citations"
assert citation_metadata.citations[0]["uri"] == "https://example.com/video-source"
assert citation_metadata.citations[1]["uri"] == "https://another-source.com/reference"
print("✅ Citation metadata transformation test passed!")
except Exception as e:
pytest.fail(f"Citation metadata transformation failed: {e}")
@pytest.mark.asyncio
async def test_generate_content_adapter_preserves_proxy_server_request():
"""
Ensure GenerateContentToCompletionHandler forwards proxy_server_request
to the downstream completion call so proxy spend logging captures the request body.
"""
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import Choices, Message, ModelResponse
handler = GenerateContentToCompletionHandler()
dummy_proxy_request: dict[str, object] = {
"url": "http://localhost:4000/v1beta/models/gemini-2.0-flash:generateContent",
"method": "POST",
"headers": {"content-type": "application/json"},
"body": {"contents": [{"role": "user", "parts": [{"text": "Hello, world!"}]}]},
}
gemini_data: list[dict[str, object]] = [{"role": "user", "parts": [{"text": "Hello, world!"}]}]
mock_response = ModelResponse(choices=[Choices(message=Message(content="Hi!", role="assistant"))])
with patch(
"litellm.google_genai.adapters.handler.litellm.acompletion",
new_callable=AsyncMock,
) as mock_acompletion:
mock_acompletion.return_value = mock_response
await handler.async_generate_content_handler(
model="gemini-2.0-flash",
contents=gemini_data,
litellm_params=GenericLiteLLMParams(),
proxy_server_request=dummy_proxy_request,
metadata={"source": "unit_test"},
)
assert mock_acompletion.called, "Inner acompletion was not called"
called_kwargs = mock_acompletion.call_args.kwargs
assert "proxy_server_request" in called_kwargs, "proxy_server_request was dropped from completion_kwargs"
assert called_kwargs["proxy_server_request"] == dummy_proxy_request