From f08fc45a0fbcf871ae2f37a6de0d74262b2031d0 Mon Sep 17 00:00:00 2001 From: Sameerlite Date: Tue, 16 Sep 2025 15:15:24 +0530 Subject: [PATCH] add base url support for gemini --- litellm/llms/vertex_ai/vertex_llm_base.py | 5 +- .../test_google_gemini_proxy_request.py | 124 +++++++++++++ .../llms/vertex_ai/test_vertex_llm_base.py | 174 ++++++++++++++++++ 3 files changed, 302 insertions(+), 1 deletion(-) diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 76998e7669..f89ee92294 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -292,6 +292,7 @@ class VertexBase: stream: Optional[bool], auth_header: Optional[str], url: str, + model: str, ) -> Tuple[Optional[str], str]: """ for cloudflare ai gateway - https://github.com/BerriAI/litellm/issues/4317 @@ -301,7 +302,8 @@ class VertexBase: """ if api_base: if custom_llm_provider == "gemini": - url = "{}:{}".format(api_base, endpoint) + # For Gemini (Google AI Studio), construct the full path like other providers + url = "{}/models/{}:{}".format(api_base, model, endpoint) if gemini_api_key is None: raise ValueError( "Missing gemini_api_key, please set `GEMINI_API_KEY`" @@ -373,6 +375,7 @@ class VertexBase: endpoint=endpoint, stream=stream, url=url, + model=model, ) def _handle_reauthentication( diff --git a/tests/proxy_unit_tests/test_google_gemini_proxy_request.py b/tests/proxy_unit_tests/test_google_gemini_proxy_request.py index a8c6741940..388dd4ab17 100644 --- a/tests/proxy_unit_tests/test_google_gemini_proxy_request.py +++ b/tests/proxy_unit_tests/test_google_gemini_proxy_request.py @@ -345,6 +345,130 @@ async def test_generationconfig_to_config_mapping(sample_request_payload): print("✅ generationConfig to config mapping test passed") +@pytest.mark.asyncio +async def test_gemini_custom_api_base_proxy_integration(): + """ + Test that Gemini models work correctly with custom API base URLs in proxy context. + + This test verifies that when a custom api_base is provided for Gemini models, + the URL is correctly constructed using the _check_custom_proxy method. + """ + from litellm.llms.vertex_ai.vertex_llm_base import VertexBase + + # Test the _check_custom_proxy method directly + vertex_base = VertexBase() + + # Test case 1: Custom API base for Gemini + custom_api_base = "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta" + model = "gemini-2.5-flash-lite" + endpoint = "generateContent" + + auth_header, result_url = vertex_base._check_custom_proxy( + api_base=custom_api_base, + custom_llm_provider="gemini", + gemini_api_key="test-api-key", + endpoint=endpoint, + stream=False, + auth_header=None, + url=f"https://generativelanguage.googleapis.com/v1beta/models/{model}:{endpoint}", + model=model, + ) + + # Verify the URL is correctly constructed + expected_url = f"{custom_api_base}/models/{model}:{endpoint}" + assert result_url == expected_url, f"Expected {expected_url}, got {result_url}" + + # Verify the auth header is set to the API key + assert auth_header == "test-api-key", f"Expected 'test-api-key', got {auth_header}" + + print(f"✅ Custom API base URL construction test passed: {result_url}") + + # Test case 2: Custom API base with streaming + auth_header_streaming, result_url_streaming = vertex_base._check_custom_proxy( + api_base=custom_api_base, + custom_llm_provider="gemini", + gemini_api_key="test-api-key", + endpoint=endpoint, + stream=True, + auth_header=None, + url=f"https://generativelanguage.googleapis.com/v1beta/models/{model}:{endpoint}", + model=model, + ) + + # Verify streaming URL has ?alt=sse parameter + expected_streaming_url = f"{custom_api_base}/models/{model}:{endpoint}?alt=sse" + assert result_url_streaming == expected_streaming_url, f"Expected {expected_streaming_url}, got {result_url_streaming}" + + print(f"✅ Custom API base streaming URL test passed: {result_url_streaming}") + + # Test case 3: Error handling - missing API key + with pytest.raises(ValueError, match="Missing gemini_api_key"): + vertex_base._check_custom_proxy( + api_base=custom_api_base, + custom_llm_provider="gemini", + gemini_api_key=None, # Missing API key + endpoint=endpoint, + stream=False, + auth_header=None, + url=f"https://generativelanguage.googleapis.com/v1beta/models/{model}:{endpoint}", + model=model, + ) + + print("✅ Missing API key error handling test passed") + + +@pytest.mark.asyncio +async def test_gemini_proxy_config_with_custom_api_base(): + """ + Test that proxy configuration correctly handles custom API base for Gemini models. + + This test simulates the proxy configuration scenario where a model is configured + with a custom api_base in the config.yaml file. + """ + from litellm.llms.vertex_ai.vertex_llm_base import VertexBase + + # Simulate proxy configuration + model_config = { + "model_name": "byok-gemini/*", + "litellm_params": { + "model": "gemini/*", + "api_key": "dummy-key-for-testing", + "api_base": "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta" + } + } + + vertex_base = VertexBase() + + # Test with different Gemini models + test_models = [ + "gemini-2.5-flash-lite", + "gemini-2.5-pro", + "gemini-1.5-flash", + "gemini-1.5-pro" + ] + + for model in test_models: + # Test generateContent endpoint + auth_header, result_url = vertex_base._check_custom_proxy( + api_base=model_config["litellm_params"]["api_base"], + custom_llm_provider="gemini", + gemini_api_key=model_config["litellm_params"]["api_key"], + endpoint="generateContent", + stream=False, + auth_header=None, + url=f"https://generativelanguage.googleapis.com/v1beta/models/{model}:generateContent", + model=model, + ) + + expected_url = f"{model_config['litellm_params']['api_base']}/models/{model}:generateContent" + assert result_url == expected_url, f"Expected {expected_url}, got {result_url} for model {model}" + assert auth_header == model_config["litellm_params"]["api_key"], f"Expected API key, got {auth_header} for model {model}" + + print(f"✅ Model {model} configuration test passed: {result_url}") + + print("✅ Proxy configuration with custom API base test passed") + + if __name__ == "__main__": # Run the tests pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py index c1cefa41ae..bec60434c2 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py @@ -704,3 +704,177 @@ class TestVertexBase: vertex_base.get_api_base(api_base=api_base, vertex_location=vertex_location) == expected ), f"Expected {expected} with api_base {api_base} and vertex_location {vertex_location}" + + @pytest.mark.parametrize( + "api_base, custom_llm_provider, gemini_api_key, endpoint, stream, auth_header, url, model, expected_auth_header, expected_url", + [ + # Test case 1: Gemini with custom API base + ( + "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta", + "gemini", + "test-api-key", + "generateContent", + False, + None, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", + "gemini-2.5-flash-lite", + "test-api-key", + "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent" + ), + # Test case 2: Gemini with custom API base and streaming + ( + "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta", + "gemini", + "test-api-key", + "generateContent", + True, + None, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", + "gemini-2.5-flash-lite", + "test-api-key", + "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" + ), + # Test case 3: Non-Gemini provider with custom API base + ( + "https://custom-vertex-api.com", + "vertex_ai", + None, + "generateContent", + False, + "Bearer token123", + "https://aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent", + "gemini-pro", + "Bearer token123", + "https://custom-vertex-api.com:generateContent" + ), + # Test case 4: No API base provided (should return original values) + ( + None, + "gemini", + "test-api-key", + "generateContent", + False, + "Bearer token123", + "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", + "gemini-2.5-flash-lite", + "Bearer token123", + "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent" + ), + # Test case 5: Gemini without API key (should raise ValueError) + ( + "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta", + "gemini", + None, + "generateContent", + False, + None, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", + "gemini-2.5-flash-lite", + None, # This should raise an exception + None + ), + ], + ) + def test_check_custom_proxy( + self, + api_base, + custom_llm_provider, + gemini_api_key, + endpoint, + stream, + auth_header, + url, + model, + expected_auth_header, + expected_url + ): + """Test the _check_custom_proxy method for handling custom API base URLs""" + vertex_base = VertexBase() + + if custom_llm_provider == "gemini" and api_base and gemini_api_key is None: + # Test case 5: Should raise ValueError for Gemini without API key + with pytest.raises(ValueError, match="Missing gemini_api_key"): + vertex_base._check_custom_proxy( + api_base=api_base, + custom_llm_provider=custom_llm_provider, + gemini_api_key=gemini_api_key, + endpoint=endpoint, + stream=stream, + auth_header=auth_header, + url=url, + model=model, + ) + else: + # Test cases 1-4: Should work correctly + result_auth_header, result_url = vertex_base._check_custom_proxy( + api_base=api_base, + custom_llm_provider=custom_llm_provider, + gemini_api_key=gemini_api_key, + endpoint=endpoint, + stream=stream, + auth_header=auth_header, + url=url, + model=model, + ) + + assert result_auth_header == expected_auth_header, f"Expected auth_header {expected_auth_header}, got {result_auth_header}" + assert result_url == expected_url, f"Expected URL {expected_url}, got {result_url}" + + def test_check_custom_proxy_gemini_url_construction(self): + """Test that Gemini URLs are constructed correctly with custom API base""" + vertex_base = VertexBase() + + # Test various Gemini models with custom API base + test_cases = [ + ("gemini-2.5-flash-lite", "generateContent", "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent"), + ("gemini-2.5-pro", "generateContent", "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent"), + ("gemini-1.5-flash", "streamGenerateContent", "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta/models/gemini-1.5-flash:streamGenerateContent"), + ] + + for model, endpoint, expected_url in test_cases: + _, result_url = vertex_base._check_custom_proxy( + api_base="https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta", + custom_llm_provider="gemini", + gemini_api_key="test-api-key", + endpoint=endpoint, + stream=False, + auth_header=None, + url=f"https://generativelanguage.googleapis.com/v1beta/models/{model}:{endpoint}", + model=model, + ) + + assert result_url == expected_url, f"Expected {expected_url}, got {result_url} for model {model}" + + def test_check_custom_proxy_streaming_parameter(self): + """Test that streaming parameter correctly adds ?alt=sse to URLs""" + vertex_base = VertexBase() + + # Test with streaming enabled + _, result_url_streaming = vertex_base._check_custom_proxy( + api_base="https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta", + custom_llm_provider="gemini", + gemini_api_key="test-api-key", + endpoint="generateContent", + stream=True, + auth_header=None, + url="https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", + model="gemini-2.5-flash-lite", + ) + + expected_streaming_url = "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" + assert result_url_streaming == expected_streaming_url, f"Expected {expected_streaming_url}, got {result_url_streaming}" + + # Test with streaming disabled + _, result_url_no_streaming = vertex_base._check_custom_proxy( + api_base="https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta", + custom_llm_provider="gemini", + gemini_api_key="test-api-key", + endpoint="generateContent", + stream=False, + auth_header=None, + url="https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", + model="gemini-2.5-flash-lite", + ) + + expected_no_streaming_url = "https://proxy.zapier.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent" + assert result_url_no_streaming == expected_no_streaming_url, f"Expected {expected_no_streaming_url}, got {result_url_no_streaming}"