diff --git a/litellm/google_genai/adapters/handler.py b/litellm/google_genai/adapters/handler.py index 8df71504850..ed06aca0809 100644 --- a/litellm/google_genai/adapters/handler.py +++ b/litellm/google_genai/adapters/handler.py @@ -45,6 +45,8 @@ class GenerateContentToCompletionHandler: # Forward extra_headers for providers that require custom headers (e.g., github_copilot) if "extra_headers" in extra_kwargs: completion_kwargs["extra_headers"] = extra_kwargs["extra_headers"] + if "proxy_server_request" in extra_kwargs: + completion_kwargs["proxy_server_request"] = extra_kwargs["proxy_server_request"] if stream: completion_kwargs["stream"] = stream diff --git a/litellm/types/google_genai/adapters.py b/litellm/types/google_genai/adapters.py index 172a45b4cbc..771b362cae3 100644 --- a/litellm/types/google_genai/adapters.py +++ b/litellm/types/google_genai/adapters.py @@ -19,3 +19,4 @@ class GenerateContentCompletionKwargs(TypedDict, total=False): stream: bool metadata: dict[str, object] extra_headers: dict[str, str] | None + proxy_server_request: dict[str, object] | None diff --git a/tests/unit/google_genai/test_google_genai_handler.py b/tests/unit/google_genai/test_google_genai_handler.py index 5361d91718d..68f24c1c0f5 100644 --- a/tests/unit/google_genai/test_google_genai_handler.py +++ b/tests/unit/google_genai/test_google_genai_handler.py @@ -2,11 +2,11 @@ """ 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 @@ -49,9 +49,7 @@ async def test_stream_response_when_stream_requested_async(): """ # Mock a stream response mock_stream = MagicMock() - mock_stream.__aiter__ = AsyncMock( - return_value=iter([]) - ) # Return an empty async iterator + 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( @@ -61,13 +59,11 @@ async def test_stream_response_when_stream_requested_async(): ) 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, - ) + 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 @@ -93,9 +89,7 @@ def test_stream_transformation_error_sync(): # 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" - ): + with pytest.raises(ValueError, match="Failed to transform streaming response"): GenerateContentToCompletionHandler.generate_content_handler( model="gemini-pro", contents=[{"role": "user", "parts": [{"text": "Hello"}]}], @@ -125,9 +119,7 @@ async def test_stream_transformation_error_async(): # 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" - ): + 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"}]}], @@ -153,11 +145,7 @@ def test_citation_metadata_transformation(): "candidates": [ { "content": { - "parts": [ - { - "text": "This is a video analysis response with citation metadata." - } - ], + "parts": [{"text": "This is a video analysis response with citation metadata."}], "role": "model", }, "finishReason": "STOP", @@ -232,28 +220,58 @@ def test_citation_metadata_transformation(): citation_metadata = candidate.citationMetadata # Check that citations field exists - assert hasattr( - citation_metadata, "citations" - ), "citations field should exist after transformation" + 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" - ) + 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