mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-21 03:21:56 +00:00
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:
+3
-1
@@ -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
|
||||
|
||||
+61
@@ -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
|
||||
|
||||
+106
-27
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user