mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(google_genai/main.py): handle stream=true being set in kwargs
This commit is contained in:
parent
5f8b1e9fd2
commit
c0c6530bb8
3 changed files with 112 additions and 65 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
45
tests/test_litellm/google_genai/test_google_genai_main.py
Normal file
45
tests/test_litellm/google_genai/test_google_genai_main.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue