feat(test_connection): Add optional AWS Account ID validation (#5361)

This commit is contained in:
Pepe Fagoaga
2024-10-10 12:45:16 -04:00
committed by GitHub
parent cad99c5e0f
commit d9c2933dc5
4 changed files with 96 additions and 2 deletions
+4 -1
View File
@@ -45,7 +45,10 @@ class ProwlerException(Exception):
def __str__(self):
"""Overriding the __str__ method"""
return f"{self.__class__.__name__}[{self.code}]: {self.message} - {self.original_exception}"
default_str = f"{self.__class__.__name__}[{self.code}]: {self.message}"
if self.original_exception:
default_str += f" - {self.original_exception}"
return default_str
class UnexpectedError(ProwlerException):
+16 -1
View File
@@ -34,6 +34,7 @@ from prowler.providers.aws.exceptions.exceptions import (
AWSIAMRoleARNPartitionEmpty,
AWSIAMRoleARNRegionNotEmtpy,
AWSIAMRoleARNServiceNotIAMnorSTS,
AWSInvalidAccountCredentials,
AWSNoCredentialsError,
AWSProfileNotFoundError,
AWSSecretAccessKeyInvalid,
@@ -1005,6 +1006,7 @@ class AwsProvider(Provider):
aws_access_key_id: str = None,
aws_secret_access_key: str = None,
aws_session_token: Optional[str] = None,
provider_id: Optional[str] = None,
) -> Connection:
"""
Test the connection to AWS with one of the Boto3 credentials methods.
@@ -1021,6 +1023,7 @@ class AwsProvider(Provider):
aws_access_key_id (str): The AWS access key ID to use for the session.
aws_secret_access_key (str): The AWS secret access key to use for the session.
aws_session_token (str): The AWS session token to use for the session. Optional.
provider_id (str): The AWS account ID to validate that the provided credentials belongs to it.
Returns:
Connection: An object tha contains the result of the test connection operation.
@@ -1049,6 +1052,8 @@ class AwsProvider(Provider):
Connection(is_connected=False, Error=NoCredentialsError('Unable to locate credentials'))
>>> AwsProvider.test_connection(aws_access_key_id="XXXXXXXX", aws_secret_access_key="XXXXXXXX", raise_on_exception=False))
Connection(is_connected=True, Error=None))
>>> AwsProvider.test_connection(aws_access_key_id="XXXXXXXX", aws_secret_access_key="XXXXXXXX", provider_id="111122223333", raise_on_exception=False))
Connection(is_connected=True, Error=None))
"""
try:
session = AwsProvider.setup_session(
@@ -1082,7 +1087,11 @@ class AwsProvider(Provider):
profile_name=profile,
)
_ = AwsProvider.validate_credentials(session, aws_region)
caller_identity = AwsProvider.validate_credentials(session, aws_region)
# Do an extra validation if the AWS account ID is provided
if provider_id and caller_identity.account != provider_id:
raise AWSInvalidAccountCredentials(file=pathlib.Path(__file__).name)
return Connection(
is_connected=True,
)
@@ -1185,6 +1194,12 @@ class AwsProvider(Provider):
raise secret_access_key_invalid_error
return Connection(error=secret_access_key_invalid_error)
except AWSInvalidAccountCredentials as invalid_account_credentials_error:
logger.error(str(invalid_account_credentials_error))
if raise_on_exception:
raise invalid_account_credentials_error
return Connection(error=invalid_account_credentials_error)
except Exception as error:
logger.critical(
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
@@ -65,6 +65,10 @@ class AWSBaseException(ProwlerException):
"message": "AWS Secret Access Key is invalid",
"remediation": "Check your AWS Secret Access Key and signing method and ensure it is valid.",
},
(1917, "AWSInvalidAccountCredentials"): {
"message": "The provided AWS credentials belong to a different account",
"remediation": "Check the provided AWS credentials and review if belong to the account you want to use.",
},
}
def __init__(self, code, file=None, original_exception=None, message=None):
@@ -197,3 +201,10 @@ class AWSSecretAccessKeyInvalid(AWSCredentialsError):
super().__init__(
1916, file=file, original_exception=original_exception, message=message
)
class AWSInvalidAccountCredentials(AWSCredentialsError):
def __init__(self, file=None, original_exception=None, message=None):
super().__init__(
1917, file=file, original_exception=original_exception, message=message
)
+65
View File
@@ -29,6 +29,7 @@ from prowler.providers.aws.config import (
from prowler.providers.aws.exceptions.exceptions import (
AWSArgumentTypeValidationError,
AWSIAMRoleARNInvalidResourceType,
AWSInvalidAccountCredentials,
AWSNoCredentialsError,
)
from prowler.providers.aws.lib.arn.models import ARN
@@ -1398,6 +1399,70 @@ aws:
assert connection.is_connected
assert connection.error is None
@mock_aws
def test_test_connection_with_own_account(self):
sts_client = client("sts", region_name=AWS_REGION_EU_WEST_1)
session_token = sts_client.get_session_token()
session_credentials = {
"aws_access_key_id": session_token["Credentials"]["AccessKeyId"],
"aws_secret_access_key": session_token["Credentials"]["SecretAccessKey"],
"aws_session_token": session_token["Credentials"]["SessionToken"],
"provider_id": AWS_ACCOUNT_NUMBER,
}
connection = AwsProvider.test_connection(**session_credentials)
assert isinstance(connection, Connection)
assert connection.is_connected
assert connection.error is None
@mock_aws
def test_test_connection_with_different_account(self):
sts_client = client("sts", region_name=AWS_REGION_EU_WEST_1)
session_token = sts_client.get_session_token()
session_credentials = {
"aws_access_key_id": session_token["Credentials"]["AccessKeyId"],
"aws_secret_access_key": session_token["Credentials"]["SecretAccessKey"],
"aws_session_token": session_token["Credentials"]["SessionToken"],
"provider_id": "111122223333",
}
with raises(AWSInvalidAccountCredentials) as exception:
AwsProvider.test_connection(**session_credentials)
assert exception.type == AWSInvalidAccountCredentials
assert (
exception.value.args[0]
== "[1917] The provided AWS credentials belong to a different account"
)
@mock_aws
def test_test_connection_with_different_account_dont_raise(self):
sts_client = client("sts", region_name=AWS_REGION_EU_WEST_1)
session_token = sts_client.get_session_token()
session_credentials = {
"aws_access_key_id": session_token["Credentials"]["AccessKeyId"],
"aws_secret_access_key": session_token["Credentials"]["SecretAccessKey"],
"aws_session_token": session_token["Credentials"]["SessionToken"],
"provider_id": "111122223333",
}
connection = AwsProvider.test_connection(
**session_credentials, raise_on_exception=False
)
assert isinstance(connection, Connection)
assert not connection.is_connected
assert isinstance(connection.error, AWSInvalidAccountCredentials)
assert (
connection.error.message
== "The provided AWS credentials belong to a different account"
)
assert connection.error.code == 1917
@mock_aws
def test_create_sts_session(self):
current_session = session.Session()