From 517eb0ee10bdad2acc105184e8d7d060f549a0f5 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 12 Nov 2025 08:46:29 +0530 Subject: [PATCH] Use safe loading of creds (#16479) --- .../llms/vertex_ai/rerank/transformation.py | 10 ++-- .../test_vertex_ai_rerank_transformation.py | 49 +++++++++++++++++++ 2 files changed, 54 insertions(+), 5 deletions(-) diff --git a/litellm/llms/vertex_ai/rerank/transformation.py b/litellm/llms/vertex_ai/rerank/transformation.py index c3cdd2b0fb6..953c6c84ea8 100644 --- a/litellm/llms/vertex_ai/rerank/transformation.py +++ b/litellm/llms/vertex_ai/rerank/transformation.py @@ -40,8 +40,8 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase): params = optional_params or {} # Get credentials to extract project ID if needed - vertex_credentials = self.get_vertex_ai_credentials(params.copy()) - vertex_project = self.get_vertex_ai_project(params.copy()) + vertex_credentials = self.safe_get_vertex_ai_credentials(params.copy()) + vertex_project = self.safe_get_vertex_ai_project(params.copy()) # Use _ensure_access_token to extract project_id from credentials # This is the same method used in vertex embeddings @@ -76,9 +76,9 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase): Validate and set up authentication for Vertex AI Discovery Engine API """ # Get credentials and project info from optional_params (which contains vertex_credentials, etc.) - litellm_params = optional_params or {} - vertex_credentials = self.get_vertex_ai_credentials(litellm_params) - vertex_project = self.get_vertex_ai_project(litellm_params) + litellm_params = optional_params.copy() if optional_params else {} + vertex_credentials = self.safe_get_vertex_ai_credentials(litellm_params) + vertex_project = self.safe_get_vertex_ai_project(litellm_params) # Get access token using the base class method access_token, project_id = self._ensure_access_token( diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py index d5e1e8b8c1c..c1de7933f95 100644 --- a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py @@ -449,3 +449,52 @@ class TestVertexAIRerankTransform: "X-Goog-User-Project": "test-project-123" } assert headers == expected_headers + + @patch('litellm.llms.vertex_ai.rerank.transformation.VertexAIRerankConfig._ensure_access_token') + def test_validate_environment_preserves_optional_params_for_get_complete_url( + self, + mock_ensure_access_token, + ): + """ + Validate that calling validate_environment does not remove vertex-specific + parameters needed later by get_complete_url. + """ + mock_ensure_access_token.return_value = ("test-access-token", "project-from-token") + + optional_params = { + "vertex_credentials": "path/to/credentials.json", + "vertex_project": "custom-project-id", + } + + # Call validate_environment first – this previously popped the values in-place + self.config.validate_environment( + headers={}, + model=self.model, + api_key=None, + optional_params=optional_params, + ) + + # Ensure the original optional_params dict still retains the vertex keys + assert optional_params["vertex_credentials"] == "path/to/credentials.json" + assert optional_params["vertex_project"] == "custom-project-id" + + # get_complete_url should still be able to access the vertex params + with patch('litellm.llms.vertex_ai.rerank.transformation.get_secret_str', return_value=None): + url = self.config.get_complete_url( + api_base=None, + model=self.model, + optional_params=optional_params, + ) + + expected_url = ( + "https://discoveryengine.googleapis.com/v1/projects/project-from-token/" + "locations/global/rankingConfigs/default_ranking_config:rank" + ) + assert url == expected_url + + # _ensure_access_token should have been called twice with the same credentials + assert mock_ensure_access_token.call_count == 2 + first_call = mock_ensure_access_token.call_args_list[0] + second_call = mock_ensure_access_token.call_args_list[1] + assert first_call.kwargs["credentials"] == "path/to/credentials.json" + assert second_call.kwargs["credentials"] == "path/to/credentials.json"