mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
consistently passing region and endpoint args into explicit credentials irsa
This commit is contained in:
parent
782df9283f
commit
be885ad236
2 changed files with 63 additions and 41 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue