fix(google_genai): preserve proxy_server_request in completion adapter (#43536)

This commit is contained in:
hsm207 2026-09-29 06:29:55 +02:00 • committed by GitHub
parent 118ce3cc91
commit 319b08b4b1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 61 additions and 40 deletions

View file

@ -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

View file

@ -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

View file

@ -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