Allow passed in vertex_ai credentials to be authorized_user type (#10899)

This commit is contained in:
Paul Selden 2025-05-19 11:52:35 -04:00 • committed by GitHub
parent d36b6fa60b
commit fde0c5e53f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 93 additions and 17 deletions

View file

@ -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 (

View file

@ -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