diff --git a/docs/user-guide/providers/aws/regions-and-partitions.mdx b/docs/user-guide/providers/aws/regions-and-partitions.mdx index bf88aafc2f..369f581826 100644 --- a/docs/user-guide/providers/aws/regions-and-partitions.mdx +++ b/docs/user-guide/providers/aws/regions-and-partitions.mdx @@ -39,6 +39,8 @@ It matters most where nothing else says. Resolving an identity means calling STS A region configured for the session still wins when it belongs to the declared partition, so a deployment in `us-gov-west-1` is not sent to `us-gov-east-1`. A region belonging to a different partition is ignored, since a partition that has been declared explicitly is the more deliberate statement of the two. +When no configured region says which one to prefer, the first region of the partition is tried, and up to two more follow if it cannot be reached. A network that routes to only one region of its partition therefore works without having to declare which one that is. Only a connection failure moves on to the next region: a credential error is reported from the first, since it would be the same everywhere. A region excluded from the scan is tried last, so it is avoided whenever another region of the partition answers. + Set it wherever the scan runs. For deployments that scan from containers, that means the environment of the containers doing the scanning, not only the one accepting the request. diff --git a/prowler/changelog.d/aws-partition-bootstrap-falls-back-to-the-next-region.fixed.md b/prowler/changelog.d/aws-partition-bootstrap-falls-back-to-the-next-region.fixed.md new file mode 100644 index 0000000000..4bd9c855d1 --- /dev/null +++ b/prowler/changelog.d/aws-partition-bootstrap-falls-back-to-the-next-region.fixed.md @@ -0,0 +1 @@ +Bootstrap STS calls now try up to two more regions of the partition declared in `PROWLER_AWS_PARTITION` when the first one cannot be reached, so a deployment that routes to only one region of its partition no longer fails on an endpoint it has no path to. This covers validating credentials, assuming a role and getting an MFA session token diff --git a/prowler/providers/aws/aws_provider.py b/prowler/providers/aws/aws_provider.py index e204318593..784daf0caa 100644 --- a/prowler/providers/aws/aws_provider.py +++ b/prowler/providers/aws/aws_provider.py @@ -3,12 +3,19 @@ import pathlib from datetime import datetime from functools import lru_cache from re import fullmatch -from typing import Optional +from typing import Any, Callable, Optional from boto3.session import Session from botocore.config import Config from botocore.credentials import RefreshableCredentials -from botocore.exceptions import ClientError, NoCredentialsError, ProfileNotFound +from botocore.exceptions import ( + ClientError, + ConnectTimeoutError, + EndpointConnectionError, + NoCredentialsError, + ProfileNotFound, + ReadTimeoutError, +) from botocore.session import Session as BotocoreSession from colorama import Fore, Style from pytz import utc @@ -263,7 +270,10 @@ class AwsProvider(Provider): caller_identity = self.validate_credentials( session=self.session.current_session, aws_region=sts_region, + excluded_regions=excluded_regions, ) + # Later STS calls go where validation got an answer, not where it timed out + sts_region = caller_identity.region logger.info("Credentials validated") ######## @@ -690,8 +700,6 @@ class AwsProvider(Provider): 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() # TODO: validate MFA ARN here @@ -699,8 +707,12 @@ class AwsProvider(Provider): "SerialNumber": mfa_info.arn, "TokenCode": mfa_info.totp, } - session_credentials = sts_client.get_session_token( - **get_session_token_arguments + _, session_credentials = AwsProvider.sts_call_with_partition_failover( + session, + sts_region, + lambda sts_client: sts_client.get_session_token( + **get_session_token_arguments + ), ) mfa_session = Session( aws_access_key_id=session_credentials["Credentials"]["AccessKeyId"], @@ -1244,10 +1256,11 @@ class AwsProvider(Provider): mfa_info = AwsProvider.input_role_mfa_token_and_code() assume_role_arguments["SerialNumber"] = mfa_info.arn assume_role_arguments["TokenCode"] = mfa_info.totp - sts_client = AwsProvider.create_sts_session( - session, assumed_role_info.sts_region + _, assumed_credentials = AwsProvider.sts_call_with_partition_failover( + session, + assumed_role_info.sts_region, + lambda sts_client: sts_client.assume_role(**assume_role_arguments), ) - assumed_credentials = sts_client.assume_role(**assume_role_arguments) # Convert the UTC datetime object to your local timezone credentials_expiration_local_time = ( assumed_credentials["Credentials"]["Expiration"] @@ -1326,30 +1339,98 @@ class AwsProvider(Provider): ) raise error + @staticmethod + def sts_call_with_partition_failover( + session: Session, + aws_region: str, + operation: Callable[[Any], Any], + excluded_regions: set[str] | None = None, + ) -> tuple[str, Any]: + """ + Run a bootstrap STS call, moving on when a region cannot be reached. + + Bootstrap calls happen before anything is known about the credentials, so + the region they go to is a guess whenever none was configured. On a network + that routes to only one region of its partition that guess is fatal, and the + remaining regions of the partition declared in PROWLER_AWS_PARTITION are the + ones worth trying. + + Args: + session (Session): The AWS session object. + aws_region (str): The region to try first. + operation (Callable[[Any], Any]): Receives an STS client and performs + the call. + excluded_regions (set[str] | None): Regions excluded from the scan, + tried after the rest of the partition. + + Returns: + tuple[str, Any]: The region that answered and whatever the operation + returned. + + Raises: + Exception: Whatever the operation raises, or the last connection error + when no region could be reached. + """ + *fallback_regions, last_region = get_partition_bootstrap_candidates( + aws_region, session.region_name, excluded_regions + ) + + for candidate_region in fallback_regions: + try: + sts_client = AwsProvider.create_sts_session(session, candidate_region) + return candidate_region, operation(sts_client) + # The credentials are not at fault, so the next region is worth trying + except ( + EndpointConnectionError, + ConnectTimeoutError, + ReadTimeoutError, + ) as unreachable: + logger.warning( + f"{unreachable.__class__.__name__}[{unreachable.__traceback__.tb_lineno}]: {unreachable}" + ) + + # Nothing is left to try after the last region, so its error is the answer + sts_client = AwsProvider.create_sts_session(session, last_region) + return last_region, operation(sts_client) + @staticmethod def validate_credentials( session: Session, aws_region: str, + excluded_regions: set[str] | None = None, ) -> AWSCallerIdentity: """ Validates the AWS credentials using the provided session and AWS region. + + When the region cannot be reached, the remaining regions of the partition + declared in PROWLER_AWS_PARTITION are tried before giving up. A credential + error is returned from the first region instead, since it would be the same + everywhere. + Args: session (Session): The AWS session object. aws_region (str): The AWS region to validate the credentials. + excluded_regions (set[str] | None): Regions excluded from the scan, + tried after the rest of the partition. Returns: - AWSCallerIdentity: An object containing the caller identity information. + AWSCallerIdentity: An object containing the caller identity information, + including the region that answered. Raises: Exception: If an error occurs during the validation process. """ try: - sts_client = AwsProvider.create_sts_session(session, aws_region) - caller_identity = sts_client.get_caller_identity() + sts_region, caller_identity = AwsProvider.sts_call_with_partition_failover( + session, + aws_region, + lambda sts_client: sts_client.get_caller_identity(), + excluded_regions, + ) # Include the region where the caller_identity has validated the credentials return AWSCallerIdentity( user_id=caller_identity.get("UserId"), account=caller_identity.get("Account"), arn=ARN(caller_identity.get("Arn")), - region=aws_region, + region=sts_region, ) except ClientError as client_error: logger.error( @@ -1846,6 +1927,49 @@ def get_env_partition_bootstrap_region( return regions[0] if regions else None +# An unreachable endpoint costs a connection timeout, so a partition with many +# regions is not walked in full +MAX_STS_BOOTSTRAP_ATTEMPTS = 3 + + +def get_partition_bootstrap_candidates( + aws_region: str, + session_region: Optional[str] = None, + excluded_regions: set[str] | None = None, +) -> list: + """ + Get the STS bootstrap regions to try, in order, starting with the chosen one. + + A deployment reached only through its own region's endpoints has no route to + the rest of its partition, and which region that is cannot be known from the + environment alone: a container may carry a region belonging to no partition + it scans. Offering the remaining regions of the declared partition lets the + bootstrap succeed without anything having to declare the right one. + + Args: + aws_region (str): The region already chosen for the bootstrap call. + session_region (Optional[str]): The region of the AWS session. + excluded_regions (set[str] | None): Regions excluded from the scan. They + go after the rest of the partition, so the bootstrap avoids them + whenever another region answers and still has them as a last resort. + + Returns: + list: The regions to try, preferred first, capped at + MAX_STS_BOOTSTRAP_ATTEMPTS. + """ + excluded_regions = set(excluded_regions or ()) + partition_regions = get_env_partition_regions(session_region) or [] + # sorted() is stable, so the partition order survives on each side of the split + ordered_regions = sorted( + partition_regions, key=lambda region: region in excluded_regions + ) + candidates = [aws_region] + for region in ordered_regions: + if region not in candidates: + candidates.append(region) + return candidates[:MAX_STS_BOOTSTRAP_ATTEMPTS] + + # TODO: This can be moved to another class since it doesn't need self def get_aws_region_for_sts( session_region: str, diff --git a/tests/providers/aws/aws_provider_test.py b/tests/providers/aws/aws_provider_test.py index f899706453..2d759f362a 100644 --- a/tests/providers/aws/aws_provider_test.py +++ b/tests/providers/aws/aws_provider_test.py @@ -17,10 +17,12 @@ from pytest import raises from tzlocal import get_localzone from prowler.providers.aws.aws_provider import ( + MAX_STS_BOOTSTRAP_ATTEMPTS, AwsProvider, get_aws_region_for_sts, get_env_partition_bootstrap_region, get_env_partition_regions, + get_partition_bootstrap_candidates, ) from prowler.providers.aws.config import ( AWS_STS_GLOBAL_ENDPOINT_REGION, @@ -32,6 +34,7 @@ from prowler.providers.aws.config import ( get_default_session_config, ) from prowler.providers.aws.exceptions.exceptions import ( + AWSAccessKeyIDInvalidError, AWSArgumentTypeValidationError, AWSIAMRoleARNInvalidResourceTypeError, AWSInvalidBoto3TimeoutError, @@ -1583,6 +1586,411 @@ aws: assert get_caller_identity.arn.resource == "test-user" assert get_caller_identity.arn.resource_type == "user" + def test_get_partition_bootstrap_candidates_adds_the_rest_of_the_partition( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + + assert get_partition_bootstrap_candidates( + AWS_REGION_GOV_CLOUD_US_EAST_1, AWS_REGION_US_EAST_1 + ) == [AWS_REGION_GOV_CLOUD_US_EAST_1, AWS_REGION_GOV_CLOUD_US_WEST_1] + + def test_get_partition_bootstrap_candidates_without_partition_offers_one_region( + self, monkeypatch + ): + monkeypatch.delenv("PROWLER_AWS_PARTITION", raising=False) + + assert get_partition_bootstrap_candidates(AWS_REGION_EU_WEST_1) == [ + AWS_REGION_EU_WEST_1 + ] + + def test_get_partition_bootstrap_candidates_is_capped(self, monkeypatch): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_COMMERCIAL_PARTITION) + + candidates = get_partition_bootstrap_candidates( + AWS_REGION_EU_WEST_1, AWS_REGION_EU_WEST_1 + ) + + assert len(candidates) == MAX_STS_BOOTSTRAP_ATTEMPTS + assert candidates[0] == AWS_REGION_EU_WEST_1 + + def test_get_partition_bootstrap_candidates_tries_excluded_regions_last( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_COMMERCIAL_PARTITION) + partition_regions = get_env_partition_regions(AWS_REGION_EU_WEST_1) + excluded_regions = set(partition_regions[1:3]) + + candidates = get_partition_bootstrap_candidates( + AWS_REGION_EU_WEST_1, AWS_REGION_EU_WEST_1, excluded_regions + ) + + assert candidates == [AWS_REGION_EU_WEST_1, *partition_regions[3:5]] + + def test_get_partition_bootstrap_candidates_keeps_excluded_regions_as_a_last_resort( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + + assert get_partition_bootstrap_candidates( + AWS_REGION_GOV_CLOUD_US_EAST_1, + AWS_REGION_US_EAST_1, + {AWS_REGION_GOV_CLOUD_US_EAST_1, AWS_REGION_GOV_CLOUD_US_WEST_1}, + ) == [AWS_REGION_GOV_CLOUD_US_EAST_1, AWS_REGION_GOV_CLOUD_US_WEST_1] + + def test_validate_credentials_falls_back_to_the_next_partition_region( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + # A container may carry a region that belongs to no partition it scans + current_session = session.Session(region_name=AWS_REGION_US_EAST_1) + attempted_regions = [] + + def create_sts_session(session, aws_region): + attempted_regions.append(aws_region) + if aws_region == AWS_REGION_GOV_CLOUD_US_EAST_1: + raise botocore.exceptions.EndpointConnectionError( + endpoint_url=f"https://sts.{aws_region}.amazonaws.com" + ) + sts_client = mock.MagicMock() + sts_client.get_caller_identity.return_value = { + "UserId": "test-user-id", + "Account": AWS_ACCOUNT_NUMBER, + "Arn": AWS_GOV_CLOUD_ACCOUNT_ARN, + } + return sts_client + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ): + caller_identity = AwsProvider.validate_credentials( + session=current_session, aws_region=AWS_REGION_GOV_CLOUD_US_EAST_1 + ) + + assert attempted_regions == [ + AWS_REGION_GOV_CLOUD_US_EAST_1, + AWS_REGION_GOV_CLOUD_US_WEST_1, + ] + assert caller_identity.region == AWS_REGION_GOV_CLOUD_US_WEST_1 + + def test_validate_credentials_falls_back_when_a_region_does_not_answer( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + current_session = session.Session(region_name=AWS_REGION_US_EAST_1) + attempted_regions = [] + + # The connection is accepted but nothing comes back before the read timeout + def create_sts_session(session, aws_region): + attempted_regions.append(aws_region) + if aws_region == AWS_REGION_GOV_CLOUD_US_EAST_1: + raise botocore.exceptions.ReadTimeoutError( + endpoint_url=f"https://sts.{aws_region}.amazonaws.com" + ) + sts_client = mock.MagicMock() + sts_client.get_caller_identity.return_value = { + "UserId": "test-user-id", + "Account": AWS_ACCOUNT_NUMBER, + "Arn": AWS_GOV_CLOUD_ACCOUNT_ARN, + } + return sts_client + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ): + caller_identity = AwsProvider.validate_credentials( + session=current_session, aws_region=AWS_REGION_GOV_CLOUD_US_EAST_1 + ) + + assert attempted_regions == [ + AWS_REGION_GOV_CLOUD_US_EAST_1, + AWS_REGION_GOV_CLOUD_US_WEST_1, + ] + assert caller_identity.region == AWS_REGION_GOV_CLOUD_US_WEST_1 + + def test_validate_credentials_raises_when_no_partition_region_answers( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + current_session = session.Session(region_name=AWS_REGION_US_EAST_1) + attempted_regions = [] + + def create_sts_session(session, aws_region): + attempted_regions.append(aws_region) + raise botocore.exceptions.EndpointConnectionError( + endpoint_url=f"https://sts.{aws_region}.amazonaws.com" + ) + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ): + with raises(botocore.exceptions.EndpointConnectionError): + AwsProvider.validate_credentials( + session=current_session, aws_region=AWS_REGION_GOV_CLOUD_US_EAST_1 + ) + + assert attempted_regions == [ + AWS_REGION_GOV_CLOUD_US_EAST_1, + AWS_REGION_GOV_CLOUD_US_WEST_1, + ] + + def test_validate_credentials_avoids_an_excluded_region_when_failing_over( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_COMMERCIAL_PARTITION) + current_session = session.Session(region_name=AWS_REGION_EU_WEST_1) + partition_regions = get_env_partition_regions(AWS_REGION_EU_WEST_1) + excluded_region, answering_region = partition_regions[1:3] + attempted_regions = [] + + def create_sts_session(session, aws_region): + attempted_regions.append(aws_region) + if aws_region == AWS_REGION_EU_WEST_1: + raise botocore.exceptions.EndpointConnectionError( + endpoint_url=f"https://sts.{aws_region}.amazonaws.com" + ) + sts_client = mock.MagicMock() + sts_client.get_caller_identity.return_value = { + "UserId": "test-user-id", + "Account": AWS_ACCOUNT_NUMBER, + "Arn": AWS_ACCOUNT_ARN, + } + return sts_client + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ): + caller_identity = AwsProvider.validate_credentials( + session=current_session, + aws_region=AWS_REGION_EU_WEST_1, + excluded_regions={excluded_region}, + ) + + assert attempted_regions == [AWS_REGION_EU_WEST_1, answering_region] + assert caller_identity.region == answering_region + + def test_validate_credentials_does_not_retry_a_credential_error(self, monkeypatch): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + current_session = session.Session(region_name=AWS_REGION_US_EAST_1) + attempted_regions = [] + + def create_sts_session(session, aws_region): + attempted_regions.append(aws_region) + sts_client = mock.MagicMock() + sts_client.get_caller_identity.side_effect = ( + botocore.exceptions.ClientError( + {"Error": {"Code": "InvalidClientTokenId", "Message": "invalid"}}, + "GetCallerIdentity", + ) + ) + return sts_client + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ): + with raises(AWSAccessKeyIDInvalidError): + AwsProvider.validate_credentials( + session=current_session, aws_region=AWS_REGION_GOV_CLOUD_US_EAST_1 + ) + + assert attempted_regions == [AWS_REGION_GOV_CLOUD_US_EAST_1] + + def test_assume_role_falls_back_to_the_next_partition_region(self, monkeypatch): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + current_session = session.Session(region_name=AWS_REGION_US_EAST_1) + attempted_regions = [] + + def create_sts_session(session, aws_region): + attempted_regions.append(aws_region) + if aws_region == AWS_REGION_GOV_CLOUD_US_EAST_1: + raise botocore.exceptions.EndpointConnectionError( + endpoint_url=f"https://sts.{aws_region}.amazonaws.com" + ) + sts_client = mock.MagicMock() + sts_client.assume_role.return_value = { + "Credentials": { + "AccessKeyId": "AKIAIOSFODNN7EXAMPLE", + "SecretAccessKey": "secret", + "SessionToken": "token", + "Expiration": datetime.now() + timedelta(seconds=3600), + } + } + return sts_client + + assumed_role_info = AWSAssumeRoleInfo( + role_arn=ARN( + arn=f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/test-role" + ), + session_duration=3600, + external_id=None, + mfa_enabled=False, + role_session_name=ROLE_SESSION_NAME, + sts_region=AWS_REGION_GOV_CLOUD_US_EAST_1, + ) + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ): + credentials = AwsProvider.assume_role(current_session, assumed_role_info) + + assert attempted_regions == [ + AWS_REGION_GOV_CLOUD_US_EAST_1, + AWS_REGION_GOV_CLOUD_US_WEST_1, + ] + assert isinstance(credentials, AWSCredentials) + assert credentials.aws_access_key_id == "AKIAIOSFODNN7EXAMPLE" + + def test_setup_session_mfa_falls_back_to_the_next_partition_region( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + monkeypatch.setenv("AWS_DEFAULT_REGION", AWS_REGION_US_EAST_1) + attempted_regions = [] + + def create_sts_session(session, aws_region): + attempted_regions.append(aws_region) + if aws_region == AWS_REGION_GOV_CLOUD_US_EAST_1: + raise botocore.exceptions.EndpointConnectionError( + endpoint_url=f"https://sts.{aws_region}.amazonaws.com" + ) + sts_client = mock.MagicMock() + sts_client.get_session_token.return_value = { + "Credentials": { + "AccessKeyId": "AKIAIOSFODNN7EXAMPLE", + "SecretAccessKey": "secret", + "SessionToken": "token", + } + } + return sts_client + + with ( + patch( + "prowler.providers.aws.aws_provider.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", + ), + ), + patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ), + ): + mfa_session = AwsProvider.setup_session( + mfa=True, + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + ) + + assert attempted_regions == [ + AWS_REGION_GOV_CLOUD_US_EAST_1, + AWS_REGION_GOV_CLOUD_US_WEST_1, + ] + assert mfa_session.get_credentials().access_key == "AKIAIOSFODNN7EXAMPLE" + assert mfa_session.get_credentials().token == "token" + + @mock_aws + def test_aws_provider_hands_excluded_regions_to_credential_validation( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + monkeypatch.setenv("AWS_DEFAULT_REGION", AWS_REGION_US_EAST_1) + handed = [] + + class Validated(Exception): + pass + + # Stops at the validation: what it was handed is all this checks + def validate_credentials(session, aws_region, excluded_regions=None): + handed.append((aws_region, set(excluded_regions or ()))) + raise Validated + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.validate_credentials", + side_effect=validate_credentials, + ): + with raises(Validated): + AwsProvider(excluded_regions={AWS_REGION_GOV_CLOUD_US_EAST_1}) + + assert handed == [ + (AWS_REGION_GOV_CLOUD_US_WEST_1, {AWS_REGION_GOV_CLOUD_US_EAST_1}) + ] + + @mock_aws + def test_aws_provider_assumes_the_role_where_validation_got_an_answer( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + # Out of the partition, so the first candidate is botocore's, not this one + monkeypatch.setenv("AWS_DEFAULT_REGION", AWS_REGION_US_EAST_1) + role_arn = ( + f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/test-role" + ) + answered = AWSCallerIdentity( + user_id="test-user-id", + account=AWS_ACCOUNT_NUMBER, + arn=ARN(AWS_GOV_CLOUD_ACCOUNT_ARN), + region=AWS_REGION_GOV_CLOUD_US_WEST_1, + ) + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.validate_credentials", + return_value=answered, + ): + aws_provider = AwsProvider(role_arn=role_arn, session_duration=900) + + assert ( + aws_provider._assumed_role_configuration.info.sts_region + == AWS_REGION_GOV_CLOUD_US_WEST_1 + ) + + @mock_aws + def test_aws_provider_assumes_the_organizations_role_where_validation_got_an_answer( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + monkeypatch.setenv("AWS_DEFAULT_REGION", AWS_REGION_US_EAST_1) + organizations_role_arn = f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/organizations-role" + answered = AWSCallerIdentity( + user_id="test-user-id", + account=AWS_ACCOUNT_NUMBER, + arn=ARN(AWS_GOV_CLOUD_ACCOUNT_ARN), + region=AWS_REGION_GOV_CLOUD_US_WEST_1, + ) + sts_regions = [] + + class RoleAssumed(Exception): + pass + + # Stops at the assumption: the region it was handed is all this checks + def assume_role(session, assumed_role_info): + sts_regions.append(assumed_role_info.sts_region) + raise RoleAssumed + + with ( + patch( + "prowler.providers.aws.aws_provider.AwsProvider.validate_credentials", + return_value=answered, + ), + patch( + "prowler.providers.aws.aws_provider.AwsProvider.assume_role", + side_effect=assume_role, + ), + ): + with raises(RoleAssumed): + AwsProvider( + organizations_role_arn=organizations_role_arn, + session_duration=900, + ) + + assert sts_regions == [AWS_REGION_GOV_CLOUD_US_WEST_1] + @mock_aws def test_test_connection_with_env_credentials(self, monkeypatch): # Create a mock IAM user