consistently passing region and endpoint args into explicit credentials irsa

This commit is contained in:
An Tang 2026-02-20 10:42:14 -08:00
parent 782df9283f
commit be885ad236
2 changed files with 63 additions and 41 deletions

View file

@ -735,6 +735,7 @@ class BaseAWSLLM:
region: str,
web_identity_token_file: str,
aws_external_id: Optional[str] = None,
aws_sts_endpoint: Optional[str] = None,
ssl_verify: Optional[Union[bool, str]] = None,
) -> dict:
"""Handle cross-account role assumption for IRSA."""
@ -746,11 +747,13 @@ class BaseAWSLLM:
with open(web_identity_token_file, "r") as f:
web_identity_token = f.read().strip()
irsa_sts_kwargs: dict = {"region_name": region, "verify": self._get_ssl_verify(ssl_verify)}
if aws_sts_endpoint is not None:
irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint
# Create an STS client without credentials
with tracer.trace("boto3.client(sts) for manual IRSA"):
sts_client = boto3.client(
"sts", region_name=region, verify=self._get_ssl_verify(ssl_verify)
)
sts_client = boto3.client("sts", **irsa_sts_kwargs)
# Manually assume the IRSA role with the session name
verbose_logger.debug(
@ -769,11 +772,10 @@ class BaseAWSLLM:
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"],
verify=self._get_ssl_verify(ssl_verify),
**irsa_sts_kwargs,
)
# Get current caller identity for debugging
@ -806,16 +808,19 @@ class BaseAWSLLM:
aws_session_name: str,
region: str,
aws_external_id: Optional[str] = None,
aws_sts_endpoint: Optional[str] = None,
ssl_verify: Optional[Union[bool, str]] = None,
) -> dict:
"""Handle same-account role assumption for IRSA."""
import boto3
irsa_sts_kwargs: dict = {"region_name": region, "verify": self._get_ssl_verify(ssl_verify)}
if aws_sts_endpoint is not None:
irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint
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, verify=self._get_ssl_verify(ssl_verify)
)
sts_client = boto3.client("sts", **irsa_sts_kwargs)
# Get current caller identity for debugging
try:
@ -884,6 +889,8 @@ class BaseAWSLLM:
web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE")
irsa_role_arn = os.getenv("AWS_ROLE_ARN")
region = aws_region_name or os.getenv("AWS_REGION") or os.getenv("AWS_DEFAULT_REGION")
# If we have IRSA environment variables and no explicit credentials,
# we need to use the web identity token flow
if (
@ -899,12 +906,8 @@ class BaseAWSLLM:
)
try:
# Get region from environment
region = (
os.getenv("AWS_REGION")
or os.getenv("AWS_DEFAULT_REGION")
or "us-east-1"
)
# Use passed-in region when set, else env, else default (align with AssumeRole path)
region = region or "us-east-1"
# Check if we need to do cross-account role assumption
if aws_role_name != irsa_role_arn:
@ -915,6 +918,7 @@ class BaseAWSLLM:
region,
web_identity_token_file,
aws_external_id,
aws_sts_endpoint=aws_sts_endpoint,
ssl_verify=ssl_verify,
)
else:
@ -923,6 +927,7 @@ class BaseAWSLLM:
aws_session_name,
region,
aws_external_id,
aws_sts_endpoint=aws_sts_endpoint,
ssl_verify=ssl_verify,
)
@ -945,8 +950,8 @@ class BaseAWSLLM:
# In EKS/IRSA environments, use ambient credentials (no explicit keys needed)
# This allows the web identity token to work automatically
sts_client_kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)}
if aws_region_name is not None:
sts_client_kwargs["region_name"] = aws_region_name
if region is not None:
sts_client_kwargs["region_name"] = region
if aws_sts_endpoint is not None:
sts_client_kwargs["endpoint_url"] = aws_sts_endpoint
if aws_access_key_id is None and aws_secret_access_key is None:

View file

@ -593,24 +593,54 @@ def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs):
assert ttl is not None
def test_explicit_credentials_used_when_provided():
@pytest.mark.parametrize(
"role_kwargs,expected_client_kwargs",
[
(
{},
{
"aws_access_key_id": "explicit-access-key",
"aws_secret_access_key": "explicit-secret-key",
"aws_session_token": "assumed-session-token",
"verify": True,
},
),
(
{"aws_region_name": "us-east-1"},
{
"region_name": "us-east-1",
"aws_access_key_id": "explicit-access-key",
"aws_secret_access_key": "explicit-secret-key",
"aws_session_token": "assumed-session-token",
"verify": True,
},
),
(
{"aws_sts_endpoint": "https://sts.eu-west-1.amazonaws.com"},
{
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
"aws_access_key_id": "explicit-access-key",
"aws_secret_access_key": "explicit-secret-key",
"aws_session_token": "assumed-session-token",
"verify": True,
},
),
],
ids=["no_region_or_endpoint", "regional_sts", "explicit_sts_endpoint"],
)
def test_explicit_credentials_used_when_provided(role_kwargs, expected_client_kwargs):
"""
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",
@ -619,36 +649,23 @@ def test_explicit_credentials_used_when_provided():
"Expiration": mock_expiry,
}
}
mock_sts_client = MagicMock()
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_session_token="assumed-session-token",
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
aws_session_name="test-session"
aws_session_name="test-session",
**role_kwargs
)
# Should create STS client with explicit credentials
# Note: verify parameter is passed for SSL verification
mock_boto3_client.assert_called_once_with(
"sts",
aws_access_key_id="explicit-access-key",
aws_secret_access_key="explicit-secret-key",
aws_session_token="assumed-session-token",
verify=True,
)
# Should call assume_role
mock_boto3_client.assert_called_once_with("sts", **expected_client_kwargs)
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"