mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-23 04:21:52 +00:00
chore(refactor): make Provider generation global (#4961)
Co-authored-by: pedrooot <pedromarting3@gmail.com>
This commit is contained in:
+1
-1
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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}"
|
||||
|
||||
@@ -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
|
||||
|
||||
+19
-12
@@ -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
|
||||
|
||||
+4
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user