From 8ace8c01cf6915cbfba3a3ecacaeca51dbfbede4 Mon Sep 17 00:00:00 2001 From: Sergio Garcia <38561120+sergargar@users.noreply.github.com> Date: Thu, 12 Sep 2024 10:56:58 -0400 Subject: [PATCH] chore(refactor): make Provider generation global (#4961) Co-authored-by: pedrooot --- prowler/__main__.py | 2 +- prowler/providers/aws/aws_provider.py | 2 ++ prowler/providers/azure/azure_provider.py | 2 ++ prowler/providers/common/provider.py | 16 ++++++---- prowler/providers/gcp/gcp_provider.py | 2 ++ .../kubernetes/kubernetes_provider.py | 2 ++ .../accessanalyzer_enabled_fixer_test.py | 31 ++++++++++++------- .../guardduty_is_enabled_fixer_test.py | 4 +++ 8 files changed, 42 insertions(+), 19 deletions(-) diff --git a/prowler/__main__.py b/prowler/__main__.py index f153a8a4bb..73477dd3a0 100644 --- a/prowler/__main__.py +++ b/prowler/__main__.py @@ -190,7 +190,7 @@ def prowler(): sys.exit() # Provider to scan - Provider.set_global_provider(args) + Provider.init_global_provider(args) global_provider = Provider.get_global_provider() # Print Provider Credentials diff --git a/prowler/providers/aws/aws_provider.py b/prowler/providers/aws/aws_provider.py index 6cce0bf582..63ea4dda2e 100644 --- a/prowler/providers/aws/aws_provider.py +++ b/prowler/providers/aws/aws_provider.py @@ -270,6 +270,8 @@ class AwsProvider(Provider): # Fixer Config self._fixer_config = fixer_config + Provider.set_global_provider(self) + @property def identity(self): return self._identity diff --git a/prowler/providers/azure/azure_provider.py b/prowler/providers/azure/azure_provider.py index 9515c1351c..4f2de1aa9d 100644 --- a/prowler/providers/azure/azure_provider.py +++ b/prowler/providers/azure/azure_provider.py @@ -160,6 +160,8 @@ class AzureProvider(Provider): # Fixer Config self._fixer_config = fixer_config + Provider.set_global_provider(self) + @property def identity(self): """Returns the identity of the Azure provider.""" diff --git a/prowler/providers/common/provider.py b/prowler/providers/common/provider.py index 4c5ee00ff0..bfffcae492 100644 --- a/prowler/providers/common/provider.py +++ b/prowler/providers/common/provider.py @@ -2,6 +2,7 @@ import importlib import pkgutil import sys from abc import ABC, abstractmethod +from argparse import Namespace from importlib import import_module from typing import Any, Optional @@ -170,7 +171,11 @@ class Provider(ABC): return Provider._global @staticmethod - def set_global_provider(arguments): + def set_global_provider(global_provider: "Provider") -> None: + Provider._global = global_provider + + @staticmethod + def init_global_provider(arguments: Namespace) -> None: try: provider_class_path = ( f"{providers_path}.{arguments.provider}.{arguments.provider}_provider" @@ -188,7 +193,7 @@ class Provider(ABC): if not isinstance(Provider._global, provider_class): if "aws" in provider_class_name.lower(): - global_provider = provider_class( + provider_class( arguments.aws_retries_max_attempts, arguments.role, arguments.session_duration, @@ -205,7 +210,7 @@ class Provider(ABC): fixer_config, ) elif "azure" in provider_class_name.lower(): - global_provider = provider_class( + provider_class( arguments.az_cli_auth, arguments.sp_env_auth, arguments.browser_auth, @@ -217,7 +222,7 @@ class Provider(ABC): fixer_config, ) elif "gcp" in provider_class_name.lower(): - global_provider = provider_class( + provider_class( arguments.project_id, arguments.excluded_project_id, arguments.credentials_file, @@ -227,7 +232,7 @@ class Provider(ABC): fixer_config, ) elif "kubernetes" in provider_class_name.lower(): - global_provider = provider_class( + provider_class( arguments.kubeconfig_file, arguments.context, arguments.namespace, @@ -235,7 +240,6 @@ class Provider(ABC): fixer_config, ) - Provider._global = global_provider except TypeError as error: logger.critical( f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}" diff --git a/prowler/providers/gcp/gcp_provider.py b/prowler/providers/gcp/gcp_provider.py index a04980a096..6115aa8a2d 100644 --- a/prowler/providers/gcp/gcp_provider.py +++ b/prowler/providers/gcp/gcp_provider.py @@ -136,6 +136,8 @@ class GcpProvider(Provider): self._audit_config = audit_config self._fixer_config = fixer_config + Provider.set_global_provider(self) + @property def identity(self): return self._identity diff --git a/prowler/providers/kubernetes/kubernetes_provider.py b/prowler/providers/kubernetes/kubernetes_provider.py index a229cdfa5c..e3e0699dd4 100644 --- a/prowler/providers/kubernetes/kubernetes_provider.py +++ b/prowler/providers/kubernetes/kubernetes_provider.py @@ -80,6 +80,8 @@ class KubernetesProvider(Provider): # Fixer Config self._fixer_config = fixer_config + Provider.set_global_provider(self) + @property def type(self): return self._type diff --git a/tests/providers/aws/services/accessanalyzer/accessanalyzer_enabled/accessanalyzer_enabled_fixer_test.py b/tests/providers/aws/services/accessanalyzer/accessanalyzer_enabled/accessanalyzer_enabled_fixer_test.py index d22394a71c..c1f55114cc 100644 --- a/tests/providers/aws/services/accessanalyzer/accessanalyzer_enabled/accessanalyzer_enabled_fixer_test.py +++ b/tests/providers/aws/services/accessanalyzer/accessanalyzer_enabled/accessanalyzer_enabled_fixer_test.py @@ -4,16 +4,23 @@ AWS_REGION = "eu-west-1" class Test_accessanalyzer_enabled_fixer: - def test_accessanalyzer_enabled_fixer(self): - accessanalyzer_client = mock.MagicMock - accessanalyzer_client.analyzers = [] - with mock.patch( - "prowler.providers.aws.services.accessanalyzer.accessanalyzer_service.AccessAnalyzer", - new=accessanalyzer_client, - ): - # Test Check - from prowler.providers.aws.services.accessanalyzer.accessanalyzer_enabled.accessanalyzer_enabled_fixer import ( - fixer, - ) + @mock.patch( + "prowler.providers.aws.services.accessanalyzer.accessanalyzer_enabled.accessanalyzer_enabled_fixer.accessanalyzer_client" + ) + def test_accessanalyzer_enabled_fixer(self, mock_accessanalyzer_client): + mock_client = mock.MagicMock() + mock_accessanalyzer_client.regional_clients = {AWS_REGION: mock_client} + mock_accessanalyzer_client.fixer_config = { + "accessanalyzer_enabled": { + "AnalyzerName": "DefaultAnalyzer", + "AnalyzerType": "ACCOUNT_UNUSED_ACCESS", + } + } + mock_client.create_analyzer.return_value = None - assert fixer(AWS_REGION) + from prowler.providers.aws.services.accessanalyzer.accessanalyzer_enabled.accessanalyzer_enabled_fixer import ( + fixer, + ) + + result = fixer(AWS_REGION) + assert result diff --git a/tests/providers/aws/services/guardduty/guardduty_is_enabled/guardduty_is_enabled_fixer_test.py b/tests/providers/aws/services/guardduty/guardduty_is_enabled/guardduty_is_enabled_fixer_test.py index bc920390ce..2765af1066 100644 --- a/tests/providers/aws/services/guardduty/guardduty_is_enabled/guardduty_is_enabled_fixer_test.py +++ b/tests/providers/aws/services/guardduty/guardduty_is_enabled/guardduty_is_enabled_fixer_test.py @@ -16,10 +16,14 @@ DETECTOR_ARN = f"arn:aws:guardduty:{AWS_REGION_EU_WEST_1}:{AWS_ACCOUNT_NUMBER}:d class Test_guardduty_is_enabled_fixer: @mock_aws def test_guardduty_is_enabled_fixer(self): + regional_client = mock.MagicMock() guardduty_client = mock.MagicMock guardduty_client.region = AWS_REGION_EU_WEST_1 guardduty_client.detectors = [] guardduty_client.audited_account_arn = AWS_ACCOUNT_ARN + regional_client.create_detector.return_value = None + guardduty_client.regional_clients = {AWS_REGION_EU_WEST_1: regional_client} + with mock.patch( "prowler.providers.aws.services.guardduty.guardduty_service.GuardDuty", guardduty_client,