From 2640400ffba04b046d1527566c0d047d0611043c Mon Sep 17 00:00:00 2001 From: Milan Date: Tue, 19 May 2026 13:27:05 +0300 Subject: [PATCH] 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 --- litellm/llms/bedrock/base_aws_llm.py | 126 ++++++++++++------ .../llms/bedrock/test_base_aws_llm.py | 114 +++++++++++++++- 2 files changed, 196 insertions(+), 44 deletions(-) diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 9dd2b055a12..cdfeabba72a 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -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 @@ -633,6 +639,63 @@ 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 the AWS region from a standard STS endpoint URL. + + Supports public regional endpoints (sts.{region}.amazonaws.com) and + VPC interface endpoints (vpce-....sts.{region}.vpce.amazonaws.com). + Returns None for the global endpoint (sts.amazonaws.com) or unparseable URLs. + """ + if not aws_sts_endpoint: + return None + try: + host = urllib.parse.urlparse(aws_sts_endpoint).hostname or "" + except Exception: + return None + match = _STS_REGION_FROM_ENDPOINT_PATTERN.search(host) + if not match: + return None + region = match.group(1) + if _VALID_AWS_REGION_PATTERN.match(region): + return region + return None + + @staticmethod + def _resolve_sts_region( + aws_sts_endpoint: Optional[str] = None, + ) -> Optional[str]: + """ + Resolve the region used for STS SigV4 signing. + + Bedrock's aws_region_name is intentionally not used here; STS follows the + caller environment or an explicit aws_sts_endpoint. + """ + if aws_sts_endpoint: + parsed_region = BaseAWSLLM._parse_sts_region_from_endpoint( + aws_sts_endpoint + ) + if parsed_region: + return parsed_region + return 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: + """Build boto3 STS client kwargs with aligned endpoint_url and region_name.""" + kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)} + sts_region = self._resolve_sts_region(aws_sts_endpoint=aws_sts_endpoint) + if aws_sts_endpoint is not None: + kwargs["endpoint_url"] = 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, @@ -787,10 +850,13 @@ 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_region = self._resolve_sts_region(aws_sts_endpoint=aws_sts_endpoint) + if aws_sts_endpoint is not None: sts_endpoint = aws_sts_endpoint + elif sts_region is not None: + sts_endpoint = f"https://sts.{sts_region}.amazonaws.com" + else: + sts_endpoint = "https://sts.amazonaws.com" oidc_token = get_secret(aws_web_identity_token) @@ -800,13 +866,13 @@ class BaseAWSLLM: status_code=401, ) + sts_client_kwargs = self._build_sts_client_kwargs( + aws_sts_endpoint=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 @@ -847,7 +913,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, @@ -862,12 +927,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"): @@ -924,7 +987,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, @@ -932,12 +994,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"): @@ -1010,12 +1070,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 ( @@ -1031,16 +1085,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, @@ -1050,7 +1100,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, @@ -1074,11 +1123,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) diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py index a4969e5dacc..359d58052fb 100644 --- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py +++ b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py @@ -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,106 @@ 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://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 + + +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 +1044,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 +1054,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 +1062,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): """