fix(aws): try the rest of the partition when the bootstrap region is unreachable (#12799)

Co-authored-by: pedrooot <pedromarting3@gmail.com>
This commit is contained in:
César Arroba
2026-09-16 09:24:41 +02:00
committed by GitHub
co-authored by pedrooot
parent c0fdd5bdf3
commit 757cd44ecb
4 changed files with 548 additions and 13 deletions
@@ -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.
<Note>
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.
</Note>
@@ -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
+137 -13
View File
@@ -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,
+408
View File
@@ -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