diff --git a/docs/user-guide/providers/aws/boto3-configuration.mdx b/docs/user-guide/providers/aws/boto3-configuration.mdx index d4e348b2b7..c10db03ef8 100644 --- a/docs/user-guide/providers/aws/boto3-configuration.mdx +++ b/docs/user-guide/providers/aws/boto3-configuration.mdx @@ -29,6 +29,22 @@ Boto3 defaults both timeouts to 60 seconds. In networks with restricted egress ( +## Retries Configuration + + + +The number of retries is set with `--aws-retries-max-attempts`, where `0` disables retries. It can also be set through an environment variable, which is the way to tune it in Prowler Cloud and other deployments without a CLI: + +```console +export PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS=0 +``` + +The CLI flag takes precedence over the environment variable. The value must be a non-negative integer; when neither is set, Prowler uses 3 retries. + + +The environment variable is process-wide: it applies to every AWS provider built in the process where it is set, not only to a connection check. A scan started in that same process picks it up too. Boto3's Standard retry mode, which Prowler uses, also retries service-side throttling responses (see the errors listed below), so `0` disables retries for those as well. On a large account a scan can hit throttling under normal load, and with retries disabled that throttling becomes a hard failure instead of a retried call. Set the variable only on the processes that run connection checks; leave scan workers on the default, or raise their retry count instead of lowering it. + + ## Retry Behavior Overview Boto3's Standard retry mode includes the following mechanisms: diff --git a/prowler/changelog.d/aws-retries-max-attempts-env.added.md b/prowler/changelog.d/aws-retries-max-attempts-env.added.md new file mode 100644 index 0000000000..d88cc59507 --- /dev/null +++ b/prowler/changelog.d/aws-retries-max-attempts-env.added.md @@ -0,0 +1 @@ +`PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS` environment variable to set the Boto3 retries for deployments without CLI flags diff --git a/prowler/changelog.d/aws-sts-reuse-answering-region.fixed.md b/prowler/changelog.d/aws-sts-reuse-answering-region.fixed.md new file mode 100644 index 0000000000..431426522a --- /dev/null +++ b/prowler/changelog.d/aws-sts-reuse-answering-region.fixed.md @@ -0,0 +1 @@ +STS calls after role assumption use the answering region, avoiding a second wait for an unreachable partition region diff --git a/prowler/providers/aws/aws_provider.py b/prowler/providers/aws/aws_provider.py index 784daf0caa..37067bb516 100644 --- a/prowler/providers/aws/aws_provider.py +++ b/prowler/providers/aws/aws_provider.py @@ -112,7 +112,7 @@ class AwsProvider(Provider): def __init__( self, - retries_max_attempts: int = 3, + retries_max_attempts: Optional[int] = None, role_arn: str = None, session_duration: int = 3600, external_id: str = None, @@ -141,6 +141,7 @@ class AwsProvider(Provider): Args: - retries_max_attempts: The maximum number of retries for the AWS client. + Defaults to the PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS environment variable or, if unset, to 3. - role_arn: The ARN of the IAM role to assume. - session_duration: The duration of the session in seconds, between 900 and 43200. - external_id: The external ID to use when assuming the IAM role. @@ -1230,7 +1231,8 @@ class AwsProvider(Provider): Args: - session: The AWS session object - - assumed_role_info: The AWSAssumeRoleInfo object + - assumed_role_info: The AWSAssumeRoleInfo object. Its sts_region is + updated to the region that answered, so later calls go straight there Returns: - AWSCredentials: The AWS credentials for the assumed role @@ -1256,11 +1258,14 @@ 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 - _, assumed_credentials = AwsProvider.sts_call_with_partition_failover( - session, - assumed_role_info.sts_region, - lambda sts_client: sts_client.assume_role(**assume_role_arguments), + 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_role_info.sts_region = sts_region # Convert the UTC datetime object to your local timezone credentials_expiration_local_time = ( assumed_credentials["Credentials"]["Expiration"] @@ -1558,6 +1563,8 @@ class AwsProvider(Provider): session, assumed_role_information, ) + # Validate where the role was assumed, not where it timed out + aws_region = assumed_role_information.sts_region session = Session( aws_access_key_id=assumed_role_credentials.aws_access_key_id, aws_secret_access_key=assumed_role_credentials.aws_secret_access_key, diff --git a/prowler/providers/aws/config.py b/prowler/providers/aws/config.py index ed2ca503d0..dd365bb8cf 100644 --- a/prowler/providers/aws/config.py +++ b/prowler/providers/aws/config.py @@ -2,7 +2,10 @@ import os from botocore.config import Config -from prowler.providers.aws.exceptions.exceptions import AWSInvalidBoto3TimeoutError +from prowler.providers.aws.exceptions.exceptions import ( + AWSInvalidBoto3RetriesError, + AWSInvalidBoto3TimeoutError, +) AWS_STS_GLOBAL_ENDPOINT_REGION = "us-east-1" AWS_REGION_US_EAST_1 = "us-east-1" @@ -27,10 +30,28 @@ def get_boto3_timeout_from_env(name: str, default: int) -> int: return int(raw) +def get_boto3_retries_from_env(name: str, default: int) -> int: + """Non-negative integer retries read from the environment, or default when unset.""" + raw = os.getenv(name, "").strip() + if not raw: + return default + if not raw.isdecimal(): + raise AWSInvalidBoto3RetriesError( + file=os.path.basename(__file__), + message=f"{name} must be a non-negative integer number of retries, got {raw!r}", + ) + return int(raw) + + def get_default_session_config() -> Config: return Config( user_agent_extra=BOTO3_USER_AGENT_EXTRA, - retries={"max_attempts": BOTO3_RETRIES_MAX_ATTEMPTS, "mode": "standard"}, + retries={ + "max_attempts": get_boto3_retries_from_env( + "PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS", BOTO3_RETRIES_MAX_ATTEMPTS + ), + "mode": "standard", + }, connect_timeout=get_boto3_timeout_from_env( "PROWLER_AWS_BOTO3_CONNECT_TIMEOUT", BOTO3_CONNECT_TIMEOUT ), diff --git a/prowler/providers/aws/exceptions/exceptions.py b/prowler/providers/aws/exceptions/exceptions.py index 089e0d99c7..dae23408e9 100644 --- a/prowler/providers/aws/exceptions/exceptions.py +++ b/prowler/providers/aws/exceptions/exceptions.py @@ -82,6 +82,10 @@ class AWSBaseException(ProwlerException): "message": "The Boto3 timeout configured through the environment is invalid", "remediation": "Set PROWLER_AWS_BOTO3_CONNECT_TIMEOUT and PROWLER_AWS_BOTO3_READ_TIMEOUT to a positive integer number of seconds.", }, + (1919, "AWSInvalidBoto3RetriesError"): { + "message": "The Boto3 retries configured through the environment are invalid", + "remediation": "Set PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS to a non-negative integer, 0 disables retries.", + }, } def __init__(self, code, file=None, original_exception=None, message=None): @@ -244,3 +248,12 @@ class AWSInvalidBoto3TimeoutError(AWSBaseException): super().__init__( 1918, file=file, original_exception=original_exception, message=message ) + + +class AWSInvalidBoto3RetriesError(AWSBaseException): + """Boto3 retries configured through the environment are not a non-negative integer.""" + + def __init__(self, file=None, original_exception=None, message=None): + super().__init__( + 1919, file=file, original_exception=original_exception, message=message + ) diff --git a/prowler/providers/aws/lib/s3/s3.py b/prowler/providers/aws/lib/s3/s3.py index a4bbb42cc7..d2f511fd09 100644 --- a/prowler/providers/aws/lib/s3/s3.py +++ b/prowler/providers/aws/lib/s3/s3.py @@ -85,7 +85,7 @@ class S3: aws_access_key_id: str = None, aws_secret_access_key: str = None, aws_session_token: Optional[str] = None, - retries_max_attempts: int = 3, + retries_max_attempts: Optional[int] = None, regions: set = set(), ) -> None: """ diff --git a/prowler/providers/aws/lib/security_hub/security_hub.py b/prowler/providers/aws/lib/security_hub/security_hub.py index bc372d1ddd..ae4c2350e4 100644 --- a/prowler/providers/aws/lib/security_hub/security_hub.py +++ b/prowler/providers/aws/lib/security_hub/security_hub.py @@ -106,7 +106,7 @@ class SecurityHub: aws_access_key_id: str = None, aws_secret_access_key: str = None, aws_session_token: Optional[str] = None, - retries_max_attempts: int = 3, + retries_max_attempts: Optional[int] = None, regions: set = set(), ) -> "SecurityHub": """ diff --git a/prowler/providers/aws/lib/session/aws_set_up_session.py b/prowler/providers/aws/lib/session/aws_set_up_session.py index 8f0b4130ca..6b00bb437b 100644 --- a/prowler/providers/aws/lib/session/aws_set_up_session.py +++ b/prowler/providers/aws/lib/session/aws_set_up_session.py @@ -40,7 +40,7 @@ class AwsSetUpSession: aws_access_key_id: str = None, aws_secret_access_key: str = None, aws_session_token: Optional[str] = None, - retries_max_attempts: int = 3, + retries_max_attempts: Optional[int] = None, regions: set = set(), connect_timeout: Optional[int] = None, read_timeout: Optional[int] = None, @@ -106,6 +106,8 @@ class AwsSetUpSession: session=self._session.current_session, aws_region=sts_region, ) + # Later STS calls go where validation got an answer, not where it timed out + sts_region = caller_identity.region logger.info("Credentials validated") ######## diff --git a/tests/providers/aws/aws_provider_test.py b/tests/providers/aws/aws_provider_test.py index 2d759f362a..2cf9d828b5 100644 --- a/tests/providers/aws/aws_provider_test.py +++ b/tests/providers/aws/aws_provider_test.py @@ -28,15 +28,19 @@ from prowler.providers.aws.config import ( AWS_STS_GLOBAL_ENDPOINT_REGION, BOTO3_CONNECT_TIMEOUT, BOTO3_READ_TIMEOUT, + BOTO3_RETRIES_MAX_ATTEMPTS, BOTO3_USER_AGENT_EXTRA, ROLE_SESSION_NAME, + get_boto3_retries_from_env, get_boto3_timeout_from_env, get_default_session_config, ) from prowler.providers.aws.exceptions.exceptions import ( AWSAccessKeyIDInvalidError, AWSArgumentTypeValidationError, + AWSAssumeRoleError, AWSIAMRoleARNInvalidResourceTypeError, + AWSInvalidBoto3RetriesError, AWSInvalidBoto3TimeoutError, AWSInvalidPartitionError, AWSInvalidProviderIdError, @@ -1845,6 +1849,143 @@ aws: ] assert isinstance(credentials, AWSCredentials) assert credentials.aws_access_key_id == "AKIAIOSFODNN7EXAMPLE" + # Refreshing the credentials later goes straight to the region that answered + assert assumed_role_info.sts_region == AWS_REGION_GOV_CLOUD_US_WEST_1 + + def test_assume_role_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.assume_role.side_effect = botocore.exceptions.ClientError( + {"Error": {"Code": "AccessDenied", "Message": "denied"}}, + "AssumeRole", + ) + 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, + ): + with raises(AWSAssumeRoleError): + AwsProvider.assume_role(current_session, assumed_role_info) + + assert attempted_regions == [AWS_REGION_GOV_CLOUD_US_EAST_1] + assert assumed_role_info.sts_region == AWS_REGION_GOV_CLOUD_US_EAST_1 + + def test_test_connection_role_validates_where_the_role_was_assumed( + self, monkeypatch + ): + monkeypatch.setenv("PROWLER_AWS_PARTITION", AWS_GOV_CLOUD_PARTITION) + monkeypatch.delenv("AWS_DEFAULT_REGION", raising=False) + attempted_calls = [] + + def create_sts_session(session, aws_region): + if aws_region == AWS_REGION_GOV_CLOUD_US_EAST_1: + attempted_calls.append(aws_region) + raise botocore.exceptions.ConnectTimeoutError( + endpoint_url=f"https://sts.{aws_region}.amazonaws.com" + ) + sts_client = mock.MagicMock() + + def assume_role(**_): + attempted_calls.append(("AssumeRole", aws_region)) + return { + "Credentials": { + "AccessKeyId": "AKIAIOSFODNN7EXAMPLE", + "SecretAccessKey": "secret", + "SessionToken": "token", + "Expiration": datetime.now() + timedelta(seconds=3600), + } + } + + def get_caller_identity(): + attempted_calls.append(("GetCallerIdentity", aws_region)) + return { + "UserId": "test-user-id", + "Account": AWS_ACCOUNT_NUMBER, + "Arn": AWS_GOV_CLOUD_ACCOUNT_ARN, + } + + sts_client.assume_role.side_effect = assume_role + sts_client.get_caller_identity.side_effect = get_caller_identity + return sts_client + + with patch( + "prowler.providers.aws.aws_provider.AwsProvider.create_sts_session", + new=create_sts_session, + ): + 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 + # The unreachable region is paid for once, not again for the validation + assert attempted_calls == [ + AWS_REGION_GOV_CLOUD_US_EAST_1, + ("AssumeRole", AWS_REGION_GOV_CLOUD_US_WEST_1), + ("GetCallerIdentity", AWS_REGION_GOV_CLOUD_US_WEST_1), + ] + + @mock_aws + def test_aws_set_up_session_assumes_the_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) + 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 + + 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): + AwsSetUpSession( + role_arn=f"arn:{AWS_GOV_CLOUD_PARTITION}:iam::{AWS_ACCOUNT_NUMBER}:role/test-role", + session_duration=900, + external_id="test-external-id", + role_session_name=ROLE_SESSION_NAME, + aws_access_key_id="testing", + aws_secret_access_key="testing", + ) + + assert sts_regions == [AWS_REGION_GOV_CLOUD_US_WEST_1] def test_setup_session_mfa_falls_back_to_the_next_partition_region( self, monkeypatch @@ -3352,6 +3493,123 @@ aws: ): get_boto3_timeout_from_env("PROWLER_AWS_BOTO3_CONNECT_TIMEOUT", 10) + def test_get_default_session_config_retries_from_env(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "1"} + ): + config = get_default_session_config() + + assert config.retries == {"max_attempts": 1, "mode": "standard"} + + def test_get_default_session_config_retries_from_env_0_disables_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "0"} + ): + config = get_default_session_config() + + assert config.retries == {"max_attempts": 0, "mode": "standard"} + + def test_set_session_config_argument_overrides_env_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "1"} + ): + config = AwsProvider.set_session_config(5) + + assert config.retries == {"max_attempts": 5, "mode": "standard"} + + @mock_aws + def test_aws_provider_without_retries_argument_uses_env_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "0"} + ): + aws_provider = AwsProvider() + client = aws_provider.session.current_session.client( + "ec2", region_name=AWS_REGION_US_EAST_1 + ) + + # botocore rewrites max_attempts into total_max_attempts (retries + 1) + assert client.meta.config.retries["total_max_attempts"] == 1 + + @mock_aws + def test_aws_provider_retries_argument_overrides_env_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "0"} + ): + aws_provider = AwsProvider(retries_max_attempts=7) + client = aws_provider.session.current_session.client( + "ec2", region_name=AWS_REGION_US_EAST_1 + ) + + assert client.meta.config.retries["total_max_attempts"] == 8 + + @mock_aws + def test_aws_set_up_session_without_retries_argument_uses_env_retries(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "1"} + ): + aws_session = AwsSetUpSession( + aws_access_key_id="testing", + aws_secret_access_key="testing", + ) + client = aws_session._session.current_session.client( + "ec2", region_name=AWS_REGION_US_EAST_1 + ) + + assert client.meta.config.retries["total_max_attempts"] == 2 + + def test_test_connection_session_uses_env_retries(self): + with ( + mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": "0"} + ), + 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, + ), + ) 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 + validated_session = mock_validate_credentials.call_args.args[0] + assert validated_session._session.get_default_client_config().retries == { + "max_attempts": 0, + "mode": "standard", + } + + @pytest.mark.parametrize("raw", ["-1", "three", "1.5"]) + def test_get_boto3_retries_from_env_rejects_anything_but_non_negative_integers( + self, raw + ): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": raw} + ): + with raises( + AWSInvalidBoto3RetriesError, + match="PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS", + ): + get_boto3_retries_from_env("PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS", 3) + + def test_get_boto3_retries_from_env_blank_falls_back_to_default(self): + with mock.patch.dict( + os.environ, {"PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS": " "} + ): + assert ( + get_boto3_retries_from_env( + "PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS", BOTO3_RETRIES_MAX_ATTEMPTS + ) + == BOTO3_RETRIES_MAX_ATTEMPTS + ) + def test_get_boto3_timeout_from_env_blank_falls_back_to_default(self): with mock.patch.dict(os.environ, {"PROWLER_AWS_BOTO3_CONNECT_TIMEOUT": " "}): assert (