From 6422178b76da904b1ecdeb30171050a686e328a1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Pedro=20Mart=C3=ADn?= Date: Mon, 31 Aug 2026 18:13:30 +0200 Subject: [PATCH] feat(sdk): AWS partition selection via PROWLER_AWS_PARTITION (#12680) Co-authored-by: David --- .../aws-partition-env-var.added.md | 1 + prowler/providers/aws/aws_provider.py | 134 ++++++++- tests/providers/aws/aws_provider_test.py | 256 ++++++++++++++++++ 3 files changed, 380 insertions(+), 11 deletions(-) create mode 100644 prowler/changelog.d/aws-partition-env-var.added.md diff --git a/prowler/changelog.d/aws-partition-env-var.added.md b/prowler/changelog.d/aws-partition-env-var.added.md new file mode 100644 index 0000000000..c94ed6b23d --- /dev/null +++ b/prowler/changelog.d/aws-partition-env-var.added.md @@ -0,0 +1 @@ +`PROWLER_AWS_PARTITION` environment variable to select the AWS partition used for STS credential validation and scan bootstrap, with a clear error when the account belongs to a different partition diff --git a/prowler/providers/aws/aws_provider.py b/prowler/providers/aws/aws_provider.py index b4c9ed3771..d6b24c748c 100644 --- a/prowler/providers/aws/aws_provider.py +++ b/prowler/providers/aws/aws_provider.py @@ -1,6 +1,7 @@ import os import pathlib from datetime import datetime +from functools import lru_cache from re import fullmatch from typing import Optional @@ -671,7 +672,12 @@ class AwsProvider(Provider): if mfa: session = Session(**session_arguments) session._session.set_default_client_config(session_config) - sts_client = session.client("sts") + sts_region = ( + get_env_partition_bootstrap_region() + or session.region_name + or AWS_STS_GLOBAL_ENDPOINT_REGION + ) + sts_client = AwsProvider.create_sts_session(session, sts_region) # TODO: pass values from the input mfa_info = AwsProvider.input_role_mfa_token_and_code() @@ -1352,7 +1358,7 @@ class AwsProvider(Provider): @staticmethod def test_connection( profile: str = None, - aws_region: str = AWS_STS_GLOBAL_ENDPOINT_REGION, + aws_region: str = None, role_arn: str = None, role_session_name: str = ROLE_SESSION_NAME, session_duration: int = 3600, @@ -1369,7 +1375,9 @@ class AwsProvider(Provider): Args: profile (str): The AWS profile to use for the session. - aws_region (str): The AWS region to validate the credentials in. + aws_region (str): The AWS region to validate the credentials in. When not + provided, it defaults to the bootstrap region of the partition set in + the PROWLER_AWS_PARTITION environment variable or, if unset, to us-east-1. role_arn (str): The ARN of the IAM role to assume. role_session_name (str): The name of the role session. session_duration (int): The duration of the assumed role session in seconds. @@ -1412,6 +1420,12 @@ class AwsProvider(Provider): Connection(is_connected=True, Error=None)) """ try: + if aws_region is None: + aws_region = ( + get_env_partition_bootstrap_region() + or AWS_STS_GLOBAL_ENDPOINT_REGION + ) + session = AwsProvider.setup_session( mfa=mfa_enabled, profile=profile, @@ -1430,6 +1444,7 @@ class AwsProvider(Provider): external_id=external_id, mfa_enabled=mfa_enabled, role_session_name=role_session_name, + sts_region=aws_region, ) assumed_role_credentials = AwsProvider.assume_role( session, @@ -1451,6 +1466,13 @@ class AwsProvider(Provider): if provider_id and caller_identity.account != provider_id: raise AWSInvalidProviderIdError(file=pathlib.Path(__file__).name) + # Validate that the account belongs to the configured partition, if any + env_partition = os.environ.get("PROWLER_AWS_PARTITION", "").strip() + if env_partition and caller_identity.arn.partition != env_partition: + raise AWSInvalidPartitionError( + message=f"The AWS account is in the {caller_identity.arn.partition} partition, but this deployment is configured for the {env_partition} partition via PROWLER_AWS_PARTITION" + ) + return Connection( is_connected=True, ) @@ -1591,6 +1613,14 @@ class AwsProvider(Provider): raise session_token_expired return Connection(error=session_token_expired) + except AWSInvalidPartitionError as invalid_partition_error: + logger.error( + f"{invalid_partition_error.__class__.__name__}[{invalid_partition_error.__traceback__.tb_lineno}]: {invalid_partition_error}" + ) + if raise_on_exception: + raise invalid_partition_error + return Connection(error=invalid_partition_error) + except Exception as error: logger.critical( f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" @@ -1618,14 +1648,9 @@ class AwsProvider(Provider): sts_client = create_sts_session(session, 'us-west-2') """ try: - if os.environ.get("AWS_ENDPOINT_URL"): - sts_endpoint_url = os.environ["AWS_ENDPOINT_URL"] - elif aws_region.startswith("cn-"): - sts_endpoint_url = f"https://sts.{aws_region}.amazonaws.com.cn" - elif aws_region.startswith("eusc-"): - sts_endpoint_url = f"https://sts.{aws_region}.amazonaws.eu" - else: - sts_endpoint_url = f"https://sts.{aws_region}.amazonaws.com" + # Botocore resolves the regional STS endpoint for every partition + # (China, EUSC, GovCloud, ISO); AWS_ENDPOINT_URL overrides it + sts_endpoint_url = os.environ.get("AWS_ENDPOINT_URL") or None return session.client("sts", aws_region, endpoint_url=sts_endpoint_url) except Exception as error: logger.critical( @@ -1702,6 +1727,80 @@ def read_aws_regions_file() -> dict: return data +@lru_cache(maxsize=1) +def get_botocore_partition_regions() -> dict: + """ + Get the AWS partitions and their bootstrap region candidates from the + botocore endpoints data. + + The region of the partition's global STS endpoint, when declared, is moved + to the front since it is never an opt-in region; the rest are sorted + alphabetically. + + Returns: + dict: A dictionary mapping each partition name to its list of regions. + """ + endpoints_data = BotocoreSession().get_data("endpoints") + partition_regions = {} + for partition in endpoints_data["partitions"]: + regions = sorted(partition.get("regions", {})) + sts_service = partition.get("services", {}).get("sts", {}) + global_endpoint = sts_service.get("partitionEndpoint") + global_region = ( + sts_service.get("endpoints", {}) + .get(global_endpoint, {}) + .get("credentialScope", {}) + .get("region") + ) + if global_region in regions: + regions.remove(global_region) + regions.insert(0, global_region) + partition_regions[partition["partition"]] = regions + return partition_regions + + +def get_env_partition_regions() -> Optional[list]: + """ + Get the bootstrap region candidates for the partition set in the + PROWLER_AWS_PARTITION environment variable. + + Returns: + Optional[list]: The regions of the configured partition, preferred + bootstrap region first, or None when the environment variable is + not set. + + Raises: + AWSInvalidPartitionError: If the value is not a partition known to botocore. + """ + raw_partition = os.environ.get("PROWLER_AWS_PARTITION", "").strip() + if not raw_partition: + return None + + partition_regions = get_botocore_partition_regions() + regions = partition_regions.get(raw_partition) + if not regions: + raise AWSInvalidPartitionError( + message=f"Invalid partition: {raw_partition} set in PROWLER_AWS_PARTITION. Valid partitions: {', '.join(sorted(partition_regions))}" + ) + return regions + + +def get_env_partition_bootstrap_region() -> Optional[str]: + """ + Get the STS bootstrap region for the partition set in the + PROWLER_AWS_PARTITION environment variable. + + Returns: + Optional[str]: The preferred bootstrap region of the configured + partition, or None when the environment variable is not set. + + Raises: + AWSInvalidPartitionError: If the value is not a partition known to botocore. + """ + regions = get_env_partition_regions() + return regions[0] if regions else None + + # TODO: This can be moved to another class since it doesn't need self def get_aws_region_for_sts( session_region: str, @@ -1711,6 +1810,10 @@ def get_aws_region_for_sts( """ Get the AWS region for the STS Assume Role operation. + The precedence is: explicit regions, the partition set in the + PROWLER_AWS_PARTITION environment variable, the session region and, + finally, the bootstrap region candidates. + Args: - session_region (str): The region configured in the AWS session. - regions (set[str]): The regions passed with the -f/--region/--filter-region option. @@ -1730,6 +1833,15 @@ def get_aws_region_for_sts( if region not in excluded_regions: return region + env_partition_regions = get_env_partition_regions() + if env_partition_regions: + # The configured partition constrains the whole fallback chain: prefer + # a non-excluded region, but never leave the partition + for region in env_partition_regions: + if region not in excluded_regions: + return region + return env_partition_regions[0] + if session_region and session_region not in excluded_regions: return session_region diff --git a/tests/providers/aws/aws_provider_test.py b/tests/providers/aws/aws_provider_test.py index f874ca8812..4bdb438b99 100644 --- a/tests/providers/aws/aws_provider_test.py +++ b/tests/providers/aws/aws_provider_test.py @@ -1783,6 +1783,29 @@ aws: assert sts_session._endpoint._endpoint_prefix == "sts" assert sts_session._endpoint.host == f"https://sts.{aws_region}.amazonaws.eu" + @mock_aws + def test_create_sts_session_empty_endpoint_url(self): + current_session = session.Session() + aws_region = AWS_REGION_US_EAST_1 + with mock.patch.dict(os.environ, {"AWS_ENDPOINT_URL": ""}): + sts_session = AwsProvider.create_sts_session(current_session, aws_region) + + assert sts_session._service_model.service_name == "sts" + assert sts_session._client_config.region_name == aws_region + assert sts_session._endpoint._endpoint_prefix == "sts" + assert sts_session._endpoint.host == f"https://sts.{aws_region}.amazonaws.com" + + @mock_aws + def test_create_sts_session_iso(self): + current_session = session.Session() + aws_region = "us-iso-east-1" + sts_session = AwsProvider.create_sts_session(current_session, aws_region) + + assert sts_session._service_model.service_name == "sts" + assert sts_session._client_config.region_name == aws_region + assert sts_session._endpoint._endpoint_prefix == "sts" + assert sts_session._endpoint.host == f"https://sts.{aws_region}.c2s.ic.gov" + @mock_aws @patch( "prowler.lib.check.utils.recover_checks_from_provider", @@ -2219,6 +2242,239 @@ aws: == AWS_REGION_US_EAST_1 ) + def test_get_aws_region_for_sts_env_partition_gov_cloud(self): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_GOV_CLOUD_PARTITION}, + clear=False, + ): + assert get_aws_region_for_sts(None, None) == AWS_REGION_GOV_CLOUD_US_EAST_1 + + def test_get_aws_region_for_sts_env_partition_china(self): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_CHINA_PARTITION}, + clear=False, + ): + assert get_aws_region_for_sts(None, None) == AWS_REGION_CN_NORTH_1 + + def test_get_aws_region_for_sts_env_partition_eusc(self): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_EUSC_PARTITION}, + clear=False, + ): + assert get_aws_region_for_sts(None, None) == AWS_REGION_EUSC_DE_EAST_1 + + def test_get_aws_region_for_sts_env_partition_iso(self): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_ISO_PARTITION}, + clear=False, + ): + assert get_aws_region_for_sts(None, None) == "us-iso-east-1" + + def test_get_aws_region_for_sts_env_partition_overrides_session_region(self): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_GOV_CLOUD_PARTITION}, + clear=False, + ): + assert ( + get_aws_region_for_sts(AWS_REGION_EU_WEST_1, None) + == AWS_REGION_GOV_CLOUD_US_EAST_1 + ) + + def test_get_aws_region_for_sts_input_regions_take_precedence_over_env_partition( + self, + ): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_GOV_CLOUD_PARTITION}, + clear=False, + ): + assert ( + get_aws_region_for_sts(None, {AWS_REGION_EU_WEST_1}) + == AWS_REGION_EU_WEST_1 + ) + + def test_get_aws_region_for_sts_env_partition_invalid_raises(self): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": "aws-invalid"}, + clear=False, + ): + with pytest.raises(AWSInvalidPartitionError): + get_aws_region_for_sts(None, None) + + @mock_aws + def test_test_connection_uses_env_partition_sts_region(self): + with ( + mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_GOV_CLOUD_PARTITION}, + clear=False, + ), + mock.patch.object( + AwsProvider, + "validate_credentials", + return_value=AWSCallerIdentity( + user_id="test-user-id", + account=AWS_ACCOUNT_NUMBER, + arn=ARN(AWS_GOV_CLOUD_ACCOUNT_ARN), + region=AWS_REGION_GOV_CLOUD_US_EAST_1, + ), + ) as mock_validate_credentials, + ): + connection = AwsProvider.test_connection( + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + raise_on_exception=False, + ) + + assert connection.is_connected + assert ( + mock_validate_credentials.call_args.args[1] + == AWS_REGION_GOV_CLOUD_US_EAST_1 + ) + + @mock_aws + def test_test_connection_role_uses_env_partition_sts_region(self): + with ( + mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_GOV_CLOUD_PARTITION}, + clear=False, + ), + mock.patch.object( + AwsProvider, + "assume_role", + return_value=AWSCredentials( + aws_access_key_id="assumed-access-key", + aws_secret_access_key="assumed-secret-key", + aws_session_token="assumed-session-token", + expiration=datetime.now(), + ), + ) as mock_assume_role, + mock.patch.object( + AwsProvider, + "validate_credentials", + return_value=AWSCallerIdentity( + user_id="test-user-id", + account=AWS_ACCOUNT_NUMBER, + arn=ARN(AWS_GOV_CLOUD_ACCOUNT_ARN), + region=AWS_REGION_GOV_CLOUD_US_EAST_1, + ), + ), + ): + connection = AwsProvider.test_connection( + role_arn=f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/test-role", + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + raise_on_exception=False, + ) + + assert connection.is_connected + assumed_role_info = mock_assume_role.call_args.args[1] + assert assumed_role_info.sts_region == AWS_REGION_GOV_CLOUD_US_EAST_1 + + def test_get_aws_region_for_sts_env_partition_commercial(self): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_COMMERCIAL_PARTITION}, + clear=False, + ): + assert get_aws_region_for_sts(None, None) == AWS_REGION_US_EAST_1 + + def test_get_aws_region_for_sts_env_partition_excluded_region_stays_in_partition( + self, + ): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_GOV_CLOUD_PARTITION}, + clear=False, + ): + assert ( + get_aws_region_for_sts(None, None, {AWS_REGION_GOV_CLOUD_US_EAST_1}) + == "us-gov-west-1" + ) + + def test_get_aws_region_for_sts_env_partition_all_regions_excluded_stays_in_partition( + self, + ): + with mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_GOV_CLOUD_PARTITION}, + clear=False, + ): + assert ( + get_aws_region_for_sts( + None, None, {AWS_REGION_GOV_CLOUD_US_EAST_1, "us-gov-west-1"} + ) + == AWS_REGION_GOV_CLOUD_US_EAST_1 + ) + + @mock_aws + def test_setup_session_mfa_uses_env_partition_sts_region(self): + with ( + mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_GOV_CLOUD_PARTITION}, + clear=False, + ), + mock.patch.object( + AwsProvider, + "input_role_mfa_token_and_code", + return_value=AWSMFAInfo( + arn=f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:mfa/test", + totp="123456", + ), + ), + mock.patch.object( + AwsProvider, + "create_sts_session", + side_effect=AwsProvider.create_sts_session, + ) as mock_create_sts_session, + ): + AwsProvider.setup_session( + mfa=True, + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + ) + + assert ( + mock_create_sts_session.call_args.args[1] + == AWS_REGION_GOV_CLOUD_US_EAST_1 + ) + + @mock_aws + def test_test_connection_env_partition_mismatch(self): + with ( + mock.patch.dict( + os.environ, + {"PROWLER_AWS_PARTITION": AWS_GOV_CLOUD_PARTITION}, + clear=False, + ), + mock.patch.object( + AwsProvider, + "validate_credentials", + return_value=AWSCallerIdentity( + user_id="test-user-id", + account=AWS_ACCOUNT_NUMBER, + arn=ARN(AWS_ACCOUNT_ARN), + region=AWS_REGION_US_EAST_1, + ), + ), + ): + connection = AwsProvider.test_connection( + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + raise_on_exception=False, + ) + + assert not connection.is_connected + assert isinstance(connection.error, AWSInvalidPartitionError) + def test_get_profile_region_avoids_excluded_session_region(self): mocked_session = mock.Mock(region_name=AWS_REGION_EU_WEST_1)