From 3ee39cff2ae52559c76276339eb545efd2ae29e6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Pedro=20Mart=C3=ADn?= Date: Wed, 9 Oct 2024 11:46:38 +0200 Subject: [PATCH] feat(scan): execute all checks if no checks are provided (#5307) --- prowler/lib/scan/scan.py | 29 +++++++++++++++++++++++++---- tests/lib/scan/scan_test.py | 35 +++++++++++++++++++++++++++++++++++ 2 files changed, 60 insertions(+), 4 deletions(-) diff --git a/prowler/lib/scan/scan.py b/prowler/lib/scan/scan.py index 4de9ca5d4e..27d2005aa7 100644 --- a/prowler/lib/scan/scan.py +++ b/prowler/lib/scan/scan.py @@ -2,6 +2,7 @@ import datetime from typing import Generator from prowler.lib.check.check import execute, import_check, update_audit_metadata +from prowler.lib.check.utils import recover_checks_from_provider from prowler.lib.logger import logger from prowler.lib.outputs.finding import Finding from prowler.providers.common.models import Audit_Metadata @@ -21,7 +22,7 @@ class Scan: _findings: list = [] _duration: int = 0 - def __init__(self, provider: Provider, checks_to_execute: list[str]): + def __init__(self, provider: Provider, checks_to_execute: list[str] = None): """ Scan is the class that executes the checks and yields the progress and the findings. @@ -31,11 +32,31 @@ class Scan: """ self._provider = provider # Remove duplicated checks and sort them - self._checks_to_execute = sorted(list(set(checks_to_execute))) + self._checks_to_execute = ( + sorted(list(set(checks_to_execute))) + if checks_to_execute + else sorted( + [check[0] for check in recover_checks_from_provider(provider.type)] + ) + ) - self._number_of_checks_to_execute = len(checks_to_execute) + # TODO This should be done depending on the scan args (future feature) + # Discard threat detection checks + if "cloudtrail_threat_detection_enumeration" in self._checks_to_execute: + self._checks_to_execute.remove("cloudtrail_threat_detection_enumeration") + if ( + "cloudtrail_threat_detection_privilege_escalation" + in self._checks_to_execute + ): + self._checks_to_execute.remove( + "cloudtrail_threat_detection_privilege_escalation" + ) - service_checks_to_execute = get_service_checks_to_execute(checks_to_execute) + self._number_of_checks_to_execute = len(self._checks_to_execute) + + service_checks_to_execute = get_service_checks_to_execute( + self._checks_to_execute + ) service_checks_completed = dict() self._service_checks_to_execute = service_checks_to_execute diff --git a/tests/lib/scan/scan_test.py b/tests/lib/scan/scan_test.py index 0efe631429..c737039917 100644 --- a/tests/lib/scan/scan_test.py +++ b/tests/lib/scan/scan_test.py @@ -1,3 +1,5 @@ +from importlib.machinery import FileFinder +from pkgutil import ModuleInfo from unittest import mock import pytest @@ -69,6 +71,23 @@ def mock_generate_output(): yield mock_gen_output +@pytest.fixture +def mock_list_modules(): + with mock.patch( + "prowler.lib.check.utils.list_modules", autospec=True + ) as mock_list_mod: + mock_list_mod.return_value = [ + ModuleInfo( + module_finder=FileFinder( + "/prowler/providers/aws/services/accessanalyzer/accessanalyzer_enabled" + ), + name="prowler.providers.aws.services.accessanalyzer.accessanalyzer_enabled.accessanalyzer_enabled", + ispkg=False, + ) + ] + yield mock_list_mod + + class TestScan: def test_init(mock_provider): checks_to_execute = { @@ -205,6 +224,22 @@ class TestScan: assert scan.get_completed_services() == set() assert scan.get_completed_checks() == set() + def test_init_with_no_checks(mock_provider, mock_list_modules): + checks_to_execute = set() + mock_provider.type = "aws" + + scan = Scan(mock_provider, checks_to_execute) + + assert scan.provider == mock_provider + assert scan.checks_to_execute == ["accessanalyzer_enabled"] + assert scan.service_checks_to_execute == get_service_checks_to_execute( + ["accessanalyzer_enabled"] + ) + assert scan.service_checks_completed == {} + assert scan.progress == 0 + assert scan.get_completed_services() == set() + assert scan.get_completed_checks() == set() + @patch("importlib.import_module") def test_scan( mock_import_module,