diff --git a/prowler/exceptions/exceptions.py b/prowler/exceptions/exceptions.py index 7ffadab946..352733afc9 100644 --- a/prowler/exceptions/exceptions.py +++ b/prowler/exceptions/exceptions.py @@ -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): diff --git a/prowler/providers/aws/aws_provider.py b/prowler/providers/aws/aws_provider.py index 6bdf5f4960..1a55a5d452 100644 --- a/prowler/providers/aws/aws_provider.py +++ b/prowler/providers/aws/aws_provider.py @@ -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}" diff --git a/prowler/providers/aws/exceptions/exceptions.py b/prowler/providers/aws/exceptions/exceptions.py index 2a4866555c..579dd88e97 100644 --- a/prowler/providers/aws/exceptions/exceptions.py +++ b/prowler/providers/aws/exceptions/exceptions.py @@ -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 + ) diff --git a/tests/providers/aws/aws_provider_test.py b/tests/providers/aws/aws_provider_test.py index 275b2c8c4a..974bba9be9 100644 --- a/tests/providers/aws/aws_provider_test.py +++ b/tests/providers/aws/aws_provider_test.py @@ -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()