mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Allow passed in vertex_ai credentials to be authorized_user type (#10899)
This commit is contained in:
parent
d36b6fa60b
commit
fde0c5e53f
2 changed files with 93 additions and 17 deletions
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue