From 3f8e985b58769aa898aef8cf3db712875b128102 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 20 Jan 2026 12:36:25 +0530 Subject: [PATCH] Fix: skip auth for custom api base in vertex ai --- .../vertex_and_google_ai_studio_gemini.py | 24 +-- .../vertex_ai_partner_models/main.py | 11 +- .../vertex_embeddings/embedding_handler.py | 14 +- .../vertex_ai/vertex_gemma_models/main.py | 11 +- litellm/llms/vertex_ai/vertex_llm_base.py | 78 ++++++++-- .../vertex_ai/vertex_model_garden/main.py | 11 +- .../llms/vertex_ai/test_vertex_llm_base.py | 143 ++++++++++++++++++ 7 files changed, 258 insertions(+), 34 deletions(-) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index f65a19ac46f..c465422ec2e 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -2329,15 +2329,17 @@ class VertexLLM(VertexBase): optional_params=optional_params ) + # Extract use_psc_endpoint_format from optional_params + use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) + _auth_header, vertex_project = await self._ensure_access_token_async( credentials=vertex_credentials, project_id=vertex_project, custom_llm_provider=custom_llm_provider, + api_base=api_base, + use_psc_endpoint_format=use_psc_endpoint_format, ) - # Extract use_psc_endpoint_format from optional_params - use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) - auth_header, api_base = self._get_token_and_url( model=model, gemini_api_key=gemini_api_key, @@ -2427,15 +2429,17 @@ class VertexLLM(VertexBase): optional_params=optional_params ) + # Extract use_psc_endpoint_format from optional_params + use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) + _auth_header, vertex_project = await self._ensure_access_token_async( credentials=vertex_credentials, project_id=vertex_project, custom_llm_provider=custom_llm_provider, + api_base=api_base, + use_psc_endpoint_format=use_psc_endpoint_format, ) - # Extract use_psc_endpoint_format from optional_params - use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) - auth_header, api_base = self._get_token_and_url( model=model, gemini_api_key=gemini_api_key, @@ -2615,15 +2619,17 @@ class VertexLLM(VertexBase): optional_params=optional_params ) + # Extract use_psc_endpoint_format from optional_params + use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) + _auth_header, vertex_project = self._ensure_access_token( credentials=vertex_credentials, project_id=vertex_project, custom_llm_provider=custom_llm_provider, + api_base=api_base, + use_psc_endpoint_format=use_psc_endpoint_format, ) - # Extract use_psc_endpoint_format from optional_params - use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) - auth_header, url = self._get_token_and_url( model=model, gemini_api_key=gemini_api_key, diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py index 123d925f7c1..d2d8b995868 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py @@ -135,19 +135,24 @@ class VertexAIPartnerModels(VertexBase): try: vertex_httpx_logic = VertexLLM() + ## CONSTRUCT API BASE + stream: bool = optional_params.get("stream", False) or False + + # Extract use_psc_endpoint_format from optional_params + use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) + access_token, project_id = vertex_httpx_logic._ensure_access_token( credentials=vertex_credentials, project_id=vertex_project, custom_llm_provider="vertex_ai", + api_base=api_base, + use_psc_endpoint_format=use_psc_endpoint_format, ) openai_like_chat_completions = OpenAILikeChatHandler() codestral_fim_completions = CodestralTextCompletion() anthropic_chat_completions = AnthropicChatCompletion() - ## CONSTRUCT API BASE - stream: bool = optional_params.get("stream", False) or False - optional_params["stream"] = stream if self.should_use_openai_handler(model): diff --git a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py index 8a03738ad78..0aceddbd5bc 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py @@ -67,13 +67,16 @@ class VertexEmbedding(VertexBase): optional_params=optional_params ) + # Extract use_psc_endpoint_format from optional_params + use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) + _auth_header, vertex_project = self._ensure_access_token( credentials=vertex_credentials, project_id=vertex_project, custom_llm_provider=custom_llm_provider, + api_base=api_base, + use_psc_endpoint_format=use_psc_endpoint_format, ) - # Extract use_psc_endpoint_format from optional_params - use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) auth_header, api_base = self._get_token_and_url( model=model, @@ -163,13 +166,16 @@ class VertexEmbedding(VertexBase): should_use_v1beta1_features = self.is_using_v1beta1_features( optional_params=optional_params ) + # Extract use_psc_endpoint_format from optional_params + use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) + _auth_header, vertex_project = await self._ensure_access_token_async( credentials=vertex_credentials, project_id=vertex_project, custom_llm_provider=custom_llm_provider, + api_base=api_base, + use_psc_endpoint_format=use_psc_endpoint_format, ) - # Extract use_psc_endpoint_format from optional_params - use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) auth_header, api_base = self._get_token_and_url( model=model, diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/main.py b/litellm/llms/vertex_ai/vertex_gemma_models/main.py index 41bd6b5431e..2a1ed9b7284 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/main.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/main.py @@ -86,16 +86,21 @@ class VertexAIGemmaModels(VertexBase): model = get_vertex_base_model_name(model=model) vertex_httpx_logic = VertexLLM() + ## CONSTRUCT API BASE + stream: bool = optional_params.get("stream", False) or False + + # Extract use_psc_endpoint_format from optional_params + use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) + access_token, project_id = vertex_httpx_logic._ensure_access_token( credentials=vertex_credentials, project_id=vertex_project, custom_llm_provider="vertex_ai", + api_base=api_base, + use_psc_endpoint_format=use_psc_endpoint_format, ) gemma_transformation = VertexGemmaConfig() - - ## CONSTRUCT API BASE - stream: bool = optional_params.get("stream", False) or False optional_params["stream"] = stream # If api_base is not provided, it should be set as an environment variable diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 826f151df35..dbbc3c1d030 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -274,17 +274,44 @@ class VertexBase: custom_llm_provider: Literal[ "vertex_ai", "vertex_ai_beta", "gemini" ], # if it's vertex_ai or gemini (google ai studio) + api_base: Optional[str] = None, + use_psc_endpoint_format: bool = False, ) -> Tuple[str, str]: """ Returns auth token and project id + + Args: + credentials: Vertex AI credentials + project_id: Google Cloud project ID + custom_llm_provider: Provider type (vertex_ai, vertex_ai_beta, or gemini) + api_base: Custom API base URL (e.g., for proxies) + use_psc_endpoint_format: Whether using PSC endpoint format + + Returns: + Tuple of (access_token, project_id) + + Note: + Authentication is skipped when using a custom api_base that is not a PSC endpoint. + PSC endpoints still require Google authentication even with custom api_base. """ if custom_llm_provider == "gemini": return "", "" - else: - return self.get_access_token( - credentials=credentials, - project_id=project_id, + + # Skip authentication if custom api_base is provided and it's not a PSC endpoint + if api_base is not None and not use_psc_endpoint_format: + verbose_logger.debug( + "Skipping Vertex AI authentication - custom api_base provided without PSC endpoint format" ) + # Return empty token and use provided project_id or empty string + return "", project_id or "" + + # Perform authentication for: + # 1. No custom api_base (standard Vertex AI) + # 2. Custom api_base with PSC endpoint format (PSC endpoints need auth) + return self.get_access_token( + credentials=credentials, + project_id=project_id, + ) def is_using_v1beta1_features(self, optional_params: dict) -> bool: """ @@ -626,20 +653,47 @@ class VertexBase: custom_llm_provider: Literal[ "vertex_ai", "vertex_ai_beta", "gemini" ], # if it's vertex_ai or gemini (google ai studio) + api_base: Optional[str] = None, + use_psc_endpoint_format: bool = False, ) -> Tuple[str, str]: """ Async version of _ensure_access_token + + Args: + credentials: Vertex AI credentials + project_id: Google Cloud project ID + custom_llm_provider: Provider type (vertex_ai, vertex_ai_beta, or gemini) + api_base: Custom API base URL (e.g., for proxies) + use_psc_endpoint_format: Whether using PSC endpoint format + + Returns: + Tuple of (access_token, project_id) + + Note: + Authentication is skipped when using a custom api_base that is not a PSC endpoint. + PSC endpoints still require Google authentication even with custom api_base. """ if custom_llm_provider == "gemini": return "", "" - else: - try: - return await asyncify(self.get_access_token)( - credentials=credentials, - project_id=project_id, - ) - except Exception as e: - raise e + + # Skip authentication if custom api_base is provided and it's not a PSC endpoint + if api_base is not None and not use_psc_endpoint_format: + verbose_logger.debug( + "Skipping Vertex AI authentication - custom api_base provided without PSC endpoint format" + ) + # Return empty token and use provided project_id or empty string + return "", project_id or "" + + # Perform authentication for: + # 1. No custom api_base (standard Vertex AI) + # 2. Custom api_base with PSC endpoint format (PSC endpoints need auth) + try: + return await asyncify(self.get_access_token)( + credentials=credentials, + project_id=project_id, + ) + except Exception as e: + raise e def set_headers( self, auth_header: Optional[str], extra_headers: Optional[dict] diff --git a/litellm/llms/vertex_ai/vertex_model_garden/main.py b/litellm/llms/vertex_ai/vertex_model_garden/main.py index c37bb449ecf..5cec60ac25f 100644 --- a/litellm/llms/vertex_ai/vertex_model_garden/main.py +++ b/litellm/llms/vertex_ai/vertex_model_garden/main.py @@ -93,16 +93,21 @@ class VertexAIModelGardenModels(VertexBase): model = get_vertex_base_model_name(model=model) vertex_httpx_logic = VertexLLM() + ## CONSTRUCT API BASE + stream: bool = optional_params.get("stream", False) or False + + # Extract use_psc_endpoint_format from optional_params + use_psc_endpoint_format = optional_params.get("use_psc_endpoint_format", False) + access_token, project_id = vertex_httpx_logic._ensure_access_token( credentials=vertex_credentials, project_id=vertex_project, custom_llm_provider="vertex_ai", + api_base=api_base, + use_psc_endpoint_format=use_psc_endpoint_format, ) openai_like_chat_completions = OpenAILikeChatHandler() - - ## CONSTRUCT API BASE - stream: bool = optional_params.get("stream", False) or False optional_params["stream"] = stream default_api_base = create_vertex_url( vertex_location=vertex_location or "us-central1", 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 389c8446135..a93f7172ca1 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 @@ -1048,3 +1048,146 @@ class TestVertexBase: MockCredentials.from_info.assert_called_once_with(json_obj) mock_creds.with_scopes.assert_called_once_with(scopes) assert result == "scoped_creds" + + @pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"]) + @pytest.mark.asyncio + async def test_skip_auth_with_custom_api_base(self, is_async): + """ + Test that authentication is skipped when using a custom api_base + without PSC endpoint format (e.g., custom proxy that doesn't require Google credentials) + """ + vertex_base = VertexBase() + + # Test case 1: Custom api_base without PSC format should skip authentication + if is_async: + token, project = await vertex_base._ensure_access_token_async( + credentials=None, + project_id="test-project", + custom_llm_provider="vertex_ai", + api_base="https://custom-proxy.example.com", + use_psc_endpoint_format=False, + ) + else: + token, project = vertex_base._ensure_access_token( + credentials=None, + project_id="test-project", + custom_llm_provider="vertex_ai", + api_base="https://custom-proxy.example.com", + use_psc_endpoint_format=False, + ) + + # Should return empty token and the provided project_id + assert token == "" + assert project == "test-project" + + @pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"]) + @pytest.mark.asyncio + async def test_require_auth_with_psc_endpoint(self, is_async): + """ + Test that authentication is still required when using PSC endpoint format, + even with custom api_base + """ + vertex_base = VertexBase() + + # Mock credentials for PSC endpoint + mock_creds = MagicMock() + mock_creds.token = "psc-token" + mock_creds.expired = False + mock_creds.project_id = "psc-project" + mock_creds.quota_project_id = "psc-project" + + # Test case 2: Custom api_base WITH PSC format should require authentication + with patch.object( + vertex_base, "load_auth", return_value=(mock_creds, "psc-project") + ): + if is_async: + token, project = await vertex_base._ensure_access_token_async( + credentials={"type": "service_account"}, + project_id="psc-project", + custom_llm_provider="vertex_ai", + api_base="https://10.0.0.1", + use_psc_endpoint_format=True, + ) + else: + token, project = vertex_base._ensure_access_token( + credentials={"type": "service_account"}, + project_id="psc-project", + custom_llm_provider="vertex_ai", + api_base="https://10.0.0.1", + use_psc_endpoint_format=True, + ) + + # Should return actual token from authentication + assert token == "psc-token" + assert project == "psc-project" + + @pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"]) + @pytest.mark.asyncio + async def test_require_auth_without_custom_api_base(self, is_async): + """ + Test that authentication is required when no custom api_base is provided + (standard Vertex AI usage) + """ + vertex_base = VertexBase() + + # Mock credentials for standard Vertex AI + mock_creds = MagicMock() + mock_creds.token = "standard-token" + mock_creds.expired = False + mock_creds.project_id = "standard-project" + mock_creds.quota_project_id = "standard-project" + + # Test case 3: No custom api_base should require authentication + with patch.object( + vertex_base, "load_auth", return_value=(mock_creds, "standard-project") + ): + if is_async: + token, project = await vertex_base._ensure_access_token_async( + credentials={"type": "service_account"}, + project_id="standard-project", + custom_llm_provider="vertex_ai", + api_base=None, + use_psc_endpoint_format=False, + ) + else: + token, project = vertex_base._ensure_access_token( + credentials={"type": "service_account"}, + project_id="standard-project", + custom_llm_provider="vertex_ai", + api_base=None, + use_psc_endpoint_format=False, + ) + + # Should return actual token from authentication + assert token == "standard-token" + assert project == "standard-project" + + @pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"]) + @pytest.mark.asyncio + async def test_skip_auth_with_custom_api_base_no_project(self, is_async): + """ + Test that authentication is skipped with custom api_base even when project_id is None + """ + vertex_base = VertexBase() + + # Test case 4: Custom api_base without project_id should still skip auth + if is_async: + token, project = await vertex_base._ensure_access_token_async( + credentials=None, + project_id=None, + custom_llm_provider="vertex_ai", + api_base="https://custom-proxy.example.com", + use_psc_endpoint_format=False, + ) + else: + token, project = vertex_base._ensure_access_token( + credentials=None, + project_id=None, + custom_llm_provider="vertex_ai", + api_base="https://custom-proxy.example.com", + use_psc_endpoint_format=False, + ) + + # Should return empty token and empty project + assert token == "" + assert project == ""