From a975e96a45a36ce573f1044daef183c62dbeb052 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Pedro=20Mart=C3=ADn?= Date: Wed, 4 Sep 2024 20:35:26 +0200 Subject: [PATCH] feat(compliance): add method list_compliance_requirements (#4890) Co-authored-by: Pepe Fagoaga --- prowler/lib/check/compliance_models.py | 78 +++++++++++- tests/lib/check/compliance_check_test.py | 149 +++++++++++++++++------ 2 files changed, 186 insertions(+), 41 deletions(-) diff --git a/prowler/lib/check/compliance_models.py b/prowler/lib/check/compliance_models.py index b415a0843c..ed7ba1341e 100644 --- a/prowler/lib/check/compliance_models.py +++ b/prowler/lib/check/compliance_models.py @@ -213,9 +213,8 @@ class Compliance(BaseModel): raise ValueError("Framework or Provider must not be empty") return values - def list_compliance_frameworks( - bulk_compliance_frameworks: dict, provider: str = None - ): + @staticmethod + def list(bulk_compliance_frameworks: dict, provider: str = None) -> list[str]: """ Returns a list of compliance frameworks from bulk compliance frameworks @@ -229,17 +228,84 @@ class Compliance(BaseModel): if provider: compliance_frameworks = [ compliance_framework - for compliance_framework in bulk_compliance_frameworks.values() - if compliance_framework.Provider == provider + for compliance_framework in bulk_compliance_frameworks.keys() + if provider in compliance_framework ] else: compliance_frameworks = [ compliance_framework - for compliance_framework in bulk_compliance_frameworks.values() + for compliance_framework in bulk_compliance_frameworks.keys() ] return compliance_frameworks + @staticmethod + def get( + bulk_compliance_frameworks: dict, compliance_framework_name: str + ) -> "Compliance": + """ + Returns a compliance framework from bulk compliance frameworks + + Args: + bulk_compliance_frameworks (dict): The bulk compliance frameworks + compliance_framework_name (str): The compliance framework name + + Returns: + Compliance: The compliance framework + """ + return bulk_compliance_frameworks.get(compliance_framework_name, None) + + @staticmethod + def list_requirements( + bulk_compliance_frameworks: dict, compliance_framework: str = None + ) -> list: + """ + Returns a list of compliance requirements from a compliance framework + + Args: + bulk_compliance_frameworks (dict): The bulk compliance frameworks + compliance_framework (str): The compliance framework name + + Returns: + list: The list of compliance requirements for the provided compliance framework + """ + compliance_requirements = [] + + if bulk_compliance_frameworks and compliance_framework: + compliance_requirements = [ + compliance_requirement.Id + for compliance_requirement in bulk_compliance_frameworks.get( + compliance_framework + ).Requirements + ] + + return compliance_requirements + + @staticmethod + def get_requirement( + bulk_compliance_frameworks: dict, compliance_framework: str, requirement_id: str + ) -> Union[Mitre_Requirement, Compliance_Requirement]: + """ + Returns a compliance requirement from a compliance framework + + Args: + bulk_compliance_frameworks (dict): The bulk compliance frameworks + compliance_framework (str): The compliance framework name + requirement_id (str): The compliance requirement ID + + Returns: + Mitre_Requirement | Compliance_Requirement: The compliance requirement + """ + requirement = None + for compliance_requirement in bulk_compliance_frameworks.get( + compliance_framework + ).Requirements: + if compliance_requirement.Id == requirement_id: + requirement = compliance_requirement + break + + return requirement + # Testing Pending def load_compliance_framework( diff --git a/tests/lib/check/compliance_check_test.py b/tests/lib/check/compliance_check_test.py index ae461b31fb..69a2ea76d7 100644 --- a/tests/lib/check/compliance_check_test.py +++ b/tests/lib/check/compliance_check_test.py @@ -13,7 +13,7 @@ class TestCompliance: def get_custom_framework(self): return { - "framework1": Compliance( + "framework1_aws": Compliance( Framework="Framework1", Provider="aws", Version="1.0", @@ -64,8 +64,8 @@ class TestCompliance: ), ], ), - "framework2": Compliance( - Framework="Framework2", + "framework1_azure": Compliance( + Framework="Framework1", Provider="azure", Version="1.0", Description="Framework 2 Description", @@ -192,49 +192,128 @@ class TestCompliance: assert check1_attribute.AdditionalInformation == "Additional" assert check1_attribute.References == "References" - def test_list_compliance_frameworks_no_provider(self): + def test_list_no_provider(self): bulk_compliance_frameworks = self.get_custom_framework() - list_compliance = Compliance.list_compliance_frameworks( - bulk_compliance_frameworks - ) + list_compliance = Compliance.list(bulk_compliance_frameworks) assert len(list_compliance) == 2 - assert list_compliance[0].Framework == "Framework1" - assert list_compliance[0].Provider == "aws" - assert list_compliance[0].Version == "1.0" - assert list_compliance[0].Description == "Framework 1 Description" - assert len(list_compliance[0].Requirements) == 2 - assert list_compliance[1].Framework == "Framework2" - assert list_compliance[1].Provider == "azure" - assert list_compliance[1].Version == "1.0" - assert list_compliance[1].Description == "Framework 2 Description" - assert len(list_compliance[1].Requirements) == 1 + assert list_compliance[0] == "framework1_aws" + assert list_compliance[1] == "framework1_azure" - def test_list_compliance_frameworks_with_provider_aws(self): + def test_list_with_provider_aws(self): bulk_compliance_frameworks = self.get_custom_framework() - list_compliance = Compliance.list_compliance_frameworks( - bulk_compliance_frameworks, provider="aws" - ) + list_compliance = Compliance.list(bulk_compliance_frameworks, provider="aws") assert len(list_compliance) == 1 - assert list_compliance[0].Framework == "Framework1" - assert list_compliance[0].Provider == "aws" - assert list_compliance[0].Version == "1.0" - assert list_compliance[0].Description == "Framework 1 Description" - assert len(list_compliance[0].Requirements) == 2 + assert list_compliance[0] == "framework1_aws" - def test_list_compliance_frameworks_with_provider_azure(self): + def test_list_with_provider_azure(self): bulk_compliance_frameworks = self.get_custom_framework() - list_compliance = Compliance.list_compliance_frameworks( - bulk_compliance_frameworks, provider="azure" - ) + list_compliance = Compliance.list(bulk_compliance_frameworks, provider="azure") assert len(list_compliance) == 1 - assert list_compliance[0].Framework == "Framework2" - assert list_compliance[0].Provider == "azure" - assert list_compliance[0].Version == "1.0" - assert list_compliance[0].Description == "Framework 2 Description" - assert len(list_compliance[0].Requirements) == 1 + assert list_compliance[0] == "framework1_azure" + + def test_get_compliance_frameworks(self): + bulk_compliance_frameworks = self.get_custom_framework() + + compliance_framework = Compliance.get( + bulk_compliance_frameworks, compliance_framework_name="framework1_aws" + ) + + assert compliance_framework.Framework == "Framework1" + assert compliance_framework.Provider == "aws" + assert compliance_framework.Version == "1.0" + assert compliance_framework.Description == "Framework 1 Description" + assert len(compliance_framework.Requirements) == 2 + + compliance_framework = Compliance.get( + bulk_compliance_frameworks, compliance_framework_name="framework1_azure" + ) + + assert compliance_framework.Framework == "Framework1" + assert compliance_framework.Provider == "azure" + assert compliance_framework.Version == "1.0" + assert compliance_framework.Description == "Framework 2 Description" + assert len(compliance_framework.Requirements) == 1 + + def test_get_non_existent_framework(self): + bulk_compliance_frameworks = self.get_custom_framework() + + compliance_framework = Compliance.get( + bulk_compliance_frameworks, compliance_framework_name="non_existent" + ) + + assert compliance_framework is None + + def test_list_compliance_requirements_no_compliance(self): + bulk_compliance_frameworks = self.get_custom_framework() + + 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() + + list_requirements = Compliance.list_requirements( + bulk_compliance_frameworks, compliance_framework="framework1_aws" + ) + + assert len(list_requirements) == 2 + assert list_requirements[0] == "1.1.1" + assert list_requirements[1] == "1.1.2" + + list_requirements = Compliance.list_requirements( + bulk_compliance_frameworks, compliance_framework="framework1_azure" + ) + + assert len(list_requirements) == 1 + assert list_requirements[0] == "1.1.1" + + def test_get_compliance_requirement(self): + bulk_compliance_frameworks = self.get_custom_framework() + + compliance_requirement = Compliance.get_requirement( + bulk_compliance_frameworks, + compliance_framework="framework1_aws", + requirement_id="1.1.1", + ) + + assert compliance_requirement.Id == "1.1.1" + assert compliance_requirement.Description == "description" + assert len(compliance_requirement.Attributes) == 1 + + compliance_requirement = Compliance.get_requirement( + bulk_compliance_frameworks, + compliance_framework="framework1_aws", + requirement_id="1.1.2", + ) + + assert compliance_requirement.Id == "1.1.2" + assert compliance_requirement.Description == "description" + assert len(compliance_requirement.Attributes) == 1 + + compliance_requirement = Compliance.get_requirement( + bulk_compliance_frameworks, + compliance_framework="framework1_azure", + requirement_id="1.1.1", + ) + + assert compliance_requirement.Id == "1.1.1" + assert compliance_requirement.Description == "description" + assert len(compliance_requirement.Attributes) == 1 + + def test_get_compliance_requirement_not_found(self): + bulk_compliance_frameworks = self.get_custom_framework() + + compliance_requirement = Compliance.get_requirement( + bulk_compliance_frameworks, + compliance_framework="framework1_aws", + requirement_id="1.1.3", + ) + + assert compliance_requirement is None