fix: role chaining with webauthentication for aws bedrock (#12607)

* fix(bedrock): auto-generate session name when only aws_role_name is provided

Fixes #12583 - AWS role assumption not working correctly when aws_role_name
is provided without aws_session_name.

Previously, if only aws_role_name was provided in the config without
aws_session_name, the code would fall back to using environment credentials
instead of assuming the specified role. This was problematic in EKS/IRSA
environments where users want to assume a different role.

The fix:
- When aws_role_name is provided without aws_session_name, we now
  auto-generate a session name with format 'litellm-session-{timestamp}'
- This ensures role assumption happens as expected
- Added comprehensive test coverage for this scenario

* style: format test file with black

---------

Co-authored-by: openhands <openhands@all-hands.dev>
This commit is contained in:
Richard Tweed
2025-07-15 20:17:09 -07:00
committed by GitHub
co-authored by openhands
parent 760ece5c8b
commit 197e7efa8f
2 changed files with 120 additions and 26 deletions
+4 -1
View File
@@ -178,7 +178,10 @@ class BaseAWSLLM:
aws_region_name=aws_region_name,
aws_sts_endpoint=aws_sts_endpoint,
)
elif aws_role_name is not None and aws_session_name is not None:
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,
@@ -22,6 +22,7 @@ from litellm.llms.bedrock.base_aws_llm import (
BaseAWSLLM,
Boto3CredentialsInfo,
)
from litellm.caching.caching import DualCache
# Global variable for the base_aws_llm.py file path
@@ -176,7 +177,8 @@ def test_get_aws_region_name_boto3_fallback():
assert result == "ap-southeast-1"
mock_boto3_session.assert_not_called()
def test_sign_request_with_env_var_bearer_token():
# Create instance of actual class
llm = BaseAWSLLM()
@@ -189,14 +191,14 @@ def test_sign_request_with_env_var_bearer_token():
api_base = "https://api.example.com"
# Mock environment variable
with patch.dict(os.environ, {'AWS_BEARER_TOKEN_BEDROCK': 'test_token'}):
with patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": "test_token"}):
# Execute
result_headers, result_body = llm._sign_request(
service_name=service_name,
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base
api_base=api_base,
)
# Assert
@@ -215,7 +217,7 @@ def test_sign_request_with_sigv4():
mock_request = MagicMock()
mock_request.headers = {
"Authorization": "AWS4-HMAC-SHA256 Credential=test",
"Content-Type": "application/json"
"Content-Type": "application/json",
}
mock_request.body = b'{"prompt": "test"}'
@@ -225,23 +227,25 @@ def test_sign_request_with_sigv4():
optional_params = {
"aws_access_key_id": "test_key",
"aws_secret_access_key": "test_secret",
"aws_region_name": "us-west-2"
"aws_region_name": "us-west-2",
}
request_data = {"prompt": "test"}
api_base = "https://api.example.com"
# Mock the necessary components
with patch('botocore.auth.SigV4Auth', return_value=mock_sigv4), \
patch('botocore.awsrequest.AWSRequest', return_value=mock_request), \
patch.object(llm, 'get_credentials', return_value=mock_credentials), \
patch.object(llm, '_get_aws_region_name', return_value="us-west-2"):
with patch("botocore.auth.SigV4Auth", return_value=mock_sigv4), patch(
"botocore.awsrequest.AWSRequest", return_value=mock_request
), patch.object(
llm, "get_credentials", return_value=mock_credentials
), patch.object(
llm, "_get_aws_region_name", return_value="us-west-2"
):
result_headers, result_body = llm._sign_request(
service_name=service_name,
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base
api_base=api_base,
)
# Assert
@@ -272,7 +276,7 @@ def test_sign_request_with_api_key_bearer_token():
optional_params=optional_params,
request_data=request_data,
api_base=api_base,
api_key=api_key
api_key=api_key,
)
# Assert
@@ -300,16 +304,16 @@ def test_get_request_headers_with_env_var_bearer_token():
return mock_request
# Test with bearer token
with patch.dict(os.environ, {'AWS_BEARER_TOKEN_BEDROCK': 'test_token'}), \
patch('botocore.awsrequest.AWSRequest', side_effect=mock_aws_request_init):
with patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": "test_token"}), patch(
"botocore.awsrequest.AWSRequest", side_effect=mock_aws_request_init
):
result = llm.get_request_headers(
credentials=credentials,
aws_region_name="us-west-2",
extra_headers=None,
endpoint_url="https://api.example.com",
data='{"prompt": "test"}',
headers=headers_dict
headers=headers_dict,
)
# Assert
@@ -331,17 +335,18 @@ def test_get_request_headers_with_sigv4():
mock_sigv4 = MagicMock()
# Test without bearer token (should use SigV4)
with patch.dict(os.environ, {}, clear=True), \
patch('botocore.auth.SigV4Auth', return_value=mock_sigv4) as mock_sigv4_class, \
patch('botocore.awsrequest.AWSRequest', return_value=mock_request):
with patch.dict(os.environ, {}, clear=True), patch(
"botocore.auth.SigV4Auth", return_value=mock_sigv4
) as mock_sigv4_class, patch(
"botocore.awsrequest.AWSRequest", return_value=mock_request
):
result = llm.get_request_headers(
credentials=credentials,
aws_region_name="us-west-2",
extra_headers=None,
endpoint_url="https://api.example.com",
data='{"prompt": "test"}',
headers=headers
headers=headers,
)
# Verify SigV4 authentication and result
@@ -372,9 +377,9 @@ def test_get_request_headers_with_api_key_bearer_token():
return mock_request
# Test with api_key parameter
with patch.dict(os.environ, {}, clear=True), \
patch('botocore.awsrequest.AWSRequest', side_effect=mock_aws_request_init):
with patch.dict(os.environ, {}, clear=True), patch(
"botocore.awsrequest.AWSRequest", side_effect=mock_aws_request_init
):
result = llm.get_request_headers(
credentials=credentials,
aws_region_name="us-west-2",
@@ -382,9 +387,95 @@ def test_get_request_headers_with_api_key_bearer_token():
endpoint_url="https://api.example.com",
data='{"prompt": "test"}',
headers=headers_dict,
api_key=api_key
api_key=api_key,
)
# Assert
assert mock_request.headers["Authorization"] == f"Bearer {api_key}"
assert result == mock_prepared_request
def test_role_assumption_without_session_name():
"""
Test for issue 12583: Role assumption should work when only aws_role_name is provided
without aws_session_name. The system should auto-generate a session name.
"""
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
# Test case 1: aws_role_name provided without aws_session_name
with patch("boto3.client", return_value=mock_sts_client):
credentials = base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
)
# Verify assume_role was called
mock_sts_client.assume_role.assert_called_once()
# Check the call arguments
call_args = mock_sts_client.assume_role.call_args
assert (
call_args[1]["RoleArn"]
== "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
)
# Session name should be auto-generated with format "litellm-session-{timestamp}"
assert call_args[1]["RoleSessionName"].startswith("litellm-session-")
# Verify credentials are returned correctly
assert isinstance(credentials, Credentials)
assert credentials.access_key == "assumed-access-key"
assert credentials.secret_key == "assumed-secret-key"
assert credentials.token == "assumed-session-token"
# Test case 2: Both aws_role_name and aws_session_name provided (existing behavior)
mock_sts_client.reset_mock()
with patch("boto3.client", return_value=mock_sts_client):
credentials = base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
aws_session_name="my-custom-session",
)
# Verify assume_role was called with custom session name
mock_sts_client.assume_role.assert_called_once()
call_args = mock_sts_client.assume_role.call_args
assert call_args[1]["RoleSessionName"] == "my-custom-session"
# Test case 3: Verify caching works with auto-generated session names
# Clear the cache first
base_aws_llm.iam_cache = DualCache()
mock_sts_client.reset_mock()
with patch("boto3.client", return_value=mock_sts_client):
# First call
credentials1 = base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
)
# Second call with same role should use cache (not call assume_role again)
credentials2 = base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
)
# Should only be called once due to caching
assert mock_sts_client.assume_role.call_count == 1