From 91b74822e9ffc758a2bc12ca059f16a57771f72f Mon Sep 17 00:00:00 2001 From: Prowler Bot Date: Wed, 12 Feb 2025 17:08:47 +0100 Subject: [PATCH] fix(kms): Amazon KMS API call error handling (#6904) Co-authored-by: Ogonna Iwunze <1915636+wunzeco@users.noreply.github.com> --- .../kms_key_not_publicly_accessible.py | 4 +- .../providers/aws/services/kms/kms_service.py | 82 +++++++---- .../kms_cmk_are_used/kms_cmk_are_used_test.py | 61 ++++++++ .../kms_cmk_rotation_enabled_test.py | 61 ++++++++ .../kms_key_not_publicly_accessible_test.py | 133 ++++++++++++++---- 5 files changed, 282 insertions(+), 59 deletions(-) diff --git a/prowler/providers/aws/services/kms/kms_key_not_publicly_accessible/kms_key_not_publicly_accessible.py b/prowler/providers/aws/services/kms/kms_key_not_publicly_accessible/kms_key_not_publicly_accessible.py index 843d22b3c9..408034647a 100644 --- a/prowler/providers/aws/services/kms/kms_key_not_publicly_accessible/kms_key_not_publicly_accessible.py +++ b/prowler/providers/aws/services/kms/kms_key_not_publicly_accessible/kms_key_not_publicly_accessible.py @@ -8,7 +8,9 @@ class kms_key_not_publicly_accessible(Check): findings = [] for key in kms_client.keys: if ( - key.manager == "CUSTOMER" and key.state == "Enabled" + key.manager == "CUSTOMER" + and key.state == "Enabled" + and key.policy is not None ): # only customer KMS have policies report = Check_Report_AWS(metadata=self.metadata(), resource=key) report.status = "PASS" diff --git a/prowler/providers/aws/services/kms/kms_service.py b/prowler/providers/aws/services/kms/kms_service.py index a4ef49ec57..1dc4988842 100644 --- a/prowler/providers/aws/services/kms/kms_service.py +++ b/prowler/providers/aws/services/kms/kms_service.py @@ -26,15 +26,20 @@ class KMS(AWSService): list_keys_paginator = regional_client.get_paginator("list_keys") for page in list_keys_paginator.paginate(): for key in page["Keys"]: - if not self.audit_resources or ( - is_resource_filtered(key["KeyArn"], self.audit_resources) - ): - self.keys.append( - Key( - id=key["KeyId"], - arn=key["KeyArn"], - region=regional_client.region, + try: + if not self.audit_resources or ( + is_resource_filtered(key["KeyArn"], self.audit_resources) + ): + self.keys.append( + Key( + id=key["KeyId"], + arn=key["KeyArn"], + region=regional_client.region, + ) ) + except Exception as error: + logger.error( + f"{regional_client.region} -- {error.__class__.__name__}:{error.__traceback__.tb_lineno} -- {error}" ) except Exception as error: logger.error( @@ -45,8 +50,8 @@ class KMS(AWSService): logger.info("KMS - Describing Key...") try: for key in self.keys: + regional_client = self.regional_clients[key.region] try: - regional_client = self.regional_clients[key.region] response = regional_client.describe_key(KeyId=key.id) key.state = response["KeyMetadata"]["KeyState"] key.origin = response["KeyMetadata"]["Origin"] @@ -73,9 +78,14 @@ class KMS(AWSService): and "AWS" not in key.manager ): regional_client = self.regional_clients[key.region] - key.rotation_enabled = regional_client.get_key_rotation_status( - KeyId=key.id - )["KeyRotationEnabled"] + try: + key.rotation_enabled = regional_client.get_key_rotation_status( + KeyId=key.id + )["KeyRotationEnabled"] + except Exception as error: + logger.error( + f"{regional_client.region} -- {error.__class__.__name__}:{error.__traceback__.tb_lineno} -- {error}" + ) except Exception as error: logger.error( f"{regional_client.region} -- {error.__class__.__name__}:{error.__traceback__.tb_lineno} -- {error}" @@ -89,11 +99,16 @@ class KMS(AWSService): key.manager and key.manager == "CUSTOMER" ): # only customer KMS have policies regional_client = self.regional_clients[key.region] - key.policy = json.loads( - regional_client.get_key_policy( - KeyId=key.id, PolicyName="default" - )["Policy"] - ) + try: + key.policy = json.loads( + regional_client.get_key_policy( + KeyId=key.id, PolicyName="default" + )["Policy"] + ) + except Exception as error: + logger.error( + f"{regional_client.region} -- {error.__class__.__name__}:{error.__traceback__.tb_lineno} -- {error}" + ) except Exception as error: logger.error( f"{regional_client.region} -- {error.__class__.__name__}:{error.__traceback__.tb_lineno} -- {error}" @@ -101,20 +116,25 @@ class KMS(AWSService): def _list_resource_tags(self): logger.info("KMS - List Tags...") - for key in self.keys: - if ( - key.manager and key.manager == "CUSTOMER" - ): # only check customer KMS keys - try: - regional_client = self.regional_clients[key.region] - response = regional_client.list_resource_tags( - KeyId=key.id, - )["Tags"] - key.tags = response - except Exception as error: - logger.error( - f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" - ) + try: + for key in self.keys: + if ( + key.manager and key.manager == "CUSTOMER" + ): # only check customer KMS keys + try: + regional_client = self.regional_clients[key.region] + response = regional_client.list_resource_tags( + KeyId=key.id, + )["Tags"] + key.tags = response + except Exception as error: + logger.error( + f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" + ) + except Exception as error: + logger.error( + f"{regional_client.region} -- {error.__class__.__name__}:{error.__traceback__.tb_lineno} -- {error}" + ) class Key(BaseModel): diff --git a/tests/providers/aws/services/kms/kms_cmk_are_used/kms_cmk_are_used_test.py b/tests/providers/aws/services/kms/kms_cmk_are_used/kms_cmk_are_used_test.py index 7bf5fd5b07..2c8e162609 100644 --- a/tests/providers/aws/services/kms/kms_cmk_are_used/kms_cmk_are_used_test.py +++ b/tests/providers/aws/services/kms/kms_cmk_are_used/kms_cmk_are_used_test.py @@ -1,7 +1,9 @@ +from typing import Any, List from unittest import mock from boto3 import client from moto import mock_aws +import pytest from tests.providers.aws.utils import AWS_REGION_US_EAST_1, set_mocked_aws_provider @@ -62,6 +64,65 @@ class Test_kms_cmk_are_used: assert result[0].resource_id == key["KeyId"] assert result[0].resource_arn == key["Arn"] + @pytest.mark.parametrize( + "no_of_keys_created,expected_no_of_results", + [ + (5, 3), + (7, 5), + (10, 8), + ], + ) + @mock_aws + def test_kms_cmk_are_used_when_describe_key_fails_on_2_keys_out_of_x_keys( + self, no_of_keys_created: int, expected_no_of_results: int + ) -> None: + # Generate KMS Client + kms_client = client("kms", region_name=AWS_REGION_US_EAST_1) + kms_client.__dict__["region"] = AWS_REGION_US_EAST_1 + # Create enabled KMS key + for i in range(no_of_keys_created): + kms_client.create_key( + Tags=[ + {"TagKey": "test", "TagValue": f"test{i}"}, + ], + ) + + orig_describe_key = kms_client.describe_key + + def mock_describe_key(KeyId: str, count: List[int] = [0]) -> Any: + if count[0] in [2, 4]: + count[0] += 1 + raise Exception("FakeClientError") + else: + count[0] += 1 + return orig_describe_key(KeyId=KeyId) + + kms_client.describe_key = mock_describe_key + + from prowler.providers.aws.services.kms.kms_service import KMS + + aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) + + with mock.patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.aws_provider.AwsProvider.generate_regional_clients", + return_value={AWS_REGION_US_EAST_1: kms_client}, + ), mock.patch( + "prowler.providers.aws.services.kms.kms_cmk_are_used.kms_cmk_are_used.kms_client", + new=KMS(aws_provider), + ): + # Test Check + from prowler.providers.aws.services.kms.kms_cmk_are_used.kms_cmk_are_used import ( + kms_cmk_are_used, + ) + + check = kms_cmk_are_used() + result = check.execute() + + assert len(result) == expected_no_of_results + @mock_aws def test_kms_key_with_deletion(self): # Generate KMS Client diff --git a/tests/providers/aws/services/kms/kms_cmk_rotation_enabled/kms_cmk_rotation_enabled_test.py b/tests/providers/aws/services/kms/kms_cmk_rotation_enabled/kms_cmk_rotation_enabled_test.py index 888a3da0f6..c5562cc6f0 100644 --- a/tests/providers/aws/services/kms/kms_cmk_rotation_enabled/kms_cmk_rotation_enabled_test.py +++ b/tests/providers/aws/services/kms/kms_cmk_rotation_enabled/kms_cmk_rotation_enabled_test.py @@ -1,7 +1,9 @@ +from typing import Any, List from unittest import mock from boto3 import client from moto import mock_aws +import pytest from tests.providers.aws.utils import AWS_REGION_US_EAST_1, set_mocked_aws_provider @@ -66,6 +68,65 @@ class Test_kms_cmk_rotation_enabled: assert result[0].resource_id == key["KeyId"] assert result[0].resource_arn == key["Arn"] + @pytest.mark.parametrize( + "no_of_keys_created,expected_no_of_passes", + [ + (5, 3), + (7, 5), + (10, 8), + ], + ) + @mock_aws + def test_kms_cmk_rotation_enabled_when_get_key_rotation_status_fails_on_2_keys_out_of_x_keys( + self, no_of_keys_created: int, expected_no_of_passes: int + ) -> None: + # Generate KMS Client + kms_client = client("kms", region_name=AWS_REGION_US_EAST_1) + kms_client.__dict__["region"] = AWS_REGION_US_EAST_1 + # Creaty KMS key with rotation + for i in range(no_of_keys_created): + key = kms_client.create_key()["KeyMetadata"] + if i not in [2, 4]: + kms_client.enable_key_rotation(KeyId=key["KeyId"]) + + orig_get_key_rotation_status = kms_client.get_key_rotation_status + + def mock_get_key_rotation_status(KeyId: str, count: List[int] = [0]) -> Any: + if count[0] in [2, 4]: + count[0] += 1 + raise Exception("FakeClientError") + else: + count[0] += 1 + return orig_get_key_rotation_status(KeyId=KeyId) + + kms_client.get_key_rotation_status = mock_get_key_rotation_status + + from prowler.providers.aws.services.kms.kms_service import KMS + + aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) + + with mock.patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), mock.patch( + "prowler.providers.aws.aws_provider.AwsProvider.generate_regional_clients", + return_value={AWS_REGION_US_EAST_1: kms_client}, + ), mock.patch( + "prowler.providers.aws.services.kms.kms_cmk_rotation_enabled.kms_cmk_rotation_enabled.kms_client", + new=KMS(aws_provider), + ): + # Test Check + from prowler.providers.aws.services.kms.kms_cmk_rotation_enabled.kms_cmk_rotation_enabled import ( + kms_cmk_rotation_enabled, + ) + + check = kms_cmk_rotation_enabled() + result = check.execute() + + assert len(result) == no_of_keys_created + statuses = [r.status for r in result] + assert statuses.count("PASS") == expected_no_of_passes + @mock_aws def test_kms_cmk_rotation_disabled(self): # Generate KMS Client diff --git a/tests/providers/aws/services/kms/kms_key_not_publicly_accessible/kms_key_not_publicly_accessible_test.py b/tests/providers/aws/services/kms/kms_key_not_publicly_accessible/kms_key_not_publicly_accessible_test.py index 6e0caa22d5..c8abdaaf5d 100644 --- a/tests/providers/aws/services/kms/kms_key_not_publicly_accessible/kms_key_not_publicly_accessible_test.py +++ b/tests/providers/aws/services/kms/kms_key_not_publicly_accessible/kms_key_not_publicly_accessible_test.py @@ -1,6 +1,8 @@ import json +from typing import Any, List from unittest import mock +import pytest from boto3 import client from moto import mock_aws @@ -14,12 +16,15 @@ class Test_kms_key_not_publicly_accessible: aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) - with mock.patch( - "prowler.providers.common.provider.Provider.get_global_provider", - return_value=aws_provider, - ), mock.patch( - "prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible.kms_client", - new=KMS(aws_provider), + with ( + mock.patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), + mock.patch( + "prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible.kms_client", + new=KMS(aws_provider), + ), ): # Test Check from prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible import ( @@ -36,18 +41,21 @@ class Test_kms_key_not_publicly_accessible: # Generate KMS Client kms_client = client("kms", region_name=AWS_REGION_US_EAST_1) # Creaty KMS key without policy - key = kms_client.create_key()["KeyMetadata"] + key = kms_client.create_key(MultiRegion=False)["KeyMetadata"] from prowler.providers.aws.services.kms.kms_service import KMS aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) - with mock.patch( - "prowler.providers.common.provider.Provider.get_global_provider", - return_value=aws_provider, - ), mock.patch( - "prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible.kms_client", - new=KMS(aws_provider), + with ( + mock.patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), + mock.patch( + "prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible.kms_client", + new=KMS(aws_provider), + ), ): # Test Check from prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible import ( @@ -72,6 +80,7 @@ class Test_kms_key_not_publicly_accessible: kms_client = client("kms", region_name=AWS_REGION_US_EAST_1) # Creaty KMS key with public policy key = kms_client.create_key( + MultiRegion=False, Policy=json.dumps( { "Version": "2012-10-17", @@ -86,19 +95,22 @@ class Test_kms_key_not_publicly_accessible: } ], } - ) + ), )["KeyMetadata"] from prowler.providers.aws.services.kms.kms_service import KMS aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) - with mock.patch( - "prowler.providers.common.provider.Provider.get_global_provider", - return_value=aws_provider, - ), mock.patch( - "prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible.kms_client", - new=KMS(aws_provider), + with ( + mock.patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), + mock.patch( + "prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible.kms_client", + new=KMS(aws_provider), + ), ): # Test Check from prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible import ( @@ -123,6 +135,7 @@ class Test_kms_key_not_publicly_accessible: kms_client = client("kms", region_name=AWS_REGION_US_EAST_1) # Creaty KMS key with public policy key = kms_client.create_key( + MultiRegion=False, Policy=json.dumps( { "Version": "2012-10-17", @@ -136,19 +149,22 @@ class Test_kms_key_not_publicly_accessible: } ], } - ) + ), )["KeyMetadata"] from prowler.providers.aws.services.kms.kms_service import KMS aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) - with mock.patch( - "prowler.providers.common.provider.Provider.get_global_provider", - return_value=aws_provider, - ), mock.patch( - "prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible.kms_client", - new=KMS(aws_provider), + with ( + mock.patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), + mock.patch( + "prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible.kms_client", + new=KMS(aws_provider), + ), ): # Test Check from prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible import ( @@ -166,3 +182,66 @@ class Test_kms_key_not_publicly_accessible: ) assert result[0].resource_id == key["KeyId"] assert result[0].resource_arn == key["Arn"] + + @pytest.mark.parametrize( + "no_of_keys_created,expected_no_of_passes", + [ + (5, 3), + (7, 5), + (10, 8), + ], + ) + @mock_aws + def test_kms_key_not_publicly_accessible_when_get_key_policy_fails_on_2_keys_out_of_x_keys( + self, no_of_keys_created: int, expected_no_of_passes: int + ) -> None: + # Generate KMS Client + kms_client = client("kms", region_name=AWS_REGION_US_EAST_1) + kms_client.__dict__["region"] = AWS_REGION_US_EAST_1 + # Creaty KMS key with public policy + for i in range(no_of_keys_created): + kms_client.create_key(MultiRegion=False) + + orig_get_key_policy = kms_client.get_key_policy + + def mock_get_key_policy( + KeyId: str, PolicyName: str, count: List[int] = [0] + ) -> Any: + if count[0] in [2, 4]: + count[0] += 1 + raise Exception("FakeClientError") + else: + count[0] += 1 + return orig_get_key_policy(KeyId=KeyId, PolicyName=PolicyName) + + kms_client.get_key_policy = mock_get_key_policy + + from prowler.providers.aws.services.kms.kms_service import KMS + + aws_provider = set_mocked_aws_provider([AWS_REGION_US_EAST_1]) + + with ( + mock.patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=aws_provider, + ), + mock.patch( + "prowler.providers.aws.aws_provider.AwsProvider.generate_regional_clients", + return_value={AWS_REGION_US_EAST_1: kms_client}, + ), + mock.patch( + "prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible.kms_client", + new=KMS(aws_provider), + ), + ): + # Test Check + from prowler.providers.aws.services.kms.kms_key_not_publicly_accessible.kms_key_not_publicly_accessible import ( + kms_key_not_publicly_accessible, + ) + + check = kms_key_not_publicly_accessible() + result = check.execute() + + assert len(result) == expected_no_of_passes + statuses = [r.status for r in result] + assert statuses.count("PASS") == expected_no_of_passes