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:
Richard Tweed
2025-08-02 08:55:35 -07:00
committed by GitHub
co-authored by openhands
parent 1e33dc50a0
commit 0ac093b59e
2 changed files with 719 additions and 16 deletions
+155 -15
View File
@@ -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'