feat(compliance): add method list_compliance_requirements (#4890)

Co-authored-by: Pepe Fagoaga <pepe@prowler.com>
This commit is contained in:
Pedro Martín
2024-09-04 20:35:26 +02:00
committed by GitHub
parent 3933440a08
commit a975e96a45
2 changed files with 186 additions and 41 deletions
+72 -6
View File
@@ -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(
+114 -35
View File
@@ -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