mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(google_genai): preserve proxy_server_request in completion adapter (#43536)
This commit is contained in:
parent
118ce3cc91
commit
319b08b4b1
3 changed files with 61 additions and 40 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue