feat(sdk): AWS partition selection via PROWLER_AWS_PARTITION (#12680)

Co-authored-by: David <david.copo@gmail.com>
This commit is contained in:
Pedro Martín
2026-08-31 18:13:30 +02:00
committed by GitHub
co-authored by David
parent 587c47bfe2
commit 6422178b76
3 changed files with 380 additions and 11 deletions
@@ -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
+123 -11
View File
@@ -1,6 +1,7 @@
import os import os
import pathlib import pathlib
from datetime import datetime from datetime import datetime
from functools import lru_cache
from re import fullmatch from re import fullmatch
from typing import Optional from typing import Optional
@@ -671,7 +672,12 @@ class AwsProvider(Provider):
if mfa: if mfa:
session = Session(**session_arguments) session = Session(**session_arguments)
session._session.set_default_client_config(session_config) 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 # TODO: pass values from the input
mfa_info = AwsProvider.input_role_mfa_token_and_code() mfa_info = AwsProvider.input_role_mfa_token_and_code()
@@ -1352,7 +1358,7 @@ class AwsProvider(Provider):
@staticmethod @staticmethod
def test_connection( def test_connection(
profile: str = None, profile: str = None,
aws_region: str = AWS_STS_GLOBAL_ENDPOINT_REGION, aws_region: str = None,
role_arn: str = None, role_arn: str = None,
role_session_name: str = ROLE_SESSION_NAME, role_session_name: str = ROLE_SESSION_NAME,
session_duration: int = 3600, session_duration: int = 3600,
@@ -1369,7 +1375,9 @@ class AwsProvider(Provider):
Args: Args:
profile (str): The AWS profile to use for the session. 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_arn (str): The ARN of the IAM role to assume.
role_session_name (str): The name of the role session. role_session_name (str): The name of the role session.
session_duration (int): The duration of the assumed role session in seconds. 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)) Connection(is_connected=True, Error=None))
""" """
try: try:
if aws_region is None:
aws_region = (
get_env_partition_bootstrap_region()
or AWS_STS_GLOBAL_ENDPOINT_REGION
)
session = AwsProvider.setup_session( session = AwsProvider.setup_session(
mfa=mfa_enabled, mfa=mfa_enabled,
profile=profile, profile=profile,
@@ -1430,6 +1444,7 @@ class AwsProvider(Provider):
external_id=external_id, external_id=external_id,
mfa_enabled=mfa_enabled, mfa_enabled=mfa_enabled,
role_session_name=role_session_name, role_session_name=role_session_name,
sts_region=aws_region,
) )
assumed_role_credentials = AwsProvider.assume_role( assumed_role_credentials = AwsProvider.assume_role(
session, session,
@@ -1451,6 +1466,13 @@ class AwsProvider(Provider):
if provider_id and caller_identity.account != provider_id: if provider_id and caller_identity.account != provider_id:
raise AWSInvalidProviderIdError(file=pathlib.Path(__file__).name) 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( return Connection(
is_connected=True, is_connected=True,
) )
@@ -1591,6 +1613,14 @@ class AwsProvider(Provider):
raise session_token_expired raise session_token_expired
return Connection(error=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: except Exception as error:
logger.critical( logger.critical(
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" 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') sts_client = create_sts_session(session, 'us-west-2')
""" """
try: try:
if os.environ.get("AWS_ENDPOINT_URL"): # Botocore resolves the regional STS endpoint for every partition
sts_endpoint_url = os.environ["AWS_ENDPOINT_URL"] # (China, EUSC, GovCloud, ISO); AWS_ENDPOINT_URL overrides it
elif aws_region.startswith("cn-"): sts_endpoint_url = os.environ.get("AWS_ENDPOINT_URL") or None
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"
return session.client("sts", aws_region, endpoint_url=sts_endpoint_url) return session.client("sts", aws_region, endpoint_url=sts_endpoint_url)
except Exception as error: except Exception as error:
logger.critical( logger.critical(
@@ -1702,6 +1727,80 @@ def read_aws_regions_file() -> dict:
return data 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 # TODO: This can be moved to another class since it doesn't need self
def get_aws_region_for_sts( def get_aws_region_for_sts(
session_region: str, session_region: str,
@@ -1711,6 +1810,10 @@ def get_aws_region_for_sts(
""" """
Get the AWS region for the STS Assume Role operation. 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: Args:
- session_region (str): The region configured in the AWS session. - session_region (str): The region configured in the AWS session.
- regions (set[str]): The regions passed with the -f/--region/--filter-region option. - 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: if region not in excluded_regions:
return region 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: if session_region and session_region not in excluded_regions:
return session_region return session_region
+256
View File
@@ -1783,6 +1783,29 @@ aws:
assert sts_session._endpoint._endpoint_prefix == "sts" assert sts_session._endpoint._endpoint_prefix == "sts"
assert sts_session._endpoint.host == f"https://sts.{aws_region}.amazonaws.eu" 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 @mock_aws
@patch( @patch(
"prowler.lib.check.utils.recover_checks_from_provider", "prowler.lib.check.utils.recover_checks_from_provider",
@@ -2219,6 +2242,239 @@ aws:
== AWS_REGION_US_EAST_1 == 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): def test_get_profile_region_avoids_excluded_session_region(self):
mocked_session = mock.Mock(region_name=AWS_REGION_EU_WEST_1) mocked_session = mock.Mock(region_name=AWS_REGION_EU_WEST_1)