From 7360395263458aedd3c5015fd43a02e35aa5a3dd Mon Sep 17 00:00:00 2001 From: Daniel Barranquero Date: Tue, 10 Jun 2025 19:15:16 +0200 Subject: [PATCH] feat(tests): add tests for new fixers --- tests/lib/fix/fixer_test.py | 207 ++++++++++++++++++ tests/providers/aws/lib/fix/awsfixer_test.py | 115 ++++++++++ .../azure/lib/fix/azurefixer_test.py | 126 +++++++++++ .../providers/m365/lib/fix/m365fixer_test.py | 103 +++++++++ 4 files changed, 551 insertions(+) create mode 100644 tests/lib/fix/fixer_test.py create mode 100644 tests/providers/aws/lib/fix/awsfixer_test.py create mode 100644 tests/providers/azure/lib/fix/azurefixer_test.py create mode 100644 tests/providers/m365/lib/fix/m365fixer_test.py diff --git a/tests/lib/fix/fixer_test.py b/tests/lib/fix/fixer_test.py new file mode 100644 index 0000000000..0bb5e62021 --- /dev/null +++ b/tests/lib/fix/fixer_test.py @@ -0,0 +1,207 @@ +import json +from unittest.mock import MagicMock, patch + +import pytest + +from prowler.lib.check.models import ( + Check_Report, + CheckMetadata, + Code, + Recommendation, + Remediation, +) +from prowler.lib.fix.fixer import Fixer + + +def get_mock_metadata( + provider="aws", check_id="test_check", service_name="testservice" +): + return CheckMetadata( + Provider=provider, + CheckID=check_id, + CheckTitle="Test Check", + CheckType=["type1"], + CheckAliases=[], + ServiceName=service_name, + SubServiceName="", + ResourceIdTemplate="", + Severity="low", + ResourceType="resource", + Description="desc", + Risk="risk", + RelatedUrl="url", + Remediation=Remediation( + Code=Code(NativeIaC="", Terraform="", CLI="", Other=""), + Recommendation=Recommendation(Text="", Url=""), + ), + Categories=["cat1"], + DependsOn=[], + RelatedTo=[], + Notes="", + Compliance=[], + ) + + +def build_metadata(provider="aws", check_id="test_check", service_name="testservice"): + return CheckMetadata( + Provider=provider, + CheckID=check_id, + CheckTitle="Test Check", + CheckType=["type1"], + CheckAliases=[], + ServiceName=service_name, + SubServiceName="", + ResourceIdTemplate="", + Severity="low", + ResourceType="resource", + Description="desc", + Risk="risk", + RelatedUrl="url", + Remediation=Remediation( + Code=Code(NativeIaC="", Terraform="", CLI="", Other=""), + Recommendation=Recommendation(Text="", Url=""), + ), + Categories=["cat1"], + DependsOn=[], + RelatedTo=[], + Notes="", + Compliance=[], + ) + + +def build_finding( + status="FAIL", provider="aws", check_id="test_check", service_name="testservice" +): + metadata = build_metadata(provider, check_id, service_name) + resource = MagicMock() + finding = Check_Report(json.dumps(metadata.dict()), resource) + finding.status = status + return finding + + +class DummyFixer(Fixer): + def fix(self, finding=None, **kwargs): + return True + + +class TestFixer: + def test_get_fixer_info(self): + fixer = DummyFixer( + description="desc", cost_impact=True, cost_description="cost" + ) + info = fixer._get_fixer_info() + assert info == { + "description": "desc", + "cost_impact": True, + "cost_description": "cost", + } + + def test_client_property(self): + fixer = DummyFixer(description="desc") + assert fixer.client is None + + @pytest.mark.parametrize( + "check_id,provider,service_name,expected_class", + [ + (None, "aws", "testservice", None), + ("test_check", None, "testservice", None), + ("nonexistent_check", "aws", "testservice", None), + ], + ) + def test_get_fixer_for_finding_edge( + self, check_id, provider, service_name, expected_class + ): + finding = MagicMock() + finding.check_metadata.CheckID = check_id + finding.check_metadata.Provider = provider + finding.check_metadata.ServiceName = service_name + with patch("prowler.lib.fix.fixer.logger"): + fixer = Fixer.get_fixer_for_finding(finding) + assert fixer is expected_class + + def test_get_fixer_for_finding_importerror_print(self): + finding = MagicMock() + finding.check_metadata.CheckID = "nonexistent_check" + finding.check_metadata.Provider = "aws" + finding.check_metadata.ServiceName = "testservice" + with patch("builtins.print") as mock_print: + fixer = Fixer.get_fixer_for_finding(finding) + assert fixer is None + assert mock_print.called + + def test_run_fixer_single_and_multiple(self): + finding = build_finding(status="FAIL") + with patch.object(Fixer, "run_individual_fixer", return_value=1) as mock_run: + assert Fixer.run_fixer(finding) == 1 + assert mock_run.called + finding.status = "PASS" + assert Fixer.run_fixer(finding) == 0 + finding1 = build_finding(status="FAIL") + finding2 = build_finding(status="FAIL") + with patch.object(Fixer, "run_individual_fixer", return_value=2) as mock_run: + assert Fixer.run_fixer([finding1, finding2]) == 2 + assert mock_run.called + + def test_run_fixer_grouping(self): + finding1 = build_finding(status="FAIL", check_id="check1") + finding2 = build_finding(status="FAIL", check_id="check1") + finding3 = build_finding(status="FAIL", check_id="check2") + calls = {} + + def fake_run_individual_fixer(check_id, findings): + calls[check_id] = len(findings) + return len(findings) + + with patch.object( + Fixer, "run_individual_fixer", side_effect=fake_run_individual_fixer + ): + total = Fixer.run_fixer([finding1, finding2, finding3]) + assert total == 3 + assert calls == {"check1": 2, "check2": 1} + + def test_run_fixer_exception(self): + finding = build_finding(status="FAIL") + with patch.object(Fixer, "run_individual_fixer", side_effect=Exception("fail")): + with patch("prowler.lib.fix.fixer.logger") as mock_logger: + assert Fixer.run_fixer(finding) == 0 + assert mock_logger.error.called + + def test_run_individual_fixer_success(self): + finding = build_finding(status="FAIL") + with ( + patch.object(Fixer, "get_fixer_for_finding") as mock_factory, + patch("builtins.print") as mock_print, + ): + fixer = DummyFixer(description="desc") + mock_factory.return_value = fixer + with patch.object(fixer, "fix", return_value=True): + total = Fixer.run_individual_fixer("test_check", [finding]) + assert total == 1 + assert mock_print.call_count > 0 + + def test_run_individual_fixer_no_fixer(self): + finding = build_finding(status="FAIL") + with patch.object(Fixer, "get_fixer_for_finding", return_value=None): + assert Fixer.run_individual_fixer("test_check", [finding]) == 0 + + def test_run_individual_fixer_fix_error(self): + finding = build_finding(status="FAIL") + with ( + patch.object(Fixer, "get_fixer_for_finding") as mock_factory, + patch("builtins.print") as mock_print, + ): + fixer = DummyFixer(description="desc") + mock_factory.return_value = fixer + with patch.object(fixer, "fix", return_value=False): + total = Fixer.run_individual_fixer("test_check", [finding]) + assert total == 0 + assert mock_print.call_count > 0 + + def test_run_individual_fixer_exception(self): + finding = build_finding(status="FAIL") + with patch.object( + Fixer, "get_fixer_for_finding", side_effect=Exception("fail") + ): + with patch("prowler.lib.fix.fixer.logger") as mock_logger: + assert Fixer.run_individual_fixer("test_check", [finding]) == 0 + assert mock_logger.error.called diff --git a/tests/providers/aws/lib/fix/awsfixer_test.py b/tests/providers/aws/lib/fix/awsfixer_test.py new file mode 100644 index 0000000000..ef6353fc35 --- /dev/null +++ b/tests/providers/aws/lib/fix/awsfixer_test.py @@ -0,0 +1,115 @@ +import json +from unittest.mock import MagicMock, patch + +import pytest + +from prowler.lib.check.models import ( + Check_Report_AWS, + CheckMetadata, + Code, + Recommendation, + Remediation, +) +from prowler.providers.aws.lib.fix.fixer import AWSFixer + + +def get_mock_aws_finding(): + metadata = CheckMetadata( + Provider="aws", + CheckID="test_check", + CheckTitle="Test Check", + CheckType=["type1"], + CheckAliases=[], + ServiceName="testservice", + SubServiceName="", + ResourceIdTemplate="", + Severity="low", + ResourceType="resource", + Description="desc", + Risk="risk", + RelatedUrl="url", + Remediation=Remediation( + Code=Code(NativeIaC="", Terraform="", CLI="", Other=""), + Recommendation=Recommendation(Text="", Url=""), + ), + Categories=["cat1"], + DependsOn=[], + RelatedTo=[], + Notes="", + Compliance=[], + ) + resource = MagicMock() + resource.id = "res_id" + resource.arn = "arn:aws:test" + resource.region = "eu-west-1" + return Check_Report_AWS(json.dumps(metadata.dict()), resource) + + +class TestAWSFixer: + def test_fix_success(self): + finding = get_mock_aws_finding() + finding.status = "FAIL" + with patch( + "prowler.providers.aws.lib.fix.fixer.AWSFixer.client" + ) as mock_client: + fixer = AWSFixer(description="desc", service="ec2") + mock_client.do_something.return_value = True + assert fixer.fix(finding=finding) + + def test_fix_failure(self, caplog): + fixer = AWSFixer(description="desc", service="ec2") + with patch("prowler.providers.aws.lib.fix.fixer.logger") as mock_logger: + with caplog.at_level("ERROR"): + result = fixer.fix(finding=None) + assert result is False + assert mock_logger.error.called + + def test_get_fixer_info(self): + fixer = AWSFixer( + description="desc", + service="ec2", + cost_impact=True, + cost_description="cost", + iam_policy_required={"Action": ["ec2:DescribeInstances"]}, + ) + info = fixer._get_fixer_info() + assert info["description"] == "desc" + assert info["cost_impact"] is True + assert info["cost_description"] == "cost" + assert info["service"] == "ec2" + assert info["iam_policy_required"] == {"Action": ["ec2:DescribeInstances"]} + + @pytest.mark.parametrize( + "region,resource_id,resource_arn", + [ + ("eu-west-1", "res_id", "arn:aws:test"), + (None, "res_id", None), + ("eu-west-1", None, None), + (None, None, "arn:aws:test"), + (None, None, None), + ], + ) + def test_fix_prints(self, region, resource_id, resource_arn): + fixer = AWSFixer(description="desc", service="ec2") + finding = get_mock_aws_finding() + finding.region = region + finding.resource_id = resource_id + finding.resource_arn = resource_arn + with ( + patch("builtins.print") as mock_print, + patch("prowler.providers.aws.lib.fix.fixer.logger") as mock_logger, + ): + result = fixer.fix(finding=finding) + if region or resource_id or resource_arn: + assert result is True + assert mock_print.called + else: + assert result is False + assert mock_logger.error.called + + def test_fix_exception(self): + fixer = AWSFixer(description="desc", service="ec2") + with patch("prowler.providers.aws.lib.fix.fixer.logger") as mock_logger: + result = fixer.fix(finding=None) + assert result is False + assert mock_logger.error.called diff --git a/tests/providers/azure/lib/fix/azurefixer_test.py b/tests/providers/azure/lib/fix/azurefixer_test.py new file mode 100644 index 0000000000..d8af6f518b --- /dev/null +++ b/tests/providers/azure/lib/fix/azurefixer_test.py @@ -0,0 +1,126 @@ +from unittest.mock import MagicMock, patch + +import pytest + +from prowler.lib.check.models import ( + Check_Report_Azure, + CheckMetadata, + Code, + Recommendation, + Remediation, +) +from prowler.providers.azure.lib.fix.fixer import AzureFixer +from tests.providers.azure.azure_fixtures import set_mocked_azure_provider + + +def get_mock_azure_finding(): + metadata = CheckMetadata( + Provider="azure", + CheckID="test_check", + CheckTitle="Test Check", + CheckType=["type1"], + CheckAliases=[], + ServiceName="testservice", + SubServiceName="", + ResourceIdTemplate="", + Severity="low", + ResourceType="resource", + Description="desc", + Risk="risk", + RelatedUrl="url", + Remediation=Remediation( + Code=Code(NativeIaC="", Terraform="", CLI="", Other=""), + Recommendation=Recommendation(Text="", Url=""), + ), + Categories=["cat1"], + DependsOn=[], + RelatedTo=[], + Notes="", + Compliance=[], + ) + resource = MagicMock() + resource.name = "res_name" + resource.id = "res_id" + resource.location = "westeurope" + return Check_Report_Azure(metadata.dict(), resource) + + +class TestAzureFixer: + def test_fix_success(self): + finding = get_mock_azure_finding() + finding.status = "FAIL" + provider = set_mocked_azure_provider() + with patch( + "prowler.providers.azure.lib.fix.azurefixer.AzureFixer.client" + ) as mock_client: + fixer = AzureFixer(description="desc", service="vm", provider=provider) + mock_client.do_something.return_value = True + assert fixer.fix(finding=finding) + + def test_fix_failure(self, caplog): + finding = get_mock_azure_finding() + finding.status = "FAIL" + provider = set_mocked_azure_provider() + with patch( + "prowler.providers.azure.lib.fix.azurefixer.AzureFixer.client", + side_effect=Exception("fail"), + ): + fixer = AzureFixer(description="desc", service="vm", provider=provider) + with caplog.at_level("ERROR"): + assert not fixer.fix(finding=finding) + assert "fail" in caplog.text + + def test_get_fixer_info(self): + fixer = AzureFixer( + description="desc", + service="vm", + cost_impact=True, + cost_description="cost", + permissions_required={"Action": ["Microsoft.Compute/virtualMachines/read"]}, + ) + info = fixer._get_fixer_info() + assert info["description"] == "desc" + assert info["cost_impact"] is True + assert info["cost_description"] == "cost" + assert info["service"] == "vm" + assert info["permissions_required"] == { + "Action": ["Microsoft.Compute/virtualMachines/read"] + } + + @pytest.mark.parametrize( + "subscription_id,resource_id,resource_group", + [ + ("subid", "res_id", "rg1"), + ("subid", "res_id", None), + ("subid", None, None), + (None, "res_id", None), + (None, None, None), + ], + ) + def test_fix_prints(self, subscription_id, resource_id, resource_group): + fixer = AzureFixer(description="desc", service="vm") + finding = get_mock_azure_finding() + finding.subscription = subscription_id + finding.resource_id = resource_id + finding.resource = ( + {"resource_group_name": resource_group} if resource_group else {} + ) + with ( + patch("builtins.print") as mock_print, + patch("prowler.providers.azure.lib.fix.fixer.logger") as mock_logger, + ): + result = fixer.fix(finding=finding) + if subscription_id or resource_id: + assert result is True + assert mock_print.called + else: + assert result is False + assert mock_logger.error.called + + def test_fix_exception(self): + fixer = AzureFixer(description="desc", service="vm") + with patch("prowler.providers.azure.lib.fix.fixer.logger") as mock_logger: + # Forzar excepción + result = fixer.fix(finding=None, subscription_id=lambda: 1) + assert result is False + assert mock_logger.error.called diff --git a/tests/providers/m365/lib/fix/m365fixer_test.py b/tests/providers/m365/lib/fix/m365fixer_test.py new file mode 100644 index 0000000000..6ac97e571f --- /dev/null +++ b/tests/providers/m365/lib/fix/m365fixer_test.py @@ -0,0 +1,103 @@ +from unittest.mock import MagicMock, patch + +import pytest + +from prowler.lib.check.models import ( + CheckMetadata, + CheckReportM365, + Code, + Recommendation, + Remediation, +) +from prowler.providers.m365.lib.fix.fixer import M365Fixer +from tests.providers.m365.m365_fixtures import set_mocked_m365_provider + + +def get_mock_m365_finding(): + metadata = CheckMetadata( + Provider="m365", + CheckID="test_check", + CheckTitle="Test Check", + CheckType=["type1"], + CheckAliases=[], + ServiceName="testservice", + SubServiceName="", + ResourceIdTemplate="", + Severity="low", + ResourceType="resource", + Description="desc", + Risk="risk", + RelatedUrl="url", + Remediation=Remediation( + Code=Code(NativeIaC="", Terraform="", CLI="", Other=""), + Recommendation=Recommendation(Text="", Url=""), + ), + Categories=["cat1"], + DependsOn=[], + RelatedTo=[], + Notes="", + Compliance=[], + ) + resource = MagicMock() + return CheckReportM365( + metadata.dict(), resource, resource_name="res_name", resource_id="res_id" + ) + + +class TestM365Fixer: + def test_fix_success(self): + finding = get_mock_m365_finding() + finding.status = "FAIL" + provider = set_mocked_m365_provider() + with patch( + "prowler.providers.m365.lib.fix.m365fixer.M365Fixer.client" + ) as mock_client: + fixer = M365Fixer(description="desc", service="mail", provider=provider) + mock_client.do_something.return_value = True + assert fixer.fix(finding=finding) + + def test_fix_failure(self, caplog): + finding = get_mock_m365_finding() + finding.status = "FAIL" + provider = set_mocked_m365_provider() + with patch( + "prowler.providers.m365.lib.fix.m365fixer.M365Fixer.client", + side_effect=Exception("fail"), + ): + fixer = M365Fixer(description="desc", service="mail", provider=provider) + with caplog.at_level("ERROR"): + assert not fixer.fix(finding=finding) + assert "fail" in caplog.text + + def test_get_fixer_info(self): + fixer = M365Fixer( + description="desc", + service="mail", + cost_impact=True, + cost_description="cost", + ) + info = fixer._get_fixer_info() + assert info["description"] == "desc" + assert info["cost_impact"] is True + assert info["cost_description"] == "cost" + assert info["service"] == "mail" + + @pytest.mark.parametrize("resource_id", ["res_id", None]) + def test_fix_prints(self, resource_id): + fixer = M365Fixer(description="desc", service="mail") + finding = get_mock_m365_finding() + finding.resource_id = resource_id + with ( + patch("builtins.print") as mock_print, + patch("prowler.providers.m365.lib.fix.fixer.logger"), + ): + result = fixer.fix(finding=finding) + assert result is True + assert mock_print.called + + def test_fix_exception(self): + fixer = M365Fixer(description="desc", service="mail") + with patch("prowler.providers.m365.lib.fix.fixer.logger") as mock_logger: + result = fixer.fix(finding=None, resource_id=lambda: 1) + assert result is False + assert mock_logger.error.called