fix(google_genai/main.py): handle stream=true being set in kwargs

This commit is contained in:
Krrish Dholakia 2025-07-01 17:24:14 -07:00
parent 5f8b1e9fd2
commit c0c6530bb8
3 changed files with 112 additions and 65 deletions

View file

@ -29,7 +29,7 @@ else:
GenerateContentConfigDict = Any
GenerateContentContentListUnionDict = Any
GenerateContentResponse = Any
####### ENVIRONMENT VARIABLES ###################
# Initialize any necessary instances or variables here
base_llm_http_handler = BaseLLMHTTPHandler()
@ -38,6 +38,7 @@ base_llm_http_handler = BaseLLMHTTPHandler()
class GenerateContentSetupResult(BaseModel):
"""Internal Type - Result of setting up a generate content call"""
model: str
request_body: Dict[str, Any]
custom_llm_provider: str
@ -53,7 +54,7 @@ class GenerateContentSetupResult(BaseModel):
class GenerateContentHelper:
"""Helper class for Google GenAI generate content operations"""
@staticmethod
def mock_generate_content_response(
mock_response: str = "This is a mock response from Google GenAI generate_content.",
@ -63,20 +64,17 @@ class GenerateContentHelper:
"text": mock_response,
"candidates": [
{
"content": {
"parts": [{"text": mock_response}],
"role": "model"
},
"content": {"parts": [{"text": mock_response}], "role": "model"},
"finishReason": "STOP",
"index": 0,
"safetyRatings": []
"safetyRatings": [],
}
],
"usageMetadata": {
"promptTokenCount": 10,
"candidatesTokenCount": 20,
"totalTokenCount": 30
}
"totalTokenCount": 30,
},
}
@staticmethod
@ -86,11 +84,11 @@ class GenerateContentHelper:
config: Optional[GenerateContentConfigDict] = None,
custom_llm_provider: Optional[str] = None,
stream: bool = False,
**kwargs
**kwargs,
) -> GenerateContentSetupResult:
"""
Common setup logic for generate_content calls
Args:
model: The model name
contents: The content to generate from
@ -99,18 +97,24 @@ class GenerateContentHelper:
stream: Whether this is a streaming call
local_vars: Local variables from the calling function
**kwargs: Additional keyword arguments
Returns:
GenerateContentSetupResult containing all setup information
"""
litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get(
"litellm_logging_obj"
)
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
# get llm provider logic
litellm_params = GenericLiteLLMParams(**kwargs)
## MOCK RESPONSE LOGIC (only for non-streaming)
if not stream and litellm_params.mock_response and isinstance(litellm_params.mock_response, str):
if (
not stream
and litellm_params.mock_response
and isinstance(litellm_params.mock_response, str)
):
raise ValueError("Mock response should be handled by caller")
(
@ -126,11 +130,11 @@ class GenerateContentHelper:
)
# get provider config
generate_content_provider_config: Optional[BaseGoogleGenAIGenerateContentConfig] = (
ProviderConfigManager.get_provider_google_genai_generate_content_config(
model=model,
provider=litellm.LlmProviders(custom_llm_provider),
)
generate_content_provider_config: Optional[
BaseGoogleGenAIGenerateContentConfig
] = ProviderConfigManager.get_provider_google_genai_generate_content_config(
model=model,
provider=litellm.LlmProviders(custom_llm_provider),
)
if generate_content_provider_config is None:
@ -146,28 +150,31 @@ class GenerateContentHelper:
generate_content_config_dict=dict(config or {}),
litellm_params=litellm_params,
litellm_logging_obj=litellm_logging_obj,
litellm_call_id=litellm_call_id
litellm_call_id=litellm_call_id,
)
#########################################################################################
# Construct request body
#########################################################################################
# Create Google Optional Params Config
generate_content_config_dict = generate_content_provider_config.map_generate_content_optional_params(
generate_content_config_dict=config or {},
model=model,
generate_content_config_dict = (
generate_content_provider_config.map_generate_content_optional_params(
generate_content_config_dict=config or {},
model=model,
)
)
request_body = generate_content_provider_config.transform_generate_content_request(
model=model,
contents=contents,
generate_content_config_dict=generate_content_config_dict,
request_body = (
generate_content_provider_config.transform_generate_content_request(
model=model,
contents=contents,
generate_content_config_dict=generate_content_config_dict,
)
)
# Pre Call logging
if litellm_logging_obj is None:
raise ValueError("litellm_logging_obj is required, but got None")
litellm_logging_obj.update_environment_variables(
model=model,
optional_params=dict(generate_content_config_dict),
@ -185,7 +192,7 @@ class GenerateContentHelper:
generate_content_config_dict=generate_content_config_dict,
litellm_params=litellm_params,
litellm_logging_obj=litellm_logging_obj,
litellm_call_id=litellm_call_id
litellm_call_id=litellm_call_id,
)
@ -202,7 +209,7 @@ async def agenerate_content(
timeout: Optional[Union[float, httpx.Timeout]] = None,
# LiteLLM specific params,
custom_llm_provider: Optional[str] = None,
**kwargs
**kwargs,
) -> Any:
"""
Async: Generate content using Google GenAI
@ -273,10 +280,12 @@ def generate_content(
local_vars = locals()
try:
_is_async = kwargs.pop("agenerate_content", False) is True
# Check for mock response first
litellm_params = GenericLiteLLMParams(**kwargs)
if litellm_params.mock_response and isinstance(litellm_params.mock_response, str):
if litellm_params.mock_response and isinstance(
litellm_params.mock_response, str
):
return GenerateContentHelper.mock_generate_content_response(
mock_response=litellm_params.mock_response
)
@ -288,7 +297,7 @@ def generate_content(
config=config,
custom_llm_provider=custom_llm_provider,
stream=False,
**kwargs
**kwargs,
)
# Check if we should use the adapter (when provider config is None)
@ -301,7 +310,7 @@ def generate_content(
stream=False,
_is_async=_is_async,
litellm_params=setup_result.litellm_params,
**kwargs
**kwargs,
)
# Call the standard handler
@ -346,7 +355,7 @@ async def agenerate_content_stream(
timeout: Optional[Union[float, httpx.Timeout]] = None,
# LiteLLM specific params,
custom_llm_provider: Optional[str] = None,
**kwargs
**kwargs,
) -> Any:
"""
Async: Generate content using Google GenAI with streaming response
@ -354,7 +363,7 @@ async def agenerate_content_stream(
local_vars = locals()
try:
kwargs["agenerate_content_stream"] = True
# get custom llm provider so we can use this for mapping exceptions
if custom_llm_provider is None:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
@ -363,24 +372,28 @@ async def agenerate_content_stream(
# Setup the call
setup_result = GenerateContentHelper.setup_generate_content_call(
model=model,
contents=contents,
config=config,
custom_llm_provider=custom_llm_provider,
stream=True,
**kwargs
**{
"model": model,
"contents": contents,
"config": config,
"custom_llm_provider": custom_llm_provider,
"stream": True,
**kwargs,
}
)
# Check if we should use the adapter (when provider config is None)
if setup_result.generate_content_provider_config is None:
# Use the adapter to convert to completion format
return await GenerateContentToCompletionHandler.async_generate_content_handler(
model=setup_result.model,
contents=contents, # type: ignore
config=setup_result.generate_content_config_dict,
litellm_params=setup_result.litellm_params,
stream=True,
**kwargs
return (
await GenerateContentToCompletionHandler.async_generate_content_handler(
model=setup_result.model,
contents=contents, # type: ignore
config=setup_result.generate_content_config_dict,
litellm_params=setup_result.litellm_params,
stream=True,
**kwargs,
)
)
# Call the handler with async enabled and streaming
@ -401,7 +414,7 @@ async def agenerate_content_stream(
stream=True,
litellm_metadata=kwargs.get("litellm_metadata", {}),
)
except Exception as e:
raise litellm.exception_type(
model=model,
@ -442,7 +455,7 @@ def generate_content_stream(
config=config,
custom_llm_provider=custom_llm_provider,
stream=True,
**kwargs
**kwargs,
)
# Check if we should use the adapter (when provider config is None)
@ -455,7 +468,7 @@ def generate_content_stream(
stream=True,
_is_async=_is_async,
litellm_params=setup_result.litellm_params,
**kwargs
**kwargs,
)
# Call the handler with streaming enabled (sync version)
@ -484,4 +497,3 @@ def generate_content_stream(
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)

View file

@ -3,13 +3,3 @@ model_list:
litellm_params:
model: "gemini/*"
api_key: os.environ/GEMINI_API_KEY
- model_name: "[IP-approved] o3-pro"
litellm_params:
model: azure/responses/o_series/webinterface-o3-pro
api_base: "https://webhook.site/fba79dae-220a-4bb7-9a3a-8caa49604e55"
api_key: "sk-1234567890"
api_version: "preview"
stream: True
model_info:
input_cost_per_token: 0.00002 # $20 per 1M tokens
output_cost_per_token: 0.00008 # $80 per 1M tokens

View file

@ -0,0 +1,45 @@
#!/usr/bin/env python3
"""
Test to verify the Google GenAI generate_content adapter functionality
"""
import json
import os
import sys
import pytest
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import json
import os
import sys
import pytest
import litellm
@pytest.mark.asyncio
async def test_agenerate_content_stream():
"""
Test that the agenerate_content_stream function works
"""
from unittest.mock import AsyncMock, patch
from litellm.google_genai.main import (
agenerate_content_stream,
base_llm_http_handler,
)
with patch.object(
base_llm_http_handler, "generate_content_handler", new=AsyncMock()
) as mock_post:
result = await agenerate_content_stream(
model="gemini/gemini-2.0-flash-001",
contents="Hello, world!",
stream=True,
)
mock_post.assert_called_once()
mock_post.call_args.kwargs["stream"] == True