chore(refactor): make Provider generation global (#4961)

Co-authored-by: pedrooot <pedromarting3@gmail.com>
This commit is contained in:
Sergio Garcia
2024-09-12 10:56:58 -04:00
committed by GitHub
parent 8f37252676
commit 8ace8c01cf
8 changed files with 42 additions and 19 deletions
+1 -1
View File
@@ -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
+2
View File
@@ -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
@@ -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."""
+10 -6
View File
@@ -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}"
+2
View File
@@ -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
@@ -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
@@ -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
@@ -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,