feat(scan): execute all checks if no checks are provided (#5307)

This commit is contained in:
Pedro Martín
2024-10-09 11:46:38 +02:00
committed by GitHub
parent 41ba118cc4
commit 3ee39cff2a
2 changed files with 60 additions and 4 deletions
+25 -4
View File
@@ -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
+35
View File
@@ -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,