diff --git a/docs/my-website/docs/providers/vertex.md b/docs/my-website/docs/providers/vertex.md
index d919d0412cd..a3eb673f039 100644
--- a/docs/my-website/docs/providers/vertex.md
+++ b/docs/my-website/docs/providers/vertex.md
@@ -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.
+
+
+
+
+```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"
+)
+```
+
+
+
+
+```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
+```
+
+
+
+
+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:
diff --git a/litellm/llms/vertex_ai/aws_credentials_supplier.py b/litellm/llms/vertex_ai/aws_credentials_supplier.py
new file mode 100644
index 00000000000..f358511311b
--- /dev/null
+++ b/litellm/llms/vertex_ai/aws_credentials_supplier.py
@@ -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
diff --git a/litellm/llms/vertex_ai/vertex_ai_aws_wif.py b/litellm/llms/vertex_ai/vertex_ai_aws_wif.py
new file mode 100644
index 00000000000..44a0016e4ec
--- /dev/null
+++ b/litellm/llms/vertex_ai/vertex_ai_aws_wif.py
@@ -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": "" 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
diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py
index 4613b6a5715..86e14a30df4 100644
--- a/litellm/llms/vertex_ai/vertex_llm_base.py
+++ b/litellm/llms/vertex_ai/vertex_llm_base.py
@@ -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,
diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py
index 389c8446135..f9fc730e1df 100644
--- a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py
+++ b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py
@@ -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)