mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(bedrock): decouple STS region from Bedrock aws_region_name (#28245)
* fix(bedrock): decouple STS region from Bedrock aws_region_name STS AssumeRole now resolves signing region from aws_sts_endpoint (parsed host) or AWS_REGION/AWS_DEFAULT_REGION instead of aws_region_name, fixing air-gapped cross-region Bedrock setups and endpoint/signature mismatches. Co-authored-by: Cursor <cursoragent@cursor.com> * test(bedrock): add regression coverage for _build_sts_client_kwargs Parametrize _resolve_sts_region and _build_sts_client_kwargs matrix cases, and assert IRSA/web-identity paths use aligned STS endpoint and region_name. Co-authored-by: Cursor <cursoragent@cursor.com> * refactor(bedrock): tighten STS region helpers and drop redundant web-identity endpoint synthesis Co-authored-by: Cursor <cursoragent@cursor.com> * test(bedrock): cover FIPS, GovCloud, and China STS endpoints Addresses greptile P2: regex sts(?:-fips)? supported sts-fips hosts but was not exercised by the parametrized parse test. Co-authored-by: Cursor <cursoragent@cursor.com> --------- Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
a3c953ed4e
commit
1b141bc588
2 changed files with 377 additions and 46 deletions
|
|
@ -44,6 +44,12 @@ else:
|
|||
# (e.g. "us-east-1", "eu-west-2", "us-gov-west-1", "cn-north-1").
|
||||
_VALID_AWS_REGION_PATTERN = re.compile(r"\A[a-z0-9-]+\Z")
|
||||
|
||||
# Regional STS hostnames, e.g. sts.eu-west-1.amazonaws.com or
|
||||
# vpce-xxx.sts.eu-west-1.vpce.amazonaws.com
|
||||
_STS_REGION_FROM_ENDPOINT_PATTERN = re.compile(
|
||||
r"(?:^|\.)sts(?:-fips)?\.([a-z0-9-]+)\.(?:amazonaws\.com(?:\.cn)?|vpce\.amazonaws\.com)"
|
||||
)
|
||||
|
||||
|
||||
class Boto3CredentialsInfo(BaseModel):
|
||||
credentials: Credentials
|
||||
|
|
@ -651,6 +657,40 @@ class BaseAWSLLM:
|
|||
"Region names must contain only lowercase letters, digits, and hyphens."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _parse_sts_region_from_endpoint(
|
||||
aws_sts_endpoint: Optional[str],
|
||||
) -> Optional[str]:
|
||||
"""Extract region from sts.{region}.amazonaws.com or vpce-x.sts.{region}.vpce.amazonaws.com."""
|
||||
if not aws_sts_endpoint:
|
||||
return None
|
||||
host = urllib.parse.urlparse(aws_sts_endpoint).hostname or ""
|
||||
match = _STS_REGION_FROM_ENDPOINT_PATTERN.search(host)
|
||||
return match.group(1) if match else None
|
||||
|
||||
@staticmethod
|
||||
def _resolve_sts_region(aws_sts_endpoint: Optional[str] = None) -> Optional[str]:
|
||||
"""STS signing region: parsed from aws_sts_endpoint else AWS_REGION / AWS_DEFAULT_REGION."""
|
||||
return (
|
||||
BaseAWSLLM._parse_sts_region_from_endpoint(aws_sts_endpoint)
|
||||
or os.getenv("AWS_REGION")
|
||||
or os.getenv("AWS_DEFAULT_REGION")
|
||||
)
|
||||
|
||||
def _build_sts_client_kwargs(
|
||||
self,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
) -> dict:
|
||||
"""STS client kwargs with aligned endpoint_url and region_name (SigV4)."""
|
||||
kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)}
|
||||
if aws_sts_endpoint is not None:
|
||||
kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
sts_region = self._resolve_sts_region(aws_sts_endpoint)
|
||||
if sts_region is not None:
|
||||
kwargs["region_name"] = sts_region
|
||||
return kwargs
|
||||
|
||||
def get_aws_region_name_for_non_llm_api_calls(
|
||||
self,
|
||||
aws_region_name: Optional[str] = None,
|
||||
|
|
@ -805,11 +845,6 @@ class BaseAWSLLM:
|
|||
f"IN Web Identity Token: {aws_web_identity_token} | Role Name: {aws_role_name} | Session Name: {aws_session_name}"
|
||||
)
|
||||
|
||||
if aws_sts_endpoint is None:
|
||||
sts_endpoint = f"https://sts.{aws_region_name}.amazonaws.com"
|
||||
else:
|
||||
sts_endpoint = aws_sts_endpoint
|
||||
|
||||
oidc_token = get_secret(aws_web_identity_token)
|
||||
|
||||
if oidc_token is None:
|
||||
|
|
@ -818,13 +853,13 @@ class BaseAWSLLM:
|
|||
status_code=401,
|
||||
)
|
||||
|
||||
sts_client_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client(
|
||||
"sts",
|
||||
region_name=aws_region_name,
|
||||
endpoint_url=sts_endpoint,
|
||||
verify=self._get_ssl_verify(ssl_verify),
|
||||
)
|
||||
sts_client = boto3.client("sts", **sts_client_kwargs)
|
||||
|
||||
# https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html
|
||||
# https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html
|
||||
|
|
@ -865,7 +900,6 @@ class BaseAWSLLM:
|
|||
irsa_role_arn: str,
|
||||
aws_role_name: str,
|
||||
aws_session_name: str,
|
||||
region: str,
|
||||
web_identity_token_file: str,
|
||||
aws_external_id: Optional[str] = None,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
|
|
@ -880,12 +914,10 @@ 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
|
||||
irsa_sts_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
# Create an STS client without credentials
|
||||
with tracer.trace("boto3.client(sts) for manual IRSA"):
|
||||
|
|
@ -942,7 +974,6 @@ class BaseAWSLLM:
|
|||
self,
|
||||
aws_role_name: str,
|
||||
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,
|
||||
|
|
@ -950,12 +981,10 @@ class BaseAWSLLM:
|
|||
"""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
|
||||
irsa_sts_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
verbose_logger.debug("Same account role assumption, using automatic IRSA")
|
||||
with tracer.trace("boto3.client(sts) with automatic IRSA"):
|
||||
|
|
@ -1028,12 +1057,6 @@ 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 (
|
||||
|
|
@ -1049,16 +1072,12 @@ class BaseAWSLLM:
|
|||
)
|
||||
|
||||
try:
|
||||
# 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:
|
||||
sts_response = self._handle_irsa_cross_account(
|
||||
irsa_role_arn,
|
||||
aws_role_name,
|
||||
aws_session_name,
|
||||
region,
|
||||
web_identity_token_file,
|
||||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
|
|
@ -1068,7 +1087,6 @@ class BaseAWSLLM:
|
|||
sts_response = self._handle_irsa_same_account(
|
||||
aws_role_name,
|
||||
aws_session_name,
|
||||
region,
|
||||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
|
|
@ -1092,11 +1110,10 @@ 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 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
|
||||
sts_client_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
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", **sts_client_kwargs)
|
||||
|
|
|
|||
|
|
@ -869,14 +869,18 @@ def test_different_roles_without_session_names_should_not_share_cache():
|
|||
({}, {"verify": True}),
|
||||
(
|
||||
{"aws_region_name": "us-east-1"},
|
||||
{"region_name": "us-east-1", "verify": True},
|
||||
{"verify": True},
|
||||
),
|
||||
(
|
||||
{"aws_sts_endpoint": "https://sts.eu-west-1.amazonaws.com"},
|
||||
{"endpoint_url": "https://sts.eu-west-1.amazonaws.com", "verify": True},
|
||||
{
|
||||
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
|
||||
"region_name": "eu-west-1",
|
||||
"verify": True,
|
||||
},
|
||||
),
|
||||
],
|
||||
ids=["no_region_or_endpoint", "regional_sts", "explicit_sts_endpoint"],
|
||||
ids=["no_region_or_endpoint", "bedrock_region_ignored_for_sts", "explicit_sts_endpoint"],
|
||||
)
|
||||
def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs):
|
||||
"""
|
||||
|
|
@ -925,6 +929,316 @@ def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs):
|
|||
assert ttl is not None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint,expected_region",
|
||||
[
|
||||
("https://sts.eu-west-1.amazonaws.com", "eu-west-1"),
|
||||
("https://sts.us-east-1.amazonaws.com", "us-east-1"),
|
||||
("https://sts-fips.us-east-1.amazonaws.com", "us-east-1"),
|
||||
("https://sts-fips.us-gov-west-1.amazonaws.com", "us-gov-west-1"),
|
||||
("https://sts.us-gov-west-1.amazonaws.com", "us-gov-west-1"),
|
||||
("https://sts.cn-north-1.amazonaws.com.cn", "cn-north-1"),
|
||||
(
|
||||
"https://vpce-abc123.sts.eu-west-1.vpce.amazonaws.com",
|
||||
"eu-west-1",
|
||||
),
|
||||
("https://sts.amazonaws.com", None),
|
||||
("https://invalid.example.com", None),
|
||||
],
|
||||
)
|
||||
def test_parse_sts_region_from_endpoint(endpoint, expected_region):
|
||||
assert BaseAWSLLM._parse_sts_region_from_endpoint(endpoint) == expected_region
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"env,aws_sts_endpoint,expected_region",
|
||||
[
|
||||
({}, None, None),
|
||||
({"AWS_REGION": "us-east-1"}, None, "us-east-1"),
|
||||
({"AWS_DEFAULT_REGION": "ap-southeast-1"}, None, "ap-southeast-1"),
|
||||
({}, "https://sts.eu-west-1.amazonaws.com", "eu-west-1"),
|
||||
(
|
||||
{"AWS_REGION": "us-east-1"},
|
||||
"https://sts.eu-west-1.amazonaws.com",
|
||||
"eu-west-1",
|
||||
),
|
||||
({}, "https://sts.amazonaws.com", None),
|
||||
(
|
||||
{},
|
||||
"https://vpce-abc.sts.eu-central-1.vpce.amazonaws.com",
|
||||
"eu-central-1",
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"no_env_no_endpoint",
|
||||
"env_region",
|
||||
"env_default_region",
|
||||
"parsed_from_endpoint",
|
||||
"parsed_endpoint_over_env",
|
||||
"global_endpoint",
|
||||
"vpce_endpoint",
|
||||
],
|
||||
)
|
||||
def test_resolve_sts_region(env, aws_sts_endpoint, expected_region):
|
||||
with patch.dict(os.environ, env, clear=True):
|
||||
assert (
|
||||
BaseAWSLLM._resolve_sts_region(aws_sts_endpoint=aws_sts_endpoint)
|
||||
== expected_region
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"env,aws_sts_endpoint,ssl_verify,expected",
|
||||
[
|
||||
({}, None, None, {"verify": True}),
|
||||
(
|
||||
{"AWS_REGION": "us-east-1"},
|
||||
None,
|
||||
None,
|
||||
{"verify": True, "region_name": "us-east-1"},
|
||||
),
|
||||
(
|
||||
{},
|
||||
"https://sts.eu-west-1.amazonaws.com",
|
||||
None,
|
||||
{
|
||||
"verify": True,
|
||||
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
|
||||
"region_name": "eu-west-1",
|
||||
},
|
||||
),
|
||||
(
|
||||
{"AWS_REGION": "us-east-1"},
|
||||
"https://sts.eu-west-1.amazonaws.com",
|
||||
None,
|
||||
{
|
||||
"verify": True,
|
||||
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
|
||||
"region_name": "eu-west-1",
|
||||
},
|
||||
),
|
||||
(
|
||||
{},
|
||||
"https://sts.amazonaws.com",
|
||||
None,
|
||||
{"verify": True, "endpoint_url": "https://sts.amazonaws.com"},
|
||||
),
|
||||
(
|
||||
{},
|
||||
"https://vpce-abc.sts.eu-central-1.vpce.amazonaws.com",
|
||||
None,
|
||||
{
|
||||
"verify": True,
|
||||
"endpoint_url": "https://vpce-abc.sts.eu-central-1.vpce.amazonaws.com",
|
||||
"region_name": "eu-central-1",
|
||||
},
|
||||
),
|
||||
({}, None, False, {"verify": False}),
|
||||
(
|
||||
{"AWS_DEFAULT_REGION": "ap-southeast-1"},
|
||||
None,
|
||||
None,
|
||||
{"verify": True, "region_name": "ap-southeast-1"},
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"default_verify_only",
|
||||
"env_region",
|
||||
"endpoint_with_parsed_region",
|
||||
"endpoint_parsed_over_env",
|
||||
"global_endpoint_no_region",
|
||||
"vpce_endpoint",
|
||||
"ssl_verify_false",
|
||||
"env_default_region",
|
||||
],
|
||||
)
|
||||
def test_build_sts_client_kwargs(env, aws_sts_endpoint, ssl_verify, expected):
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
with patch.dict(os.environ, env, clear=True):
|
||||
assert (
|
||||
base_aws_llm._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
== expected
|
||||
)
|
||||
|
||||
|
||||
def test_irsa_cross_account_sts_client_uses_resolved_region():
|
||||
"""IRSA cross-account path must use _build_sts_client_kwargs (env region, not Bedrock)."""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
import tempfile
|
||||
|
||||
with tempfile.NamedTemporaryFile(mode="w", delete=False) as f:
|
||||
f.write("test-web-identity-token")
|
||||
token_file = f.name
|
||||
|
||||
try:
|
||||
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": "eu-west-1",
|
||||
},
|
||||
clear=True,
|
||||
):
|
||||
mock_sts_client = MagicMock()
|
||||
mock_sts_client.assume_role_with_web_identity.return_value = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "temp-key",
|
||||
"SecretAccessKey": "temp-secret",
|
||||
"SessionToken": "temp-token",
|
||||
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
|
||||
}
|
||||
}
|
||||
mock_sts_client.assume_role.return_value = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "assumed-key",
|
||||
"SecretAccessKey": "assumed-secret",
|
||||
"SessionToken": "assumed-token",
|
||||
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
|
||||
}
|
||||
}
|
||||
|
||||
with patch(
|
||||
"boto3.client", return_value=mock_sts_client
|
||||
) as mock_boto3_client:
|
||||
base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_session_token=None,
|
||||
aws_role_name="arn:aws:iam::222222222222:role/target-role",
|
||||
aws_session_name="test-session",
|
||||
aws_region_name="eu-central-1",
|
||||
)
|
||||
|
||||
for call in mock_boto3_client.call_args_list:
|
||||
assert call.args == ("sts",)
|
||||
assert call.kwargs["region_name"] == "eu-west-1"
|
||||
assert call.kwargs["verify"] is True
|
||||
finally:
|
||||
os.unlink(token_file)
|
||||
|
||||
|
||||
def test_web_identity_token_sts_client_uses_build_sts_client_kwargs():
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
mock_sts_client = MagicMock()
|
||||
mock_sts_client.assume_role_with_web_identity.return_value = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "key",
|
||||
"SecretAccessKey": "secret",
|
||||
"SessionToken": "token",
|
||||
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
|
||||
},
|
||||
"PackedPolicySize": 0,
|
||||
}
|
||||
|
||||
with patch.dict(os.environ, {"AWS_REGION": "eu-west-1"}, clear=True):
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
with patch(
|
||||
"litellm.llms.bedrock.base_aws_llm.get_secret",
|
||||
return_value="oidc-token",
|
||||
):
|
||||
base_aws_llm._auth_with_web_identity_token(
|
||||
aws_web_identity_token="my-token",
|
||||
aws_role_name="arn:aws:iam::111111111111:role/target",
|
||||
aws_session_name="test-session",
|
||||
aws_region_name="eu-central-1",
|
||||
aws_sts_endpoint="https://sts.eu-west-1.amazonaws.com",
|
||||
)
|
||||
|
||||
mock_boto3_client.assert_called_once_with(
|
||||
"sts",
|
||||
verify=True,
|
||||
endpoint_url="https://sts.eu-west-1.amazonaws.com",
|
||||
region_name="eu-west-1",
|
||||
)
|
||||
|
||||
|
||||
def test_sts_uses_workload_region_not_bedrock_region():
|
||||
"""Air-gapped: Bedrock in eu-central-1, STS VPC endpoint in eu-west-1 via AWS_REGION."""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
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_client = MagicMock()
|
||||
mock_sts_client.assume_role.return_value = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "assumed-access-key",
|
||||
"SecretAccessKey": "assumed-secret-key",
|
||||
"SessionToken": "assumed-session-token",
|
||||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
|
||||
with patch.dict(os.environ, {"AWS_REGION": "eu-west-1"}, clear=True):
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_session_token=None,
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="test-session",
|
||||
aws_region_name="eu-central-1",
|
||||
)
|
||||
mock_boto3_client.assert_called_with(
|
||||
"sts",
|
||||
region_name="eu-west-1",
|
||||
verify=True,
|
||||
)
|
||||
|
||||
|
||||
def test_sts_endpoint_region_matches_bedrock_region_param():
|
||||
"""aws_sts_endpoint signing region must not follow aws_region_name when they differ."""
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
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_client = MagicMock()
|
||||
mock_sts_client.assume_role.return_value = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "assumed-access-key",
|
||||
"SecretAccessKey": "assumed-secret-key",
|
||||
"SessionToken": "assumed-session-token",
|
||||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
|
||||
env_without_irsa = {
|
||||
k: v
|
||||
for k, v in os.environ.items()
|
||||
if k
|
||||
not in (
|
||||
"AWS_ROLE_ARN",
|
||||
"AWS_WEB_IDENTITY_TOKEN_FILE",
|
||||
"AWS_REGION",
|
||||
"AWS_DEFAULT_REGION",
|
||||
)
|
||||
}
|
||||
with patch.dict(env_without_irsa, clear=True):
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_session_token=None,
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="test-session",
|
||||
aws_region_name="eu-central-1",
|
||||
aws_sts_endpoint="https://sts.eu-west-1.amazonaws.com",
|
||||
)
|
||||
mock_boto3_client.assert_called_with(
|
||||
"sts",
|
||||
endpoint_url="https://sts.eu-west-1.amazonaws.com",
|
||||
region_name="eu-west-1",
|
||||
verify=True,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"role_kwargs,expected_client_kwargs",
|
||||
[
|
||||
|
|
@ -940,7 +1254,6 @@ def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs):
|
|||
(
|
||||
{"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",
|
||||
|
|
@ -951,6 +1264,7 @@ def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs):
|
|||
{"aws_sts_endpoint": "https://sts.eu-west-1.amazonaws.com"},
|
||||
{
|
||||
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
|
||||
"region_name": "eu-west-1",
|
||||
"aws_access_key_id": "explicit-access-key",
|
||||
"aws_secret_access_key": "explicit-secret-key",
|
||||
"aws_session_token": "assumed-session-token",
|
||||
|
|
@ -958,7 +1272,7 @@ def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs):
|
|||
},
|
||||
),
|
||||
],
|
||||
ids=["no_region_or_endpoint", "regional_sts", "explicit_sts_endpoint"],
|
||||
ids=["no_region_or_endpoint", "bedrock_region_ignored_for_sts", "explicit_sts_endpoint"],
|
||||
)
|
||||
def test_explicit_credentials_used_when_provided(role_kwargs, expected_client_kwargs):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue