diff --git a/prowler/providers/aws/lib/service/service.py b/prowler/providers/aws/lib/service/service.py index 8f9b2bdcfd..b4e2a27e3d 100644 --- a/prowler/providers/aws/lib/service/service.py +++ b/prowler/providers/aws/lib/service/service.py @@ -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 diff --git a/prowler/providers/aws/services/ec2/ec2_service.py b/prowler/providers/aws/services/ec2/ec2_service.py index bb1863bdab..60a0f3faef 100644 --- a/prowler/providers/aws/services/ec2/ec2_service.py +++ b/prowler/providers/aws/services/ec2/ec2_service.py @@ -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(