fix(aws): reuse the STS region that answered and add PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS (#12870)

This commit is contained in:
César Arroba
2026-09-24 10:24:44 +02:00
committed by GitHub
parent 706603fe4d
commit 15630f54d2
10 changed files with 330 additions and 11 deletions
@@ -29,6 +29,22 @@ Boto3 defaults both timeouts to 60 seconds. In networks with restricted egress (
</Note>
## Retries Configuration
<VersionBadge version="5.44.0" />
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.
<Warning>
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.
</Warning>
## Retry Behavior Overview
Boto3's Standard retry mode includes the following mechanisms:
@@ -0,0 +1 @@
`PROWLER_AWS_BOTO3_RETRIES_MAX_ATTEMPTS` environment variable to set the Boto3 retries for deployments without CLI flags
@@ -0,0 +1 @@
STS calls after role assumption use the answering region, avoiding a second wait for an unreachable partition region
+10 -3
View File
@@ -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(
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,
+23 -2
View File
@@ -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
),
@@ -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
)
+1 -1
View File
@@ -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:
"""
@@ -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":
"""
@@ -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")
########
+258
View File
@@ -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 (