fix(s3): enhance threading in s3 service (#4530)

This commit is contained in:
Sergio Garcia
2024-07-26 09:16:47 -04:00
committed by GitHub
parent d244475578
commit f7729381e0
+27 -34
View File
@@ -1,5 +1,4 @@
import json
import threading
from typing import Optional
from botocore.client import ClientError
@@ -18,30 +17,15 @@ class S3(AWSService):
self.account_arn_template = f"arn:{self.audited_partition}:s3:{self.region}:{self.audited_account}:account"
self.regions_with_buckets = []
self.buckets = self.__list_buckets__(provider)
self.__threading_call__(self.__get_bucket_versioning__)
self.__threading_call__(self.__get_bucket_logging__)
self.__threading_call__(self.__get_bucket_policy__)
self.__threading_call__(self.__get_bucket_acl__)
self.__threading_call__(self.__get_public_access_block__)
self.__threading_call__(self.__get_bucket_encryption__)
self.__threading_call__(self.__get_bucket_ownership_controls__)
self.__threading_call__(self.__get_object_lock_configuration__)
self.__threading_call__(self.__get_bucket_tagging__)
# In the S3 service we override the "__threading_call__" method because we spawn a process per bucket instead of per region
# TODO: Replace the above function with the service __threading_call__ using the buckets as the iterator
def __threading_call__(self, call):
threads = []
for bucket in self.buckets:
if bucket.region in self.regional_clients:
regional_client = self.regional_clients[bucket.region]
threads.append(
threading.Thread(target=call, args=(bucket, regional_client))
)
for t in threads:
t.start()
for t in threads:
t.join()
self.__threading_call__(self.__get_bucket_versioning__, self.buckets)
self.__threading_call__(self.__get_bucket_logging__, self.buckets)
self.__threading_call__(self.__get_bucket_policy__, self.buckets)
self.__threading_call__(self.__get_bucket_acl__, self.buckets)
self.__threading_call__(self.__get_public_access_block__, self.buckets)
self.__threading_call__(self.__get_bucket_encryption__, self.buckets)
self.__threading_call__(self.__get_bucket_ownership_controls__, self.buckets)
self.__threading_call__(self.__get_object_lock_configuration__, self.buckets)
self.__threading_call__(self.__get_bucket_tagging__, self.buckets)
def __list_buckets__(self, provider):
logger.info("S3 - Listing buckets...")
@@ -108,9 +92,10 @@ class S3(AWSService):
)
return buckets
def __get_bucket_versioning__(self, bucket, regional_client):
def __get_bucket_versioning__(self, bucket):
logger.info("S3 - Get buckets versioning...")
try:
regional_client = self.regional_clients[bucket.region]
bucket_versioning = regional_client.get_bucket_versioning(
Bucket=bucket.name
)
@@ -139,9 +124,10 @@ class S3(AWSService):
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def __get_bucket_encryption__(self, bucket, regional_client):
def __get_bucket_encryption__(self, bucket):
logger.info("S3 - Get buckets encryption...")
try:
regional_client = self.regional_clients[bucket.region]
bucket.encryption = regional_client.get_bucket_encryption(
Bucket=bucket.name
)["ServerSideEncryptionConfiguration"]["Rules"][0][
@@ -170,9 +156,10 @@ class S3(AWSService):
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def __get_bucket_logging__(self, bucket, regional_client):
def __get_bucket_logging__(self, bucket):
logger.info("S3 - Get buckets logging...")
try:
regional_client = self.regional_clients[bucket.region]
bucket_logging = regional_client.get_bucket_logging(Bucket=bucket.name)
if "LoggingEnabled" in bucket_logging:
bucket.logging = True
@@ -198,9 +185,10 @@ class S3(AWSService):
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def __get_public_access_block__(self, bucket, regional_client):
def __get_public_access_block__(self, bucket):
logger.info("S3 - Get buckets public access block...")
try:
regional_client = self.regional_clients[bucket.region]
public_access_block = regional_client.get_public_access_block(
Bucket=bucket.name
)["PublicAccessBlockConfiguration"]
@@ -240,9 +228,10 @@ class S3(AWSService):
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def __get_bucket_acl__(self, bucket, regional_client):
def __get_bucket_acl__(self, bucket):
logger.info("S3 - Get buckets acl...")
try:
regional_client = self.regional_clients[bucket.region]
grantees = []
acl_grants = regional_client.get_bucket_acl(Bucket=bucket.name)["Grants"]
for grant in acl_grants:
@@ -276,9 +265,10 @@ class S3(AWSService):
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def __get_bucket_policy__(self, bucket, regional_client):
def __get_bucket_policy__(self, bucket):
logger.info("S3 - Get buckets policy...")
try:
regional_client = self.regional_clients[bucket.region]
bucket.policy = json.loads(
regional_client.get_bucket_policy(Bucket=bucket.name)["Policy"]
)
@@ -303,9 +293,10 @@ class S3(AWSService):
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def __get_bucket_ownership_controls__(self, bucket, regional_client):
def __get_bucket_ownership_controls__(self, bucket):
logger.info("S3 - Get buckets ownership controls...")
try:
regional_client = self.regional_clients[bucket.region]
bucket.ownership = regional_client.get_bucket_ownership_controls(
Bucket=bucket.name
)["OwnershipControls"]["Rules"][0]["ObjectOwnership"]
@@ -330,9 +321,10 @@ class S3(AWSService):
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def __get_object_lock_configuration__(self, bucket, regional_client):
def __get_object_lock_configuration__(self, bucket):
logger.info("S3 - Get buckets ownership controls...")
try:
regional_client = self.regional_clients[bucket.region]
regional_client.get_object_lock_configuration(Bucket=bucket.name)
bucket.object_lock = True
except Exception as error:
@@ -359,9 +351,10 @@ class S3(AWSService):
f"{error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}"
)
def __get_bucket_tagging__(self, bucket, regional_client):
def __get_bucket_tagging__(self, bucket):
logger.info("S3 - Get buckets logging...")
try:
regional_client = self.regional_clients[bucket.region]
bucket_tags = regional_client.get_bucket_tagging(Bucket=bucket.name)[
"TagSet"
]