mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(vertex_ai): support explicit AWS credentials for WIF auth (#21472)
* feat(vertex_ai): support explicit AWS credentials for WIF auth The current Vertex AI AWS Workload Identity Federation implementation exclusively uses google.auth.aws.Credentials.from_info(), which requires EC2 instance metadata access to obtain AWS credentials. In environments where the metadata service is blocked for security reasons, this makes WIF unusable. Add support for explicit AWS credentials by implementing a custom AwsSecurityCredentialsSupplier (google-auth >= 2.29.0). When aws_* keys (e.g. aws_role_name, aws_region_name) are present in the WIF credential JSON, LiteLLM uses BaseAWSLLM.get_credentials() to obtain AWS creds via STS AssumeRole (or any other supported AWS auth flow), wraps them in the custom supplier, and passes them to aws.Credentials() — bypassing the metadata service entirely. When no aws_* keys are present, the existing from_info() flow is used unchanged, preserving full backward compatibility. * refactor(vertex_ai): extract AWS WIF auth to own class + add docs Address PR review feedback: - Move _AWS_CREDENTIAL_KEYS, _extract_aws_params(), and _credentials_from_aws_with_explicit_auth() from VertexBase into new VertexAIAwsWifAuth class in vertex_ai_aws_wif.py - Add documentation for explicit AWS credentials WIF auth method in vertex.md (supported params, JSON example, SDK/Proxy tabs) Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * fix(vertex_ai): use lazy credentials provider to prevent stale STS tokens --------- Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
8b0375f99c
commit
5e34fdce77
5 changed files with 569 additions and 3 deletions
|
|
@ -1472,6 +1472,82 @@ Your WIF credentials JSON file typically looks like this (for AWS federation):
|
|||
|
||||
For more details on setting up Workload Identity Federation, see [Google Cloud WIF documentation](https://cloud.google.com/iam/docs/workload-identity-federation).
|
||||
|
||||
#### Explicit AWS Credentials for WIF
|
||||
|
||||
By default, AWS-based WIF relies on the EC2 instance metadata service to obtain AWS credentials. This works when LiteLLM runs on an EC2 instance or ECS task with an IAM role attached.
|
||||
|
||||
If your environment **does not have access to the EC2 metadata service** (e.g., running on-premises, in a container without host networking, or in a different cloud with security restrictions), you can provide explicit AWS credentials directly in the WIF credential JSON file. LiteLLM will use these to authenticate to AWS before performing the GCP token exchange.
|
||||
|
||||
Add the `aws_*` keys at the **top level** of your WIF credential JSON (alongside `type`, `audience`, etc.):
|
||||
|
||||
```json
|
||||
{
|
||||
"type": "external_account",
|
||||
"audience": "//iam.googleapis.com/projects/PROJECT_NUMBER/locations/global/workloadIdentityPools/POOL_ID/providers/PROVIDER_ID",
|
||||
"subject_token_type": "urn:ietf:params:aws:token-type:aws4_request",
|
||||
"service_account_impersonation_url": "https://iamcredentials.googleapis.com/v1/projects/-/serviceAccounts/SERVICE_ACCOUNT_EMAIL:generateAccessToken",
|
||||
"token_url": "https://sts.googleapis.com/v1/token",
|
||||
"credential_source": {
|
||||
"environment_id": "aws1",
|
||||
"region_url": "http://169.254.169.254/latest/meta-data/placement/availability-zone",
|
||||
"url": "http://169.254.169.254/latest/meta-data/iam/security-credentials",
|
||||
"regional_cred_verification_url": "https://sts.{region}.amazonaws.com?Action=GetCallerIdentity&Version=2011-06-15"
|
||||
},
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyWifRole",
|
||||
"aws_region_name": "us-east-1"
|
||||
}
|
||||
```
|
||||
|
||||
**Supported `aws_*` parameters:**
|
||||
|
||||
| Parameter | Required | Description |
|
||||
|---|---|---|
|
||||
| `aws_region_name` | Yes | AWS region for credential verification (e.g. `us-east-1`) |
|
||||
| `aws_role_name` | No | IAM role ARN for STS AssumeRole |
|
||||
| `aws_access_key_id` | No | Static AWS access key ID |
|
||||
| `aws_secret_access_key` | No | Static AWS secret access key |
|
||||
| `aws_session_token` | No | Temporary session token |
|
||||
| `aws_profile_name` | No | AWS CLI profile name |
|
||||
| `aws_session_name` | No | Session name for AssumeRole |
|
||||
| `aws_web_identity_token` | No | Web identity token for STS |
|
||||
| `aws_sts_endpoint` | No | Custom STS endpoint URL |
|
||||
| `aws_external_id` | No | External ID for cross-account AssumeRole |
|
||||
|
||||
`aws_region_name` is always required when using explicit AWS credentials. The other parameters follow the same authentication flows as [Bedrock AWS auth](/docs/providers/bedrock#authentication) -- you can use role assumption, static keys, profiles, or web identity tokens.
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/gemini-1.5-pro",
|
||||
messages=[{"role": "user", "content": "Hello!"}],
|
||||
vertex_credentials="/path/to/wif-credentials-with-aws.json", # WIF JSON with aws_* keys
|
||||
vertex_project="your-gcp-project-id",
|
||||
vertex_location="us-central1"
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gemini-model
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-1.5-pro
|
||||
vertex_project: your-gcp-project-id
|
||||
vertex_location: us-central1
|
||||
vertex_credentials: /path/to/wif-credentials-with-aws.json # WIF JSON with aws_* keys
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
When `aws_*` keys are present in the JSON, LiteLLM automatically uses explicit AWS authentication instead of the EC2 metadata service. When they are absent, the standard metadata-based flow is used unchanged.
|
||||
|
||||
### **Environment Variables**
|
||||
|
||||
You can set:
|
||||
|
|
|
|||
52
litellm/llms/vertex_ai/aws_credentials_supplier.py
Normal file
52
litellm/llms/vertex_ai/aws_credentials_supplier.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
"""
|
||||
Custom AWS Security Credentials Supplier for Vertex AI WIF.
|
||||
|
||||
Wraps boto3/botocore credentials so that google-auth can use them
|
||||
for the AWS-to-GCP Workload Identity Federation token exchange
|
||||
without hitting the EC2 instance metadata service.
|
||||
|
||||
Requires google-auth >= 2.29.0.
|
||||
"""
|
||||
|
||||
from typing import Callable
|
||||
|
||||
from google.auth import aws
|
||||
|
||||
|
||||
class AwsCredentialsSupplier(aws.AwsSecurityCredentialsSupplier):
|
||||
"""
|
||||
Supplies AWS credentials to google-auth's aws.Credentials for WIF
|
||||
token exchange.
|
||||
|
||||
This bypasses the default metadata-based credential retrieval,
|
||||
allowing WIF to work in environments where EC2 metadata is blocked.
|
||||
|
||||
Accepts a credentials_provider callable that is invoked on every
|
||||
get_aws_security_credentials() call, so that refreshed/rotated
|
||||
credentials are picked up automatically (important for temporary
|
||||
STS tokens).
|
||||
"""
|
||||
|
||||
def __init__(self, credentials_provider: Callable, aws_region: str):
|
||||
"""
|
||||
Args:
|
||||
credentials_provider: A zero-arg callable that returns a
|
||||
botocore.credentials.Credentials object (with access_key,
|
||||
secret_key, and token attributes).
|
||||
aws_region: The AWS region string (e.g. "us-east-1").
|
||||
"""
|
||||
self._credentials_provider = credentials_provider
|
||||
self._region = aws_region
|
||||
|
||||
def get_aws_security_credentials(self, context, request):
|
||||
"""Return current AWS credentials for the GCP token exchange."""
|
||||
current = self._credentials_provider()
|
||||
return aws.AwsSecurityCredentials(
|
||||
access_key_id=current.access_key,
|
||||
secret_access_key=current.secret_key,
|
||||
session_token=current.token,
|
||||
)
|
||||
|
||||
def get_aws_region(self, context, request):
|
||||
"""Return the AWS region for credential verification."""
|
||||
return self._region
|
||||
125
litellm/llms/vertex_ai/vertex_ai_aws_wif.py
Normal file
125
litellm/llms/vertex_ai/vertex_ai_aws_wif.py
Normal file
|
|
@ -0,0 +1,125 @@
|
|||
"""
|
||||
AWS Workload Identity Federation (WIF) auth for Vertex AI.
|
||||
|
||||
Handles explicit AWS credentials for GCP WIF token exchange,
|
||||
bypassing the EC2 instance metadata service.
|
||||
|
||||
When aws_* keys are present in the WIF credential JSON, this module
|
||||
uses BaseAWSLLM to obtain AWS credentials and wraps them in a custom
|
||||
AwsSecurityCredentialsSupplier for google-auth.
|
||||
"""
|
||||
|
||||
from typing import Dict
|
||||
|
||||
GOOGLE_IMPORT_ERROR_MESSAGE = (
|
||||
"Google Cloud SDK not found. Install it with: pip install 'litellm[google]' "
|
||||
"or pip install google-cloud-aiplatform"
|
||||
)
|
||||
|
||||
# AWS params recognized in WIF credential JSON for explicit auth.
|
||||
# These match the kwargs accepted by BaseAWSLLM.get_credentials().
|
||||
_AWS_CREDENTIAL_KEYS = frozenset({
|
||||
"aws_access_key_id",
|
||||
"aws_secret_access_key",
|
||||
"aws_session_token",
|
||||
"aws_region_name",
|
||||
"aws_session_name",
|
||||
"aws_profile_name",
|
||||
"aws_role_name",
|
||||
"aws_web_identity_token",
|
||||
"aws_sts_endpoint",
|
||||
"aws_external_id",
|
||||
})
|
||||
|
||||
|
||||
class VertexAIAwsWifAuth:
|
||||
"""
|
||||
Handles AWS-to-GCP Workload Identity Federation credential creation
|
||||
for Vertex AI, using explicit AWS credentials rather than EC2 metadata.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def extract_aws_params(json_obj: dict) -> Dict[str, str]:
|
||||
"""
|
||||
Extract LiteLLM-specific aws_* keys from a WIF credential JSON dict.
|
||||
|
||||
Returns a dict of {param_name: value} for any recognized aws_* keys
|
||||
found in the JSON. Returns empty dict if none are present.
|
||||
"""
|
||||
return {
|
||||
key: json_obj[key]
|
||||
for key in _AWS_CREDENTIAL_KEYS
|
||||
if key in json_obj
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def credentials_from_explicit_aws(json_obj, aws_params, scopes):
|
||||
"""
|
||||
Create GCP credentials using explicit AWS credentials for WIF.
|
||||
|
||||
Uses BaseAWSLLM to obtain AWS credentials (via STS AssumeRole, profile,
|
||||
static keys, etc.), then wraps them in a custom AwsSecurityCredentialsSupplier
|
||||
so that google-auth bypasses the EC2 metadata service.
|
||||
|
||||
Args:
|
||||
json_obj: The WIF credential JSON dict (contains audience, token_url, etc.)
|
||||
aws_params: Dict of aws_* params extracted from json_obj
|
||||
scopes: OAuth scopes for the GCP credentials
|
||||
"""
|
||||
try:
|
||||
from google.auth import aws
|
||||
except ImportError:
|
||||
raise ImportError(GOOGLE_IMPORT_ERROR_MESSAGE)
|
||||
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.llms.vertex_ai.aws_credentials_supplier import (
|
||||
AwsCredentialsSupplier,
|
||||
)
|
||||
|
||||
# Validate region first — required for the GCP token exchange.
|
||||
# Check before get_credentials() to avoid unnecessary AWS API calls
|
||||
# (e.g. STS AssumeRole) on misconfiguration.
|
||||
aws_region = aws_params.get("aws_region_name")
|
||||
if not aws_region:
|
||||
raise ValueError(
|
||||
"aws_region_name is required in the WIF credential JSON "
|
||||
"when using explicit AWS authentication. Add "
|
||||
'"aws_region_name": "<your-region>" to your credential file.'
|
||||
)
|
||||
|
||||
# Build a credentials provider that re-resolves AWS creds on each call.
|
||||
# This ensures rotated/refreshed STS tokens are picked up during
|
||||
# long-running processes when google-auth refreshes the GCP token.
|
||||
base_aws = BaseAWSLLM()
|
||||
aws_params_copy = dict(aws_params) # avoid mutating caller's dict
|
||||
|
||||
def _get_aws_credentials():
|
||||
return base_aws.get_credentials(**aws_params_copy)
|
||||
|
||||
# Create the custom supplier with a lazy credentials provider
|
||||
supplier = AwsCredentialsSupplier(
|
||||
credentials_provider=_get_aws_credentials,
|
||||
aws_region=aws_region,
|
||||
)
|
||||
|
||||
# Build kwargs for aws.Credentials — forward optional fields from JSON
|
||||
creds_kwargs = dict(
|
||||
audience=json_obj.get("audience"),
|
||||
subject_token_type=json_obj.get("subject_token_type"),
|
||||
token_url=json_obj.get("token_url"),
|
||||
credential_source=None, # Not using metadata endpoints
|
||||
aws_security_credentials_supplier=supplier,
|
||||
service_account_impersonation_url=json_obj.get(
|
||||
"service_account_impersonation_url"
|
||||
),
|
||||
)
|
||||
# Forward universe_domain if present (defaults to googleapis.com)
|
||||
if "universe_domain" in json_obj:
|
||||
creds_kwargs["universe_domain"] = json_obj["universe_domain"]
|
||||
|
||||
creds = aws.Credentials(**creds_kwargs)
|
||||
|
||||
if scopes and hasattr(creds, "requires_scopes") and creds.requires_scopes:
|
||||
creds = creds.with_scopes(scopes)
|
||||
|
||||
return creds
|
||||
|
|
@ -96,10 +96,23 @@ class VertexBase:
|
|||
else ""
|
||||
)
|
||||
if isinstance(environment_id, str) and "aws" in environment_id:
|
||||
creds = self._credentials_from_identity_pool_with_aws(
|
||||
json_obj,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"],
|
||||
# Check if explicit AWS params are in the JSON (bypasses metadata)
|
||||
from litellm.llms.vertex_ai.vertex_ai_aws_wif import (
|
||||
VertexAIAwsWifAuth,
|
||||
)
|
||||
|
||||
aws_params = VertexAIAwsWifAuth.extract_aws_params(json_obj)
|
||||
if aws_params:
|
||||
creds = VertexAIAwsWifAuth.credentials_from_explicit_aws(
|
||||
json_obj,
|
||||
aws_params=aws_params,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"],
|
||||
)
|
||||
else:
|
||||
creds = self._credentials_from_identity_pool_with_aws(
|
||||
json_obj,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"],
|
||||
)
|
||||
else:
|
||||
creds = self._credentials_from_identity_pool(
|
||||
json_obj,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ sys.path.insert(
|
|||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
from litellm.llms.vertex_ai.vertex_ai_aws_wif import VertexAIAwsWifAuth
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
|
||||
|
||||
|
|
@ -1048,3 +1049,302 @@ class TestVertexBase:
|
|||
MockCredentials.from_info.assert_called_once_with(json_obj)
|
||||
mock_creds.with_scopes.assert_called_once_with(scopes)
|
||||
assert result == "scoped_creds"
|
||||
|
||||
def test_extract_aws_params(self):
|
||||
"""Test _extract_aws_params: extraction, empty case, and unrecognized keys."""
|
||||
# Case 1: Extracts recognized aws_* keys, ignores GCP-standard fields
|
||||
json_with_role = {
|
||||
"type": "external_account",
|
||||
"audience": "//iam.googleapis.com/...",
|
||||
"token_url": "https://sts.googleapis.com/v1/token",
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
result = VertexAIAwsWifAuth.extract_aws_params(json_with_role)
|
||||
assert result == {
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
|
||||
# Case 2: Returns empty dict for standard WIF JSON (no aws_* keys)
|
||||
json_standard = {
|
||||
"type": "external_account",
|
||||
"audience": "//iam.googleapis.com/...",
|
||||
"credential_source": {"environment_id": "aws1"},
|
||||
}
|
||||
assert VertexAIAwsWifAuth.extract_aws_params(json_standard) == {}
|
||||
|
||||
# Case 3: Ignores unrecognized aws_* keys (e.g. aws_bedrock_runtime_endpoint)
|
||||
json_with_unknown = {
|
||||
"type": "external_account",
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
"aws_region_name": "us-east-1",
|
||||
"aws_unknown_field": "should-be-ignored",
|
||||
"aws_bedrock_runtime_endpoint": "should-also-be-ignored",
|
||||
}
|
||||
result = VertexAIAwsWifAuth.extract_aws_params(json_with_unknown)
|
||||
assert result == {
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
|
||||
def test_credentials_from_aws_with_explicit_auth(self):
|
||||
"""Test that explicit AWS auth creates credentials via supplier, not metadata."""
|
||||
json_obj = {
|
||||
"type": "external_account",
|
||||
"audience": "//iam.googleapis.com/projects/123/locations/global/workloadIdentityPools/pool/providers/aws",
|
||||
"subject_token_type": "urn:ietf:params:aws:token-type:aws4_request",
|
||||
"token_url": "https://sts.googleapis.com/v1/token",
|
||||
"service_account_impersonation_url": "https://iamcredentials.googleapis.com/v1/projects/-/serviceAccounts/sa@proj.iam.gserviceaccount.com:generateAccessToken",
|
||||
"credential_source": {"environment_id": "aws1"},
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
aws_params = {
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
scopes = ["https://www.googleapis.com/auth/cloud-platform"]
|
||||
|
||||
# Mock BaseAWSLLM.get_credentials to return fake boto3 credentials
|
||||
mock_boto3_creds = MagicMock()
|
||||
mock_boto3_creds.access_key = "AKIAIOSFODNN7EXAMPLE"
|
||||
mock_boto3_creds.secret_key = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
|
||||
mock_boto3_creds.token = "FwoGZXIvYXdzEBYaDHqa0AP"
|
||||
|
||||
# Mock aws.Credentials constructor
|
||||
mock_gcp_creds = MagicMock()
|
||||
mock_gcp_creds.requires_scopes = True
|
||||
mock_gcp_creds.with_scopes.return_value = mock_gcp_creds
|
||||
|
||||
# IMPORTANT: Patch at the SOURCE modules, not at vertex_llm_base level.
|
||||
# The imports happen inside the function via `from X import Y`, so
|
||||
# the mock must replace the class in its defining module.
|
||||
with patch(
|
||||
"litellm.llms.bedrock.base_aws_llm.BaseAWSLLM"
|
||||
) as MockBaseAWSLLM, patch(
|
||||
"google.auth.aws.Credentials",
|
||||
) as MockAwsCredentials:
|
||||
mock_base_aws = MagicMock()
|
||||
mock_base_aws.get_credentials.return_value = mock_boto3_creds
|
||||
MockBaseAWSLLM.return_value = mock_base_aws
|
||||
MockAwsCredentials.return_value = mock_gcp_creds
|
||||
|
||||
result = VertexAIAwsWifAuth.credentials_from_explicit_aws(
|
||||
json_obj, aws_params, scopes
|
||||
)
|
||||
|
||||
# Verify aws.Credentials was called with supplier (not from_info)
|
||||
MockAwsCredentials.assert_called_once()
|
||||
call_kwargs = MockAwsCredentials.call_args[1]
|
||||
assert call_kwargs["audience"] == json_obj["audience"]
|
||||
assert call_kwargs["subject_token_type"] == json_obj["subject_token_type"]
|
||||
assert call_kwargs["token_url"] == json_obj["token_url"]
|
||||
assert call_kwargs["credential_source"] is None
|
||||
assert call_kwargs["service_account_impersonation_url"] == json_obj["service_account_impersonation_url"]
|
||||
|
||||
# Verify the supplier is a lazy credentials provider (calls
|
||||
# get_credentials on demand, not at construction time)
|
||||
supplier = call_kwargs["aws_security_credentials_supplier"]
|
||||
assert supplier is not None
|
||||
# Trigger the lazy provider — this should call get_credentials
|
||||
supplier.get_aws_security_credentials(context=None, request=None)
|
||||
mock_base_aws.get_credentials.assert_called_once_with(**aws_params)
|
||||
|
||||
# Verify scopes were applied
|
||||
mock_gcp_creds.with_scopes.assert_called_once_with(scopes)
|
||||
assert result == mock_gcp_creds
|
||||
|
||||
def test_credentials_from_aws_with_explicit_auth_requires_region(self):
|
||||
"""Test that explicit AWS auth raises ValueError when region is missing."""
|
||||
json_obj = {
|
||||
"type": "external_account",
|
||||
"audience": "//iam.googleapis.com/...",
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
}
|
||||
aws_params = {
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
# No aws_region_name — should fail
|
||||
}
|
||||
scopes = ["https://www.googleapis.com/auth/cloud-platform"]
|
||||
|
||||
with pytest.raises(ValueError, match="aws_region_name is required"):
|
||||
VertexAIAwsWifAuth.credentials_from_explicit_aws(
|
||||
json_obj, aws_params, scopes
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_aws_wif_routes_to_explicit_auth_when_aws_params_present(
|
||||
self, is_async
|
||||
):
|
||||
"""Test that load_auth routes to explicit auth when aws_* keys are in JSON."""
|
||||
vertex_base = VertexBase()
|
||||
|
||||
credentials = {
|
||||
"type": "external_account",
|
||||
"credential_source": {"environment_id": "aws1"},
|
||||
"audience": "//iam.googleapis.com/...",
|
||||
"subject_token_type": "urn:ietf:params:aws:token-type:aws4_request",
|
||||
"token_url": "https://sts.googleapis.com/v1/token",
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
|
||||
mock_creds = MagicMock()
|
||||
mock_creds.token = "explicit-auth-token"
|
||||
mock_creds.expired = False
|
||||
mock_creds.project_id = "test-project"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.vertex_ai_aws_wif.VertexAIAwsWifAuth.credentials_from_explicit_aws",
|
||||
return_value=mock_creds,
|
||||
) as mock_explicit_auth, patch.object(
|
||||
vertex_base,
|
||||
"_credentials_from_identity_pool_with_aws",
|
||||
) as mock_metadata_auth, 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
|
||||
|
||||
if is_async:
|
||||
token, _ = await vertex_base._ensure_access_token_async(
|
||||
credentials=credentials,
|
||||
project_id=None,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
else:
|
||||
token, _ = vertex_base._ensure_access_token(
|
||||
credentials=credentials,
|
||||
project_id=None,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
# Explicit auth should be called, NOT metadata auth
|
||||
assert mock_explicit_auth.called
|
||||
mock_metadata_auth.assert_not_called()
|
||||
# Verify correct kwargs were passed to explicit auth
|
||||
call_kwargs = mock_explicit_auth.call_args[1]
|
||||
assert call_kwargs["aws_params"] == {
|
||||
"aws_role_name": "arn:aws:iam::123456789012:role/MyRole",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
assert call_kwargs["scopes"] == ["https://www.googleapis.com/auth/cloud-platform"]
|
||||
assert token == "refreshed-token"
|
||||
|
||||
@pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_aws_wif_falls_back_to_metadata_when_no_aws_params(self, is_async):
|
||||
"""Test that load_auth falls back to metadata flow when no aws_* keys in JSON."""
|
||||
vertex_base = VertexBase()
|
||||
|
||||
# Standard WIF JSON — no aws_* keys
|
||||
credentials = {
|
||||
"type": "external_account",
|
||||
"credential_source": {"environment_id": "aws1"},
|
||||
"audience": "//iam.googleapis.com/...",
|
||||
"token_url": "https://sts.googleapis.com/v1/token",
|
||||
}
|
||||
|
||||
mock_creds = MagicMock()
|
||||
mock_creds.token = "metadata-token"
|
||||
mock_creds.expired = False
|
||||
mock_creds.project_id = "test-project"
|
||||
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.vertex_ai_aws_wif.VertexAIAwsWifAuth.credentials_from_explicit_aws",
|
||||
) as mock_explicit_auth, patch.object(
|
||||
vertex_base,
|
||||
"_credentials_from_identity_pool_with_aws",
|
||||
return_value=mock_creds,
|
||||
) as mock_metadata_auth, 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
|
||||
|
||||
if is_async:
|
||||
token, _ = await vertex_base._ensure_access_token_async(
|
||||
credentials=credentials,
|
||||
project_id=None,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
else:
|
||||
token, _ = vertex_base._ensure_access_token(
|
||||
credentials=credentials,
|
||||
project_id=None,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
# Metadata auth should be called, NOT explicit auth
|
||||
mock_explicit_auth.assert_not_called()
|
||||
assert mock_metadata_auth.called
|
||||
assert token == "refreshed-token"
|
||||
|
||||
def test_aws_credentials_supplier(self):
|
||||
"""Test AwsCredentialsSupplier: wraps credentials provider, handles token=None."""
|
||||
from litellm.llms.vertex_ai.aws_credentials_supplier import (
|
||||
AwsCredentialsSupplier,
|
||||
)
|
||||
|
||||
# Case 1: With session token (STS temporary credentials)
|
||||
mock_boto3_creds = MagicMock()
|
||||
mock_boto3_creds.access_key = "AKIAIOSFODNN7EXAMPLE"
|
||||
mock_boto3_creds.secret_key = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
|
||||
mock_boto3_creds.token = "FwoGZXIvYXdzEBYaDHqa0AP"
|
||||
|
||||
supplier = AwsCredentialsSupplier(
|
||||
credentials_provider=lambda: mock_boto3_creds,
|
||||
aws_region="us-east-1",
|
||||
)
|
||||
|
||||
aws_creds = supplier.get_aws_security_credentials(context=None, request=None)
|
||||
assert aws_creds.access_key_id == "AKIAIOSFODNN7EXAMPLE"
|
||||
assert aws_creds.secret_access_key == "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
|
||||
assert aws_creds.session_token == "FwoGZXIvYXdzEBYaDHqa0AP"
|
||||
assert supplier.get_aws_region(context=None, request=None) == "us-east-1"
|
||||
|
||||
# Case 2: Without session token (static IAM credentials)
|
||||
mock_static_creds = MagicMock()
|
||||
mock_static_creds.access_key = "AKIAIOSFODNN7EXAMPLE"
|
||||
mock_static_creds.secret_key = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
|
||||
mock_static_creds.token = None
|
||||
|
||||
supplier_static = AwsCredentialsSupplier(
|
||||
credentials_provider=lambda: mock_static_creds,
|
||||
aws_region="eu-west-1",
|
||||
)
|
||||
|
||||
aws_creds_static = supplier_static.get_aws_security_credentials(
|
||||
context=None, request=None
|
||||
)
|
||||
assert aws_creds_static.access_key_id == "AKIAIOSFODNN7EXAMPLE"
|
||||
assert aws_creds_static.session_token is None
|
||||
|
||||
def test_aws_credentials_supplier_returns_correct_type(self):
|
||||
"""Test that AwsCredentialsSupplier returns AwsSecurityCredentials dataclass."""
|
||||
from google.auth.aws import AwsSecurityCredentials
|
||||
|
||||
from litellm.llms.vertex_ai.aws_credentials_supplier import (
|
||||
AwsCredentialsSupplier,
|
||||
)
|
||||
|
||||
mock_boto3_creds = MagicMock()
|
||||
mock_boto3_creds.access_key = "AKID"
|
||||
mock_boto3_creds.secret_key = "SECRET"
|
||||
mock_boto3_creds.token = "TOKEN"
|
||||
|
||||
supplier = AwsCredentialsSupplier(
|
||||
credentials_provider=lambda: mock_boto3_creds,
|
||||
aws_region="us-east-1",
|
||||
)
|
||||
|
||||
aws_creds = supplier.get_aws_security_credentials(context=None, request=None)
|
||||
assert isinstance(aws_creds, AwsSecurityCredentials)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue