Pull rebase from master

This commit is contained in:
Fennerr
2023-12-19 21:55:18 +02:00
parent 78505cb0a8
commit abaa7855d7
2 changed files with 17 additions and 31 deletions
+12 -29
View File
@@ -1,6 +1,4 @@
from concurrent.futures import ThreadPoolExecutor, as_completed
from prowler.lib.logger import logger
from prowler.providers.aws.aws_provider import (
generate_regional_clients,
get_default_region,
@@ -52,34 +50,19 @@ class AWSService:
def __get_session__(self):
return self.session
def __threading_call__(self, call, iterator=None):
def __threading_call__(self, call, iterator=None, max_workers=10):
# Use the provided iterator, or default to self.regional_clients
items = iterator if iterator is not None else self.regional_clients.values()
# Determine the total count for logging
item_count = len(items)
# Trim leading and trailing underscores from the call's name
call_name = call.__name__.strip("_")
# Add Capitalization
call_name = " ".join([x.capitalize() for x in call_name.split("_")])
# Using ThreadPoolExecutor for managing threads
with ThreadPoolExecutor(max_workers=max_workers) as executor:
# Submit tasks to the executor
futures = [executor.submit(call, item) for item in items]
# Print a message based on the call's name, and if its regional or processing a list of items
if iterator is None:
logger.info(
f"{self.service.upper()} - Starting threads for '{call_name}' function across {item_count} regions..."
)
else:
logger.info(
f"{self.service.upper()} - Starting threads for '{call_name}' function to process {item_count} items..."
)
# Submit tasks to the thread pool
futures = [self.thread_pool.submit(call, item) for item in items]
# Wait for all tasks to complete
for future in as_completed(futures):
try:
future.result() # Raises exceptions from the thread, if any
except Exception:
# Handle exceptions if necessary
pass # Replace 'pass' with any additional exception handling logic. Currently handled within the called function
# Wait for all tasks to complete
for future in as_completed(futures):
try:
future.result() # Raises exceptions from the thread, if any
except Exception as e:
# Handle exceptions if necessary
pass # Replace 'pass' with any additional exception handling logic
@@ -18,6 +18,7 @@ class EC2(AWSService):
self.instances = []
self.__threading_call__(self.__describe_instances__)
self.__threading_call__(self.__get_instance_user_data__, self.instances)
self.__threading_call__(self.__get_instance_user_data__, self.instances)
self.security_groups = []
self.regions_with_sgs = []
self.__threading_call__(self.__describe_security_groups__)
@@ -27,7 +28,7 @@ class EC2(AWSService):
self.volumes_with_snapshots = {}
self.regions_with_snapshots = {}
self.__threading_call__(self.__describe_snapshots__)
self.__threading_call__(self.__determine_public_snapshots__, self.snapshots)
self.__threading_call__(self.__get_snapshot_public__, self.snapshots)
self.network_interfaces = []
self.__threading_call__(self.__describe_public_network_interfaces__)
self.__threading_call__(self.__describe_sg_network_interfaces__)
@@ -215,7 +216,8 @@ class EC2(AWSService):
f"{regional_client.region} -- {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def __determine_public_snapshots__(self, snapshot):
def __get_snapshot_public__(self, snapshot):
logger.info("EC2 - Getting snapshot volume attribute permissions...")
try:
regional_client = self.regional_clients[snapshot.region]
snapshot_public = regional_client.describe_snapshot_attribute(
@@ -290,6 +292,7 @@ class EC2(AWSService):
)
def __get_instance_user_data__(self, instance):
logger.info("EC2 - Getting instance user data...")
try:
regional_client = self.regional_clients[instance.region]
user_data = regional_client.describe_instance_attribute(