From fde0c5e53f0d9351d396bf522274e946cbd38ddc Mon Sep 17 00:00:00 2001 From: Paul Selden Date: Mon, 19 May 2025 11:52:35 -0400 Subject: [PATCH] Allow passed in vertex_ai credentials to be authorized_user type (#10899) --- litellm/llms/vertex_ai/vertex_llm_base.py | 50 ++++++++++------ .../llms/vertex_ai/test_vertex_llm_base.py | 60 +++++++++++++++++++ 2 files changed, 93 insertions(+), 17 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index b4b04df0a90..8bc10a975bc 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -40,15 +40,7 @@ class VertexBase: def load_auth( self, credentials: Optional[VERTEX_CREDENTIALS_TYPES], project_id: Optional[str] ) -> Tuple[Any, str]: - import google.auth as google_auth - from google.auth import identity_pool - from google.auth.transport.requests import ( - Request, # type: ignore[import-untyped] - ) - if credentials is not None: - import google.oauth2.service_account - if isinstance(credentials, str): verbose_logger.debug( "Vertex: Loading vertex credentials from %s", credentials @@ -80,25 +72,31 @@ class VertexBase: # Check if the JSON object contains Workload Identity Federation configuration if "type" in json_obj and json_obj["type"] == "external_account": - creds = identity_pool.Credentials.from_info(json_obj) + creds = self._credentials_from_identity_pool(json_obj) + # Check if the JSON object contains Authorized User configuration (via gcloud auth application-default login) + elif "type" in json_obj and json_obj["type"] == "authorized_user": + creds = self._credentials_from_authorized_user( + json_obj, + scopes=["https://www.googleapis.com/auth/cloud-platform"], + ) + if project_id is None: + project_id = creds.quota_project_id # authorized user credentials don't have a project_id, only quota_project_id else: - creds = ( - google.oauth2.service_account.Credentials.from_service_account_info( - json_obj, - scopes=["https://www.googleapis.com/auth/cloud-platform"], - ) + creds = self._credentials_from_service_account( + json_obj, + scopes=["https://www.googleapis.com/auth/cloud-platform"], ) if project_id is None: project_id = getattr(creds, "project_id", None) else: - creds, creds_project_id = google_auth.default( - scopes=["https://www.googleapis.com/auth/cloud-platform"], + creds, creds_project_id = self._credentials_from_default_auth( + scopes=["https://www.googleapis.com/auth/cloud-platform"] ) if project_id is None: project_id = creds_project_id - creds.refresh(Request()) # type: ignore + self.refresh_auth(creds) if not project_id: raise ValueError("Could not resolve project_id") @@ -109,6 +107,24 @@ class VertexBase: ) return creds, project_id + + # Google Auth Helpers -- extracted for mocking purposes in tests + def _credentials_from_identity_pool(self, json_obj): + from google.auth import identity_pool + return identity_pool.Credentials.from_info(json_obj) + + def _credentials_from_authorized_user(self, json_obj, scopes): + import google.oauth2.credentials + return google.oauth2.credentials.Credentials.from_authorized_user_info(json_obj, scopes=scopes) + + def _credentials_from_service_account(self, json_obj, scopes): + import google.oauth2.service_account + return google.oauth2.service_account.Credentials.from_service_account_info(json_obj, scopes=scopes) + + def _credentials_from_default_auth(self, scopes): + import google.auth as google_auth + return google_auth.default(scopes=scopes) + def refresh_auth(self, credentials: Any) -> None: from google.auth.transport.requests import ( diff --git a/tests/litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/litellm/llms/vertex_ai/test_vertex_llm_base.py index 135dc5b616f..e1bdb31c7dd 100644 --- a/tests/litellm/llms/vertex_ai/test_vertex_llm_base.py +++ b/tests/litellm/llms/vertex_ai/test_vertex_llm_base.py @@ -174,3 +174,63 @@ class TestVertexBase: ) assert token == "" assert project == "" + + @pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"]) + @pytest.mark.asyncio + async def test_authorized_user_credentials(self, is_async): + vertex_base = VertexBase() + + quota_project_id = "test-project" + + credentials = { + "account": "", + "client_id": "fake-client-id", + "client_secret": "fake-secret", + "quota_project_id": "test-project", + "refresh_token": "fake-refresh-token", + "type": "authorized_user", + "universe_domain": "googleapis.com" + } + + mock_creds = MagicMock() + mock_creds.token = "token-1" + mock_creds.expired = False + mock_creds.quota_project_id = quota_project_id + + + with patch.object( + vertex_base, "_credentials_from_authorized_user", return_value=mock_creds + ) as mock_credentials_from_authorized_user, patch.object(vertex_base, "refresh_auth") as mock_refresh: + def mock_refresh_impl(creds): + creds.token = "refreshed-token" + + mock_refresh.side_effect = mock_refresh_impl + + # 1. Test that authorized_user-style credentials are correctly handled and uses quota_project_id + if is_async: + token, project = await vertex_base._ensure_access_token_async( + credentials=credentials, project_id=None, custom_llm_provider="vertex_ai" + ) + else: + token, project = vertex_base._ensure_access_token( + credentials=credentials, project_id=None, custom_llm_provider="vertex_ai" + ) + + assert mock_credentials_from_authorized_user.called + assert token == "refreshed-token" + assert project == quota_project_id + + # 2. Test that authorized_user-style credentials are correctly handled and uses passed in project_id + not_quota_project_id = "new-project" + if is_async: + token, project = await vertex_base._ensure_access_token_async( + credentials=credentials, project_id=not_quota_project_id, custom_llm_provider="vertex_ai" + ) + else: + token, project = vertex_base._ensure_access_token( + credentials=credentials, project_id=not_quota_project_id, custom_llm_provider="vertex_ai" + ) + + + assert token == "refreshed-token" + assert project == not_quota_project_id \ No newline at end of file