mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-10-04 02:04:06 +00:00
feat(sdk): AWS partition selection via PROWLER_AWS_PARTITION (#12680)
Co-authored-by: David <david.copo@gmail.com>
This commit is contained in:
@@ -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
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user