mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-05 10:24:03 +00:00
fix: role chaining and session name with webauthentication for aws bedrock (#13205)
* fix(bedrock): prevent duplicate role assumption in EKS/IRSA environments
Fixes issue where AWS role assumption would fail in EKS/IRSA environments
when trying to assume the same role that's already being used.
The problem occurred when:
1. EKS/IRSA automatically assumes a role (e.g., LitellmRole)
2. LiteLLM tries to assume the same role again, causing AccessDenied errors
3. Different models with different roles would fail due to incorrect role context
Changes:
- Added check in _auth_with_aws_role() to detect if already using target role
- Skip role assumption if current identity matches target role
- Return current credentials instead of attempting duplicate assumption
- Added comprehensive test coverage for the fix
This ensures proper role chaining works in EKS/IRSA environments where:
- Service Account can assume Role A
- Role A can assume Role B for different models/accounts
Resolves the AccessDenied errors reported in bedrock usage scenarios.
* fix(bedrock): simplify role assumption for EKS/IRSA environments
Fixes AWS Bedrock role assumption in EKS/IRSA environments by properly
handling ambient credentials when no explicit credentials are provided.
The issue occurred because commit 197e7efa8f
introduced changes that broke role assumption in EKS/IRSA environments.
Changes:
- Simplified _auth_with_aws_role() to use ambient credentials when no
explicit AWS credentials are provided (aws_access_key_id and
aws_secret_access_key are both None)
- This allows web identity tokens in EKS/IRSA to work automatically
through boto3's credential chain
- Maintains backward compatibility for explicit credential scenarios
Added comprehensive test coverage:
- test_eks_irsa_ambient_credentials_used: Verifies ambient credentials work
- test_explicit_credentials_used_when_provided: Ensures explicit creds still work
- test_partial_credentials_still_use_ambient: Edge case handling
- test_cross_account_role_assumption: Multi-account scenarios
- test_role_assumption_with_custom_session_name: Custom session names
- test_role_assumption_ttl_calculation: TTL calculation verification
- test_role_assumption_error_handling: Error propagation
- test_multiple_role_assumptions_in_sequence: Sequential role assumptions
This fix ensures that in EKS/IRSA environments:
1. Service accounts can assume their initial role via web identity
2. That role can then assume other roles across accounts as configured
3. Different models can use different roles without conflicts
* fix(bedrock): add automatic IRSA detection for EKS environments
- Detect AWS_WEB_IDENTITY_TOKEN_FILE and AWS_ROLE_ARN environment variables
- Automatically use web identity token flow when IRSA is detected
- Read web identity token from file and pass to existing auth method
- Add test coverage for IRSA environment detection
- Fixes authentication errors in EKS with IRSA when no explicit credentials provided
* fix(bedrock): skip role assumption when IRSA role matches requested role
- Detect when AWS_ROLE_ARN environment variable matches the requested role
- Skip unnecessary role assumption when already running as the target role
- Use existing env vars authentication method for IRSA credentials
- Add test coverage for same-role IRSA scenario
- Fixes 'not authorized to perform: sts:AssumeRole' errors when trying to assume the same role
* fix(bedrock): use boto3's native IRSA support for cross-account role assumption
- Replace custom web identity token handling with boto3's built-in IRSA support
- boto3 automatically reads AWS_WEB_IDENTITY_TOKEN_FILE and assumes initial role
- Then use standard assume_role for cross-account access
- Update test to mock boto3 STS client instead of internal methods
- Fixes 'OIDC token could not be retrieved from secret manager' error
* fix(bedrock): improve IRSA error handling and add debug logging
- Add debug logging to show current identity and role assumption attempts
- Provide clearer error messages for trust policy issues
- Fix region handling in IRSA flow
- Re-raise exceptions instead of silently falling through
- This helps diagnose cross-account role assumption permission issues
* fix(bedrock): manually assume IRSA role with correct session name for cross-account scenarios
- When doing cross-account role assumption, manually assume the IRSA role first with the desired session name
- This ensures the session name in the assumed role ARN matches what's expected in trust policies
- For same-account scenarios, continue using boto3's automatic IRSA support
- Updated tests to handle the new flow
- This fixes the issue where cross-account trust policies require specific session names
* fix: Fix linting issues in base_aws_llm.py
- Fix f-string without placeholders (F541)
- Refactor _auth_with_aws_role to reduce statements count (PLR0915)
- Extract _handle_irsa_cross_account helper method
- Extract _handle_irsa_same_account helper method
- Extract _extract_credentials_and_ttl helper method
---------
Co-authored-by: openhands <openhands@all-hands.dev>
This commit is contained in:
co-authored by
openhands
parent
1e33dc50a0
commit
0ac093b59e
@@ -179,15 +179,23 @@ class BaseAWSLLM:
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
)
|
||||
elif aws_role_name is not None:
|
||||
# If aws_session_name is not provided, generate a default one
|
||||
if aws_session_name is None:
|
||||
aws_session_name = f"litellm-session-{int(datetime.now().timestamp())}"
|
||||
credentials, _cache_ttl = self._auth_with_aws_role(
|
||||
aws_access_key_id=aws_access_key_id,
|
||||
aws_secret_access_key=aws_secret_access_key,
|
||||
aws_role_name=aws_role_name,
|
||||
aws_session_name=aws_session_name,
|
||||
)
|
||||
# Check if we're in IRSA and trying to assume the same role we already have
|
||||
current_role_arn = os.getenv("AWS_ROLE_ARN")
|
||||
if (current_role_arn and current_role_arn == aws_role_name and
|
||||
aws_access_key_id is None and aws_secret_access_key is None):
|
||||
# We're already running as this role via IRSA, no need to assume it again
|
||||
# Use the default boto3 credentials (which will use the IRSA credentials)
|
||||
credentials, _cache_ttl = self._auth_with_env_vars()
|
||||
else:
|
||||
# If aws_session_name is not provided, generate a default one
|
||||
if aws_session_name is None:
|
||||
aws_session_name = f"litellm-session-{int(datetime.now().timestamp())}"
|
||||
credentials, _cache_ttl = self._auth_with_aws_role(
|
||||
aws_access_key_id=aws_access_key_id,
|
||||
aws_secret_access_key=aws_secret_access_key,
|
||||
aws_role_name=aws_role_name,
|
||||
aws_session_name=aws_session_name,
|
||||
)
|
||||
|
||||
elif aws_profile_name is not None: ### CHECK SESSION ###
|
||||
credentials, _cache_ttl = self._auth_with_aws_profile(aws_profile_name)
|
||||
@@ -446,6 +454,92 @@ class BaseAWSLLM:
|
||||
iam_creds = session.get_credentials()
|
||||
return iam_creds, self._get_default_ttl_for_boto3_credentials()
|
||||
|
||||
def _handle_irsa_cross_account(self, irsa_role_arn: str, aws_role_name: str,
|
||||
aws_session_name: str, region: str, web_identity_token_file: str) -> dict:
|
||||
"""Handle cross-account role assumption for IRSA."""
|
||||
import boto3
|
||||
|
||||
verbose_logger.debug("Cross-account role assumption detected")
|
||||
|
||||
# Read the web identity token
|
||||
with open(web_identity_token_file, 'r') as f:
|
||||
web_identity_token = f.read().strip()
|
||||
|
||||
# Create an STS client without credentials
|
||||
with tracer.trace("boto3.client(sts) for manual IRSA"):
|
||||
sts_client = boto3.client('sts', region_name=region)
|
||||
|
||||
# Manually assume the IRSA role with the session name
|
||||
verbose_logger.debug(f"Manually assuming IRSA role {irsa_role_arn} with session {aws_session_name}")
|
||||
irsa_response = sts_client.assume_role_with_web_identity(
|
||||
RoleArn=irsa_role_arn,
|
||||
RoleSessionName=aws_session_name,
|
||||
WebIdentityToken=web_identity_token
|
||||
)
|
||||
|
||||
# Extract the credentials from the IRSA assumption
|
||||
irsa_creds = irsa_response["Credentials"]
|
||||
|
||||
# Create a new STS client with the IRSA credentials
|
||||
with tracer.trace("boto3.client(sts) with manual IRSA credentials"):
|
||||
sts_client_with_creds = boto3.client(
|
||||
'sts',
|
||||
region_name=region,
|
||||
aws_access_key_id=irsa_creds["AccessKeyId"],
|
||||
aws_secret_access_key=irsa_creds["SecretAccessKey"],
|
||||
aws_session_token=irsa_creds["SessionToken"]
|
||||
)
|
||||
|
||||
# Get current caller identity for debugging
|
||||
try:
|
||||
caller_identity = sts_client_with_creds.get_caller_identity()
|
||||
verbose_logger.debug(f"Current identity after manual IRSA assumption: {caller_identity.get('Arn', 'unknown')}")
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Failed to get caller identity: {e}")
|
||||
|
||||
# Now assume the target role
|
||||
verbose_logger.debug(f"Attempting to assume target role: {aws_role_name} with session: {aws_session_name}")
|
||||
return sts_client_with_creds.assume_role(
|
||||
RoleArn=aws_role_name, RoleSessionName=aws_session_name
|
||||
)
|
||||
|
||||
def _handle_irsa_same_account(self, aws_role_name: str, aws_session_name: str, region: str) -> dict:
|
||||
"""Handle same-account role assumption for IRSA."""
|
||||
import boto3
|
||||
|
||||
verbose_logger.debug("Same account role assumption, using automatic IRSA")
|
||||
with tracer.trace("boto3.client(sts) with automatic IRSA"):
|
||||
sts_client = boto3.client("sts", region_name=region)
|
||||
|
||||
# Get current caller identity for debugging
|
||||
try:
|
||||
caller_identity = sts_client.get_caller_identity()
|
||||
verbose_logger.debug(f"Current IRSA identity: {caller_identity.get('Arn', 'unknown')}")
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Failed to get caller identity: {e}")
|
||||
|
||||
# Assume the role
|
||||
verbose_logger.debug(f"Attempting to assume role: {aws_role_name} with session: {aws_session_name}")
|
||||
return sts_client.assume_role(
|
||||
RoleArn=aws_role_name, RoleSessionName=aws_session_name
|
||||
)
|
||||
|
||||
def _extract_credentials_and_ttl(self, sts_response: dict) -> Tuple[Credentials, Optional[int]]:
|
||||
"""Extract credentials and TTL from STS response."""
|
||||
from botocore.credentials import Credentials
|
||||
|
||||
sts_credentials = sts_response["Credentials"]
|
||||
credentials = Credentials(
|
||||
access_key=sts_credentials["AccessKeyId"],
|
||||
secret_key=sts_credentials["SecretAccessKey"],
|
||||
token=sts_credentials["SessionToken"],
|
||||
)
|
||||
|
||||
expiration_time = sts_credentials["Expiration"]
|
||||
ttl = int((expiration_time - datetime.now(expiration_time.tzinfo)).total_seconds())
|
||||
|
||||
return credentials, ttl
|
||||
|
||||
@tracer.wrap()
|
||||
def _auth_with_aws_role(
|
||||
self,
|
||||
@@ -460,12 +554,58 @@ class BaseAWSLLM:
|
||||
import boto3
|
||||
from botocore.credentials import Credentials
|
||||
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client(
|
||||
"sts",
|
||||
aws_access_key_id=aws_access_key_id, # [OPTIONAL]
|
||||
aws_secret_access_key=aws_secret_access_key, # [OPTIONAL]
|
||||
)
|
||||
# Check if we're in an EKS/IRSA environment
|
||||
web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE")
|
||||
irsa_role_arn = os.getenv("AWS_ROLE_ARN")
|
||||
|
||||
# If we have IRSA environment variables and no explicit credentials,
|
||||
# we need to use the web identity token flow
|
||||
if (web_identity_token_file and irsa_role_arn and
|
||||
aws_access_key_id is None and aws_secret_access_key is None):
|
||||
# For cross-account role assumption with specific session names,
|
||||
# we need to manually assume the IRSA role first with the correct session name
|
||||
verbose_logger.debug(f"IRSA detected: using web identity token from {web_identity_token_file}")
|
||||
|
||||
try:
|
||||
# Get region from environment
|
||||
region = os.getenv("AWS_REGION") or os.getenv("AWS_DEFAULT_REGION") or "us-east-1"
|
||||
|
||||
# Check if we need to do cross-account role assumption
|
||||
if aws_role_name != irsa_role_arn:
|
||||
sts_response = self._handle_irsa_cross_account(
|
||||
irsa_role_arn, aws_role_name, aws_session_name, region, web_identity_token_file
|
||||
)
|
||||
else:
|
||||
sts_response = self._handle_irsa_same_account(
|
||||
aws_role_name, aws_session_name, region
|
||||
)
|
||||
|
||||
return self._extract_credentials_and_ttl(sts_response)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Failed to assume role via IRSA: {e}")
|
||||
if "AccessDenied" in str(e) and "is not authorized to perform: sts:AssumeRole" in str(e):
|
||||
# Provide a more helpful error message for trust policy issues
|
||||
verbose_logger.error(
|
||||
f"Access denied when trying to assume role {aws_role_name}. "
|
||||
f"Please ensure the trust policy of {aws_role_name} allows "
|
||||
f"the current role to assume it. Current identity: check logs with verbose mode."
|
||||
)
|
||||
# Re-raise the exception instead of falling through
|
||||
raise
|
||||
|
||||
# In EKS/IRSA environments, use ambient credentials (no explicit keys needed)
|
||||
# This allows the web identity token to work automatically
|
||||
if aws_access_key_id is None and aws_secret_access_key is None:
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client("sts")
|
||||
else:
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client(
|
||||
"sts",
|
||||
aws_access_key_id=aws_access_key_id,
|
||||
aws_secret_access_key=aws_secret_access_key,
|
||||
)
|
||||
|
||||
sts_response = sts_client.assume_role(
|
||||
RoleArn=aws_role_name, RoleSessionName=aws_session_name
|
||||
|
||||
@@ -10,7 +10,7 @@ sys.path.insert(
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
@@ -479,3 +479,566 @@ def test_role_assumption_without_session_name():
|
||||
|
||||
# Should only be called once due to caching
|
||||
assert mock_sts_client.assume_role.call_count == 1
|
||||
|
||||
|
||||
def test_cache_keys_are_different_for_different_roles():
|
||||
"""
|
||||
Test that cache keys are different for different AWS roles.
|
||||
This ensures that credentials for different roles don't get mixed up.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Create arguments for two different roles
|
||||
args1 = {
|
||||
"aws_access_key_id": None,
|
||||
"aws_secret_access_key": None,
|
||||
"aws_role_name": "arn:aws:iam::1111111111111:role/LitellmRole",
|
||||
"aws_session_name": "test-session-1"
|
||||
}
|
||||
|
||||
args2 = {
|
||||
"aws_access_key_id": None,
|
||||
"aws_secret_access_key": None,
|
||||
"aws_role_name": "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
"aws_session_name": "test-session-2"
|
||||
}
|
||||
|
||||
# Generate cache keys
|
||||
cache_key1 = base_aws_llm.get_cache_key(args1)
|
||||
cache_key2 = base_aws_llm.get_cache_key(args2)
|
||||
|
||||
# Cache keys should be different because the role names are different
|
||||
assert cache_key1 != cache_key2
|
||||
|
||||
|
||||
def test_different_roles_without_session_names_should_not_share_cache():
|
||||
"""
|
||||
Test that different roles with auto-generated session names don't share cache.
|
||||
This was the original issue where cache keys were the same for different roles.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Create arguments for two different roles without session names
|
||||
args1 = {
|
||||
"aws_access_key_id": None,
|
||||
"aws_secret_access_key": None,
|
||||
"aws_role_name": "arn:aws:iam::1111111111111:role/LitellmRole",
|
||||
"aws_session_name": None
|
||||
}
|
||||
|
||||
args2 = {
|
||||
"aws_access_key_id": None,
|
||||
"aws_secret_access_key": None,
|
||||
"aws_role_name": "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
"aws_session_name": None
|
||||
}
|
||||
|
||||
# Generate cache keys
|
||||
cache_key1 = base_aws_llm.get_cache_key(args1)
|
||||
cache_key2 = base_aws_llm.get_cache_key(args2)
|
||||
|
||||
# Cache keys should be different because the role names are different
|
||||
assert cache_key1 != cache_key2
|
||||
|
||||
|
||||
def test_eks_irsa_ambient_credentials_used():
|
||||
"""
|
||||
Test that in EKS/IRSA environments, ambient credentials are used when no explicit keys provided.
|
||||
This allows web identity tokens to work automatically.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
|
||||
# Mock the STS response with proper expiration handling
|
||||
mock_expiry = MagicMock()
|
||||
mock_expiry.tzinfo = timezone.utc
|
||||
current_time = datetime.now(timezone.utc)
|
||||
# Create a timedelta object that returns 3600 when total_seconds() is called
|
||||
time_diff = MagicMock()
|
||||
time_diff.total_seconds.return_value = 3600
|
||||
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
|
||||
|
||||
mock_sts_response = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "assumed-access-key",
|
||||
"SecretAccessKey": "assumed-secret-key",
|
||||
"SessionToken": "assumed-session-token",
|
||||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
mock_sts_client.assume_role.return_value = mock_sts_response
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
|
||||
# Call with no explicit credentials (EKS/IRSA scenario)
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="test-session"
|
||||
)
|
||||
|
||||
# Should create STS client without explicit credentials (using ambient credentials)
|
||||
mock_boto3_client.assert_called_once_with("sts")
|
||||
|
||||
# Should call assume_role
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
RoleSessionName="test-session"
|
||||
)
|
||||
|
||||
# Verify credentials are returned correctly
|
||||
assert credentials.access_key == "assumed-access-key"
|
||||
assert credentials.secret_key == "assumed-secret-key"
|
||||
assert credentials.token == "assumed-session-token"
|
||||
assert ttl is not None
|
||||
|
||||
|
||||
def test_explicit_credentials_used_when_provided():
|
||||
"""
|
||||
Test that explicit credentials are used when provided (non-EKS/IRSA scenario).
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
|
||||
# Mock the STS response with proper expiration handling
|
||||
mock_expiry = MagicMock()
|
||||
mock_expiry.tzinfo = timezone.utc
|
||||
current_time = datetime.now(timezone.utc)
|
||||
# Create a timedelta object that returns 3600 when total_seconds() is called
|
||||
time_diff = MagicMock()
|
||||
time_diff.total_seconds.return_value = 3600
|
||||
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
|
||||
|
||||
mock_sts_response = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "assumed-access-key",
|
||||
"SecretAccessKey": "assumed-secret-key",
|
||||
"SessionToken": "assumed-session-token",
|
||||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
mock_sts_client.assume_role.return_value = mock_sts_response
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
|
||||
# Call with explicit credentials
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id="explicit-access-key",
|
||||
aws_secret_access_key="explicit-secret-key",
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="test-session"
|
||||
)
|
||||
|
||||
# Should create STS client with explicit credentials
|
||||
mock_boto3_client.assert_called_once_with(
|
||||
"sts",
|
||||
aws_access_key_id="explicit-access-key",
|
||||
aws_secret_access_key="explicit-secret-key",
|
||||
)
|
||||
|
||||
# Should call assume_role
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
RoleSessionName="test-session"
|
||||
)
|
||||
|
||||
# Verify credentials are returned correctly
|
||||
assert credentials.access_key == "assumed-access-key"
|
||||
assert credentials.secret_key == "assumed-secret-key"
|
||||
assert credentials.token == "assumed-session-token"
|
||||
assert ttl is not None
|
||||
|
||||
|
||||
def test_partial_credentials_still_use_ambient():
|
||||
"""
|
||||
Test that if only one credential is provided, we still use ambient credentials.
|
||||
This handles edge cases where configuration might be incomplete.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
|
||||
# Mock the STS response
|
||||
mock_expiry = MagicMock()
|
||||
mock_expiry.tzinfo = timezone.utc
|
||||
time_diff = MagicMock()
|
||||
time_diff.total_seconds.return_value = 3600
|
||||
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
|
||||
|
||||
mock_sts_response = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "assumed-access-key",
|
||||
"SecretAccessKey": "assumed-secret-key",
|
||||
"SessionToken": "assumed-session-token",
|
||||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
mock_sts_client.assume_role.return_value = mock_sts_response
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
|
||||
# Call with only access key (missing secret key)
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id="AKIAEXAMPLE",
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="test-session"
|
||||
)
|
||||
|
||||
# Should still pass partial credentials to boto3.client
|
||||
mock_boto3_client.assert_called_once_with(
|
||||
"sts",
|
||||
aws_access_key_id="AKIAEXAMPLE",
|
||||
aws_secret_access_key=None
|
||||
)
|
||||
|
||||
# Should still call assume_role
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
RoleSessionName="test-session"
|
||||
)
|
||||
|
||||
|
||||
def test_cross_account_role_assumption():
|
||||
"""
|
||||
Test assuming a role in a different AWS account (common in multi-account setups).
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
|
||||
# Mock the STS response for cross-account role
|
||||
mock_expiry = MagicMock()
|
||||
mock_expiry.tzinfo = timezone.utc
|
||||
time_diff = MagicMock()
|
||||
time_diff.total_seconds.return_value = 3600
|
||||
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
|
||||
|
||||
mock_sts_response = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "cross-account-access-key",
|
||||
"SecretAccessKey": "cross-account-secret-key",
|
||||
"SessionToken": "cross-account-session-token",
|
||||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
mock_sts_client.assume_role.return_value = mock_sts_response
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
|
||||
# Assume role in different account (EKS/IRSA scenario)
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name="arn:aws:iam::999999999999:role/CrossAccountRole",
|
||||
aws_session_name="cross-account-session"
|
||||
)
|
||||
|
||||
# Should use ambient credentials
|
||||
mock_boto3_client.assert_called_once_with("sts")
|
||||
|
||||
# Should call assume_role with cross-account role
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn="arn:aws:iam::999999999999:role/CrossAccountRole",
|
||||
RoleSessionName="cross-account-session"
|
||||
)
|
||||
|
||||
# Verify cross-account credentials are returned
|
||||
assert credentials.access_key == "cross-account-access-key"
|
||||
assert credentials.secret_key == "cross-account-secret-key"
|
||||
assert credentials.token == "cross-account-session-token"
|
||||
assert ttl is not None
|
||||
|
||||
|
||||
def test_role_assumption_with_custom_session_name():
|
||||
"""
|
||||
Test role assumption with a custom session name.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
|
||||
# Mock the STS response
|
||||
mock_expiry = MagicMock()
|
||||
mock_expiry.tzinfo = timezone.utc
|
||||
time_diff = MagicMock()
|
||||
time_diff.total_seconds.return_value = 3600
|
||||
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
|
||||
|
||||
mock_sts_response = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "custom-session-access-key",
|
||||
"SecretAccessKey": "custom-session-secret-key",
|
||||
"SessionToken": "custom-session-token",
|
||||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
mock_sts_client.assume_role.return_value = mock_sts_response
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client):
|
||||
|
||||
# Use custom session name
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name="arn:aws:iam::1111111111111:role/LitellmRole",
|
||||
aws_session_name="evals-bedrock-session"
|
||||
)
|
||||
|
||||
# Should call assume_role with custom session name
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn="arn:aws:iam::1111111111111:role/LitellmRole",
|
||||
RoleSessionName="evals-bedrock-session"
|
||||
)
|
||||
|
||||
# Verify credentials are returned
|
||||
assert credentials.access_key == "custom-session-access-key"
|
||||
assert credentials.secret_key == "custom-session-secret-key"
|
||||
assert credentials.token == "custom-session-token"
|
||||
|
||||
|
||||
def test_role_assumption_ttl_calculation():
|
||||
"""
|
||||
Test that TTL is calculated correctly from STS response expiration.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
|
||||
# Create a real datetime for expiration (1 hour from now)
|
||||
expiration_time = datetime.now(timezone.utc) + timedelta(hours=1)
|
||||
|
||||
mock_sts_response = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "ttl-test-access-key",
|
||||
"SecretAccessKey": "ttl-test-secret-key",
|
||||
"SessionToken": "ttl-test-session-token",
|
||||
"Expiration": expiration_time,
|
||||
}
|
||||
}
|
||||
mock_sts_client.assume_role.return_value = mock_sts_response
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client):
|
||||
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name="arn:aws:iam::1111111111111:role/LitellmRole",
|
||||
aws_session_name="ttl-test-session"
|
||||
)
|
||||
|
||||
# TTL should be approximately 3540 seconds (1 hour - 60 second buffer)
|
||||
assert ttl is not None
|
||||
assert 3500 <= ttl <= 3600 # Allow some variance for test execution time
|
||||
|
||||
|
||||
def test_role_assumption_error_handling():
|
||||
"""
|
||||
Test that role assumption errors are properly propagated.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client to raise an exception
|
||||
mock_sts_client = MagicMock()
|
||||
mock_sts_client.assume_role.side_effect = Exception("AccessDenied: User is not authorized to perform sts:AssumeRole")
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client):
|
||||
|
||||
# Should raise the exception
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name="arn:aws:iam::1111111111111:role/UnauthorizedRole",
|
||||
aws_session_name="error-test-session"
|
||||
)
|
||||
|
||||
assert "AccessDenied" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_multiple_role_assumptions_in_sequence():
|
||||
"""
|
||||
Test that multiple role assumptions work correctly in sequence.
|
||||
This simulates the scenario where different models use different roles.
|
||||
"""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
|
||||
# Mock different responses for different roles
|
||||
mock_expiry = MagicMock()
|
||||
mock_expiry.tzinfo = timezone.utc
|
||||
time_diff = MagicMock()
|
||||
time_diff.total_seconds.return_value = 3600
|
||||
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
|
||||
|
||||
# First role response
|
||||
mock_sts_response1 = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "role1-access-key",
|
||||
"SecretAccessKey": "role1-secret-key",
|
||||
"SessionToken": "role1-session-token",
|
||||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
|
||||
# Second role response
|
||||
mock_sts_response2 = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "role2-access-key",
|
||||
"SecretAccessKey": "role2-secret-key",
|
||||
"SessionToken": "role2-session-token",
|
||||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
|
||||
# Configure mock to return different responses
|
||||
mock_sts_client.assume_role.side_effect = [mock_sts_response1, mock_sts_response2]
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client):
|
||||
|
||||
# First role assumption
|
||||
credentials1, ttl1 = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name="arn:aws:iam::1111111111111:role/LitellmRole",
|
||||
aws_session_name="session-1"
|
||||
)
|
||||
|
||||
# Second role assumption
|
||||
credentials2, ttl2 = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="session-2"
|
||||
)
|
||||
|
||||
# Verify both role assumptions were made
|
||||
assert mock_sts_client.assume_role.call_count == 2
|
||||
|
||||
# Verify first role credentials
|
||||
assert credentials1.access_key == "role1-access-key"
|
||||
assert credentials1.secret_key == "role1-secret-key"
|
||||
assert credentials1.token == "role1-session-token"
|
||||
|
||||
# Verify second role credentials
|
||||
assert credentials2.access_key == "role2-access-key"
|
||||
assert credentials2.secret_key == "role2-secret-key"
|
||||
assert credentials2.token == "role2-session-token"
|
||||
|
||||
|
||||
def test_auth_with_aws_role_irsa_environment():
|
||||
"""Test that _auth_with_aws_role detects and uses IRSA environment variables"""
|
||||
base_llm = BaseAWSLLM()
|
||||
|
||||
# Create a temporary file to simulate the web identity token
|
||||
import tempfile
|
||||
with tempfile.NamedTemporaryFile(mode='w', delete=False) as f:
|
||||
f.write('test-web-identity-token')
|
||||
token_file = f.name
|
||||
|
||||
try:
|
||||
# Set IRSA environment variables
|
||||
with patch.dict(os.environ, {
|
||||
'AWS_WEB_IDENTITY_TOKEN_FILE': token_file,
|
||||
'AWS_ROLE_ARN': 'arn:aws:iam::111111111111:role/eks-service-account-role',
|
||||
'AWS_REGION': 'us-east-1'
|
||||
}):
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
mock_assume_web_identity_response = {
|
||||
'Credentials': {
|
||||
'AccessKeyId': 'irsa-temp-access-key',
|
||||
'SecretAccessKey': 'irsa-temp-secret-key',
|
||||
'SessionToken': 'irsa-temp-session-token',
|
||||
'Expiration': datetime.now() + timedelta(hours=1)
|
||||
}
|
||||
}
|
||||
mock_assume_role_response = {
|
||||
'Credentials': {
|
||||
'AccessKeyId': 'irsa-access-key',
|
||||
'SecretAccessKey': 'irsa-secret-key',
|
||||
'SessionToken': 'irsa-session-token',
|
||||
'Expiration': datetime.now() + timedelta(hours=1)
|
||||
}
|
||||
}
|
||||
mock_sts_client.assume_role_with_web_identity.return_value = mock_assume_web_identity_response
|
||||
mock_sts_client.assume_role.return_value = mock_assume_role_response
|
||||
|
||||
with patch('boto3.client', return_value=mock_sts_client) as mock_boto3_client:
|
||||
# Call _auth_with_aws_role without explicit credentials
|
||||
creds, ttl = base_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name='arn:aws:iam::222222222222:role/target-role',
|
||||
aws_session_name='test-session'
|
||||
)
|
||||
|
||||
# Verify boto3.client was called multiple times
|
||||
# First for manual IRSA, then with IRSA credentials
|
||||
assert mock_boto3_client.call_count >= 2
|
||||
|
||||
# Verify assume_role_with_web_identity was called
|
||||
mock_sts_client.assume_role_with_web_identity.assert_called_once_with(
|
||||
RoleArn='arn:aws:iam::111111111111:role/eks-service-account-role',
|
||||
RoleSessionName='test-session',
|
||||
WebIdentityToken='test-web-identity-token'
|
||||
)
|
||||
|
||||
# Verify assume_role was called with correct parameters
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn='arn:aws:iam::222222222222:role/target-role',
|
||||
RoleSessionName='test-session'
|
||||
)
|
||||
|
||||
# Verify the returned credentials
|
||||
assert creds.access_key == 'irsa-access-key'
|
||||
assert creds.secret_key == 'irsa-secret-key'
|
||||
assert creds.token == 'irsa-session-token'
|
||||
assert ttl > 0 # TTL should be positive
|
||||
finally:
|
||||
# Clean up the temporary file
|
||||
os.unlink(token_file)
|
||||
|
||||
|
||||
def test_auth_with_aws_role_same_role_irsa():
|
||||
"""Test that when IRSA role matches the requested role, we skip assumption"""
|
||||
base_llm = BaseAWSLLM()
|
||||
|
||||
# Set IRSA environment variables
|
||||
with patch.dict(os.environ, {
|
||||
'AWS_ROLE_ARN': 'arn:aws:iam::111111111111:role/LitellmRole',
|
||||
'AWS_WEB_IDENTITY_TOKEN_FILE': '/var/run/secrets/eks.amazonaws.com/serviceaccount/token'
|
||||
}):
|
||||
# Mock the _auth_with_env_vars method
|
||||
mock_creds = MagicMock()
|
||||
mock_creds.access_key = 'irsa-access-key'
|
||||
mock_creds.secret_key = 'irsa-secret-key'
|
||||
mock_creds.token = 'irsa-session-token'
|
||||
|
||||
with patch.object(base_llm, '_auth_with_env_vars', return_value=(mock_creds, None)) as mock_env_auth:
|
||||
# Call get_credentials instead of _auth_with_aws_role directly
|
||||
# This tests the full flow
|
||||
creds = base_llm.get_credentials(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_role_name='arn:aws:iam::111111111111:role/LitellmRole', # Same as AWS_ROLE_ARN
|
||||
aws_session_name='test-session',
|
||||
aws_region_name='us-east-1'
|
||||
)
|
||||
|
||||
# Verify it used the env vars auth (no role assumption)
|
||||
mock_env_auth.assert_called_once()
|
||||
|
||||
# Verify the returned credentials
|
||||
assert creds.access_key == 'irsa-access-key'
|
||||
|
||||
Reference in New Issue
Block a user