diff --git a/litellm/google_genai/main.py b/litellm/google_genai/main.py index 5bf85bfc55d..b87d5f143e5 100644 --- a/litellm/google_genai/main.py +++ b/litellm/google_genai/main.py @@ -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, ) - diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 93047334180..e0380e26cc3 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -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 \ No newline at end of file diff --git a/tests/test_litellm/google_genai/test_google_genai_main.py b/tests/test_litellm/google_genai/test_google_genai_main.py new file mode 100644 index 00000000000..5854e4b55af --- /dev/null +++ b/tests/test_litellm/google_genai/test_google_genai_main.py @@ -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