mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-10-04 02:04:06 +00:00
feat(test_connection): Add optional AWS Account ID validation (#5361)
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user