fix(kms): Amazon KMS API call error handling (#6904)

Co-authored-by: Ogonna Iwunze <1915636+wunzeco@users.noreply.github.com>
This commit is contained in:
Prowler Bot
2025-02-12 17:08:47 +01:00
committed by GitHub
parent 287eef5085
commit 91b74822e9
5 changed files with 282 additions and 59 deletions
@@ -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"
@@ -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):
@@ -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
@@ -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
@@ -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