feat(check): add check methods (#5462)

This commit is contained in:
Pedro Martín
2024-10-28 12:22:34 +01:00
committed by GitHub
parent 6502330512
commit 0114d0462f
8 changed files with 637 additions and 222 deletions
+9 -11
View File
@@ -64,13 +64,8 @@ class TestCheckLoader:
categories = None
with patch(
"prowler.lib.check.checks_loader.recover_checks_from_provider",
return_value=[
(
f"{S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME}",
"path/to/{S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME}",
)
],
"prowler.lib.check.checks_loader.CheckMetadata.list",
return_value={S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME},
):
assert {S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME} == load_checks_to_execute(
bulk_checks_metatada,
@@ -145,7 +140,7 @@ class TestCheckLoader:
categories = None
with patch(
"prowler.lib.check.checks_loader.recover_checks_from_service",
"prowler.lib.check.checks_loader.CheckMetadata.list_by_service",
return_value={S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME},
):
assert {S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME} == load_checks_to_execute(
@@ -175,7 +170,10 @@ class TestCheckLoader:
categories = None
with patch(
"prowler.lib.check.checks_loader.recover_checks_from_service",
"prowler.lib.check.checks_loader.CheckMetadata.list_by_severity",
return_value={S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME},
), patch(
"prowler.lib.check.checks_loader.CheckMetadata.list_by_service",
return_value={"ec2_ami_public"},
):
assert set() == load_checks_to_execute(
@@ -235,7 +233,7 @@ class TestCheckLoader:
categories = None
with patch(
"prowler.lib.check.checks_loader.recover_checks_from_service",
"prowler.lib.check.checks_loader.CheckMetadata.list_by_service",
return_value={S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME},
):
assert {S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME} == load_checks_to_execute(
@@ -265,7 +263,7 @@ class TestCheckLoader:
categories = None
with patch(
"prowler.lib.check.checks_loader.parse_checks_from_compliance_framework",
"prowler.lib.check.checks_loader.CheckMetadata.list",
return_value={S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME},
):
assert {S3_BUCKET_LEVEL_PUBLIC_ACCESS_BLOCK_NAME} == load_checks_to_execute(
+132 -122
View File
@@ -10,98 +10,100 @@ from prowler.lib.check.compliance_models import (
)
from prowler.lib.check.models import CheckMetadata
custom_compliance_metadata = {
"framework1_aws": Compliance(
Framework="Framework1",
Provider="aws",
Version="1.0",
Description="Framework 1 Description",
Requirements=[
Compliance_Requirement(
Id="1.1.1",
Description="description",
Attributes=[
CIS_Requirement_Attribute(
Section="1. Identity",
Profile=CIS_Requirement_Attribute_Profile("Level 1"),
AssessmentStatus=CIS_Requirement_Attribute_AssessmentStatus(
"Manual"
),
Description="Description",
RationaleStatement="Rationale",
ImpactStatement="Impact",
RemediationProcedure="Remediation",
AuditProcedure="Audit",
AdditionalInformation="Additional",
References="References",
)
],
Checks=[
"accessanalyzer_enabled",
"iam_user_mfa_enabled_console_access",
],
),
# Manual requirement
Compliance_Requirement(
Id="1.1.2",
Description="description",
Attributes=[
CIS_Requirement_Attribute(
Section="1. Identity",
Profile=CIS_Requirement_Attribute_Profile("Level 1"),
AssessmentStatus=CIS_Requirement_Attribute_AssessmentStatus(
"Manual"
),
Description="Description",
RationaleStatement="Rationale",
ImpactStatement="Impact",
RemediationProcedure="Remediation",
AuditProcedure="Audit",
AdditionalInformation="Additional",
References="References",
)
],
Checks=[],
),
],
),
"framework1_azure": Compliance(
Framework="Framework1",
Provider="azure",
Version="1.0",
Description="Framework 2 Description",
Requirements=[
Compliance_Requirement(
Id="1.1.1",
Description="description",
Attributes=[
CIS_Requirement_Attribute(
Section="1. Identity",
Profile=CIS_Requirement_Attribute_Profile("Level 1"),
AssessmentStatus=CIS_Requirement_Attribute_AssessmentStatus(
"Manual"
),
Description="Description",
RationaleStatement="Rationale",
ImpactStatement="Impact",
RemediationProcedure="Remediation",
AuditProcedure="Audit",
AdditionalInformation="Additional",
References="References",
)
],
Checks=[],
)
],
),
}
class TestCompliance:
def get_custom_framework(self):
return {
"framework1_aws": Compliance(
Framework="Framework1",
Provider="aws",
Version="1.0",
Description="Framework 1 Description",
Requirements=[
Compliance_Requirement(
Id="1.1.1",
Description="description",
Attributes=[
CIS_Requirement_Attribute(
Section="1. Identity",
Profile=CIS_Requirement_Attribute_Profile("Level 1"),
AssessmentStatus=CIS_Requirement_Attribute_AssessmentStatus(
"Manual"
),
Description="Description",
RationaleStatement="Rationale",
ImpactStatement="Impact",
RemediationProcedure="Remediation",
AuditProcedure="Audit",
AdditionalInformation="Additional",
References="References",
)
],
Checks=["check1", "check2"],
),
# Manual requirement
Compliance_Requirement(
Id="1.1.2",
Description="description",
Attributes=[
CIS_Requirement_Attribute(
Section="1. Identity",
Profile=CIS_Requirement_Attribute_Profile("Level 1"),
AssessmentStatus=CIS_Requirement_Attribute_AssessmentStatus(
"Manual"
),
Description="Description",
RationaleStatement="Rationale",
ImpactStatement="Impact",
RemediationProcedure="Remediation",
AuditProcedure="Audit",
AdditionalInformation="Additional",
References="References",
)
],
Checks=[],
),
],
),
"framework1_azure": Compliance(
Framework="Framework1",
Provider="azure",
Version="1.0",
Description="Framework 2 Description",
Requirements=[
Compliance_Requirement(
Id="1.1.1",
Description="description",
Attributes=[
CIS_Requirement_Attribute(
Section="1. Identity",
Profile=CIS_Requirement_Attribute_Profile("Level 1"),
AssessmentStatus=CIS_Requirement_Attribute_AssessmentStatus(
"Manual"
),
Description="Description",
RationaleStatement="Rationale",
ImpactStatement="Impact",
RemediationProcedure="Remediation",
AuditProcedure="Audit",
AdditionalInformation="Additional",
References="References",
)
],
Checks=[],
)
],
),
}
def get_custom_check_metadata(self):
return {
"check1": CheckMetadata(
"accessanalyzer_enabled": CheckMetadata(
Provider="aws",
CheckID="check1",
CheckID="accessanalyzer_enabled",
CheckTitle="Check 1",
CheckType=["type1"],
ServiceName="service1",
@@ -127,9 +129,9 @@ class TestCompliance:
Notes="notes1",
Compliance=[],
),
"check2": CheckMetadata(
"iam_user_mfa_enabled_console_access": CheckMetadata(
Provider="aws",
CheckID="check2",
CheckID="iam_user_mfa_enabled_console_access",
CheckTitle="Check 2",
CheckType=["type2"],
ServiceName="service2",
@@ -158,44 +160,52 @@ class TestCompliance:
}
def test_update_checks_metadata(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
bulk_checks_metadata = self.get_custom_check_metadata()
updated_metadata = update_checks_metadata_with_compliance(
bulk_compliance_frameworks, bulk_checks_metadata
)
assert "check1" in updated_metadata
assert "check2" in updated_metadata
assert "accessanalyzer_enabled" in updated_metadata
assert "iam_user_mfa_enabled_console_access" in updated_metadata
check1_compliance = updated_metadata["check1"].Compliance[0]
accessanalyzer_enabled_compliance = updated_metadata[
"accessanalyzer_enabled"
].Compliance[0]
assert len(updated_metadata["check1"].Compliance) == 1
assert check1_compliance.Framework == "Framework1"
assert check1_compliance.Provider == "aws"
assert check1_compliance.Version == "1.0"
assert check1_compliance.Description == "Framework 1 Description"
assert len(check1_compliance.Requirements) == 1
assert len(updated_metadata["accessanalyzer_enabled"].Compliance) == 1
assert accessanalyzer_enabled_compliance.Framework == "Framework1"
assert accessanalyzer_enabled_compliance.Provider == "aws"
assert accessanalyzer_enabled_compliance.Version == "1.0"
assert (
accessanalyzer_enabled_compliance.Description == "Framework 1 Description"
)
assert len(accessanalyzer_enabled_compliance.Requirements) == 1
check1_requirement = check1_compliance.Requirements[0]
assert check1_requirement.Id == "1.1.1"
assert check1_requirement.Description == "description"
assert len(check1_requirement.Attributes) == 1
accessanalyzer_enabled_requirement = (
accessanalyzer_enabled_compliance.Requirements[0]
)
assert accessanalyzer_enabled_requirement.Id == "1.1.1"
assert accessanalyzer_enabled_requirement.Description == "description"
assert len(accessanalyzer_enabled_requirement.Attributes) == 1
check1_attribute = check1_requirement.Attributes[0]
assert check1_attribute.Section == "1. Identity"
assert check1_attribute.Profile == "Level 1"
assert check1_attribute.AssessmentStatus == "Manual"
assert check1_attribute.Description == "Description"
assert check1_attribute.RationaleStatement == "Rationale"
assert check1_attribute.ImpactStatement == "Impact"
assert check1_attribute.RemediationProcedure == "Remediation"
assert check1_attribute.AuditProcedure == "Audit"
assert check1_attribute.AdditionalInformation == "Additional"
assert check1_attribute.References == "References"
accessanalyzer_enabled_attribute = (
accessanalyzer_enabled_requirement.Attributes[0]
)
assert accessanalyzer_enabled_attribute.Section == "1. Identity"
assert accessanalyzer_enabled_attribute.Profile == "Level 1"
assert accessanalyzer_enabled_attribute.AssessmentStatus == "Manual"
assert accessanalyzer_enabled_attribute.Description == "Description"
assert accessanalyzer_enabled_attribute.RationaleStatement == "Rationale"
assert accessanalyzer_enabled_attribute.ImpactStatement == "Impact"
assert accessanalyzer_enabled_attribute.RemediationProcedure == "Remediation"
assert accessanalyzer_enabled_attribute.AuditProcedure == "Audit"
assert accessanalyzer_enabled_attribute.AdditionalInformation == "Additional"
assert accessanalyzer_enabled_attribute.References == "References"
def test_list_no_provider(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
list_compliance = Compliance.list(bulk_compliance_frameworks)
@@ -204,7 +214,7 @@ class TestCompliance:
assert list_compliance[1] == "framework1_azure"
def test_list_with_provider_aws(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
list_compliance = Compliance.list(bulk_compliance_frameworks, provider="aws")
@@ -212,7 +222,7 @@ class TestCompliance:
assert list_compliance[0] == "framework1_aws"
def test_list_with_provider_azure(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
list_compliance = Compliance.list(bulk_compliance_frameworks, provider="azure")
@@ -220,7 +230,7 @@ class TestCompliance:
assert list_compliance[0] == "framework1_azure"
def test_get_compliance_frameworks(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
compliance_framework = Compliance.get(
bulk_compliance_frameworks, compliance_framework_name="framework1_aws"
@@ -243,7 +253,7 @@ class TestCompliance:
assert len(compliance_framework.Requirements) == 1
def test_get_non_existent_framework(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
compliance_framework = Compliance.get(
bulk_compliance_frameworks, compliance_framework_name="non_existent"
@@ -252,14 +262,14 @@ class TestCompliance:
assert compliance_framework is None
def test_list_compliance_requirements_no_compliance(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
list_requirements = Compliance.list_requirements(bulk_compliance_frameworks)
assert len(list_requirements) == 0
def test_list_compliance_requirements_with_compliance(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
list_requirements = Compliance.list_requirements(
bulk_compliance_frameworks, compliance_framework="framework1_aws"
@@ -277,7 +287,7 @@ class TestCompliance:
assert list_requirements[0] == "1.1.1"
def test_get_compliance_requirement(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
compliance_requirement = Compliance.get_requirement(
bulk_compliance_frameworks,
@@ -310,7 +320,7 @@ class TestCompliance:
assert len(compliance_requirement.Attributes) == 1
def test_get_compliance_requirement_not_found(self):
bulk_compliance_frameworks = self.get_custom_framework()
bulk_compliance_frameworks = custom_compliance_metadata
compliance_requirement = Compliance.get_requirement(
bulk_compliance_frameworks,
+232 -31
View File
@@ -1,6 +1,36 @@
from unittest import mock
from prowler.lib.check.models import CheckMetadata
from tests.lib.check.compliance_check_test import custom_compliance_metadata
mock_metadata = CheckMetadata(
Provider="aws",
CheckID="accessanalyzer_enabled",
CheckTitle="Check 1",
CheckType=["type1"],
ServiceName="service1",
SubServiceName="subservice1",
ResourceIdTemplate="template1",
Severity="high",
ResourceType="resource1",
Description="Description 1",
Risk="risk1",
RelatedUrl="url1",
Remediation={
"Code": {
"CLI": "cli1",
"NativeIaC": "native1",
"Other": "other1",
"Terraform": "terraform1",
},
"Recommendation": {"Text": "text1", "Url": "url1"},
},
Categories=["categoryone"],
DependsOn=["dependency1"],
RelatedTo=["related1"],
Notes="notes1",
Compliance=[],
)
class TestCheckMetada:
@@ -13,44 +43,215 @@ class TestCheckMetada:
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
check_metadata = CheckMetadata(
Provider="aws",
CheckID="accessanalyzer_enabled",
CheckTitle="Check 1",
CheckType=["type1"],
ServiceName="service1",
SubServiceName="subservice1",
ResourceIdTemplate="template1",
Severity="high",
ResourceType="resource1",
Description="Description 1",
Risk="risk1",
RelatedUrl="url1",
Remediation={
"Code": {
"CLI": "cli1",
"NativeIaC": "native1",
"Other": "other1",
"Terraform": "terraform1",
},
"Recommendation": {"Text": "text1", "Url": "url1"},
},
Categories=["categoryone"],
DependsOn=["dependency1"],
RelatedTo=["related1"],
Notes="notes1",
Compliance=[],
)
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = check_metadata
mock_load_metadata.return_value = mock_metadata
result = CheckMetadata.get_bulk(provider="aws")
# Assertions
assert "accessanalyzer_enabled" in result.keys()
assert result["accessanalyzer_enabled"] == check_metadata
assert result["accessanalyzer_enabled"] == mock_metadata
mock_recover_checks.assert_called_once_with("aws")
mock_load_metadata.assert_called_once_with(
"/path/to/accessanalyzer_enabled/accessanalyzer_enabled.metadata.json"
)
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_list(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of recover_checks_from_provider
mock_recover_checks.return_value = [
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(bulk_checks_metadata=bulk_metadata)
# Assertions
assert result == {"accessanalyzer_enabled"}
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_get(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of recover_checks_from_provider
mock_recover_checks.return_value = [
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(bulk_checks_metadata=bulk_metadata)
# Assertions
assert result == {"accessanalyzer_enabled"}
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_list_by_severity(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of recover_checks_from_provider
mock_recover_checks.return_value = [
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(bulk_checks_metadata=bulk_metadata, severity="high")
# Assertions
assert result == {"accessanalyzer_enabled"}
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_list_by_severity_not_values(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of recover_checks_from_provider
mock_recover_checks.return_value = [
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(bulk_checks_metadata=bulk_metadata, severity="low")
# Assertions
assert result == set()
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_list_by_category(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of recover_checks_from_provider
mock_recover_checks.return_value = [
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(
bulk_checks_metadata=bulk_metadata, category="categoryone"
)
# Assertions
assert result == {"accessanalyzer_enabled"}
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_list_by_category_not_valid(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of recover_checks_from_provider
mock_recover_checks.return_value = [
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(
bulk_checks_metadata=bulk_metadata, category="categorytwo"
)
# Assertions
assert result == set()
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_list_by_service(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of recover_checks_from_provider
mock_recover_checks.return_value = [
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(
bulk_checks_metadata=bulk_metadata, service="service1"
)
# Assertions
assert result == {"accessanalyzer_enabled"}
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_list_by_service_invalid(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of recover_checks_from_provider
mock_recover_checks.return_value = [
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(
bulk_checks_metadata=bulk_metadata, service="service2"
)
# Assertions
assert result == set()
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_list_by_compliance(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of recover_checks_from_provider
mock_recover_checks.return_value = [
("accessanalyzer_enabled", "/path/to/accessanalyzer_enabled")
]
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
bulk_compliance_frameworks = custom_compliance_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(
bulk_checks_metadata=bulk_metadata,
bulk_compliance_frameworks=bulk_compliance_frameworks,
compliance_framework="framework1_aws",
)
# Assertions
assert result == {"accessanalyzer_enabled"}
def test_list_by_compliance_empty(self):
bulk_compliance_frameworks = custom_compliance_metadata
result = CheckMetadata.list(
bulk_compliance_frameworks=bulk_compliance_frameworks,
compliance_framework="framework1_azure",
)
# Assertions
assert result == set()
@mock.patch("prowler.lib.check.models.load_check_metadata")
@mock.patch("prowler.lib.check.models.recover_checks_from_provider")
def test_list_only_check_metadata(self, mock_recover_checks, mock_load_metadata):
# Mock the return value of load_check_metadata
mock_load_metadata.return_value = mock_metadata
bulk_metadata = CheckMetadata.get_bulk(provider="aws")
result = CheckMetadata.list(bulk_checks_metadata=bulk_metadata)
assert result == set()
+11 -2
View File
@@ -121,6 +121,15 @@ def mock_load_check_metadata():
yield mock_load
@pytest.fixture
def mock_load_checks_to_execute():
with mock.patch(
"prowler.lib.check.models.CheckMetadata.list", autospec=True
) as mock_load:
mock_load.return_value = {"accessanalyzer_enabled"}
yield mock_load
class TestScan:
def test_init(mock_provider):
checks_to_execute = {
@@ -260,16 +269,16 @@ class TestScan:
def test_init_with_no_checks(
mock_provider,
mock_list_modules,
mock_recover_checks_from_provider,
mock_load_check_metadata,
mock_load_checks_to_execute,
):
checks_to_execute = set()
mock_provider.type = "aws"
scan = Scan(mock_provider, checks=checks_to_execute)
mock_list_modules.assert_called_once_with("aws", None)
mock_load_check_metadata.assert_called_once()
mock_load_checks_to_execute.assert_called_once()
mock_recover_checks_from_provider.assert_called_once_with("aws")
assert scan.provider == mock_provider