from pathlib import PurePosixPath from typing import Iterable, Tuple, Collection import humanize from airflow.utils.log.logging_mixin import LoggingMixin from common.utils.general import iter_batch from common.utils.types import S3Path __all__ = ["S3", "S3MultiPartUploader"] class S3MultiPartUploader(LoggingMixin): min_part_size = 5242880 # 5MB def __init__(self, conn, bucket_name: str, key: S3Path, **kwargs): super().__init__(**kwargs) self.conn = conn self._bucket_name = bucket_name self._key = str(key) self._upload_id: str | None = None self._parts_info: list[dict[str, str]] = [] self.__buffer: bytes = b"" self.__total_size: int = 0 def __enter__(self): self._create_upload() return self def __exit__(self, exc_type, exc_val, exc_tb): if not exc_val: self._complete_upload() else: self._abort_upload() def _get_next_part_number(self) -> int: return len(self._parts_info) + 1 @staticmethod def _humanize_size(size: int) -> str: return humanize.naturalsize(size, gnu=True) def _get_parts_info(self) -> dict[str, list[dict[str, str]]]: return { "Parts": self._parts_info } def _create_upload(self): self.log.info(f"Prepare multipart upload of '{self._key}'") upload = self.conn.create_multipart_upload(Bucket=self._bucket_name, Key=self._key) self._upload_id = upload["UploadId"] def _complete_upload(self): if self.__buffer: self._upload_part() self.conn.complete_multipart_upload(Bucket=self._bucket_name, Key=self._key, UploadId=self._upload_id, MultipartUpload=self._get_parts_info()) self.log.info(f"Completed multipart upload of '{self._key}'. Total size: {self._humanize_size(self.__total_size)}") def _abort_upload(self): self.log.info(f"Aborting multipart upload of '{self._key}'") self.conn.abort_multipart_upload(Bucket=self._bucket_name, Key=self._key, UploadId=self._upload_id) def _upload_part(self): part_number = self._get_next_part_number() self.log.info(f"Uploading part #{part_number} with size {self._humanize_size(len(self.__buffer))} of '{self._key}'") part = self.conn.upload_part(Bucket=self._bucket_name, Key=self._key, PartNumber=part_number, UploadId=self._upload_id, Body=self.__buffer) self._parts_info.append({ "PartNumber": part_number, "ETag": part["ETag"] }) self.__total_size += len(self.__buffer) self.__buffer = b"" def add_data(self, data: bytes): self.__buffer += data self.log.info(f"Adding data to internal buffer. Current size is {self._humanize_size(len(self.__buffer))}") if len(self.__buffer) >= self.min_part_size: self._upload_part() class S3(LoggingMixin): def __init__(self, conn, run_id: str, bucket_name: str | None = None, **kwargs): super().__init__(**kwargs) self.conn = conn self._run_id = run_id self._bucket_name = bucket_name def _get_bucket_name(self, bucket_name: str | None) -> str: if not bucket_name and not self._bucket_name: raise ValueError("bucket_name should be specified as init argument or as method argument") return bucket_name or self._bucket_name def get_prefixed_key(self, key: S3Path) -> S3Path: return PurePosixPath(self._run_id) / key @staticmethod def _get_prefix(prefix: S3Path) -> str: prefix = str(prefix) if not prefix.endswith("/"): prefix += "/" return prefix def is_data_available(self, prefix: S3Path, bucket_name: str | None = None): prefix = self._get_prefix(prefix) objects = self.conn.list_objects(Bucket=self._get_bucket_name(bucket_name), Prefix=prefix) return "Contents" in objects def get_common_keys(self, prefix: S3Path, bucket_name: str | None = None): prefix = self._get_prefix(prefix) paginator = self.conn.get_paginator("list_objects_v2") response = paginator.paginate(Bucket=self._get_bucket_name(bucket_name), Prefix=prefix, Delimiter="/") for page in response: if "CommonPrefixes" not in page: continue for key in page["CommonPrefixes"]: yield key["Prefix"] def get_keys(self, prefix: S3Path, bucket_name: str | None = None): prefix = self._get_prefix(prefix) paginator = self.conn.get_paginator("list_objects_v2") response = paginator.paginate(Bucket=self._get_bucket_name(bucket_name), Prefix=prefix, Delimiter="/") for page in response: if "Contents" not in page: continue for key in page["Contents"]: yield key["Key"] def get_batched_keys(self, prefix: S3Path, batch_size: int, batch_index: int, bucket_name: str | None = None ) -> Iterable[Tuple[int, int, str]]: paginator = self.conn.get_paginator("list_objects_v2") bucket_name = self._get_bucket_name(bucket_name) for index, part_index in iter_batch(batch_size, batch_index): response = paginator.paginate(Bucket=bucket_name, Prefix=f"{prefix}/{part_index}/", Delimiter="/") for page in response: if "Contents" not in page: continue for key in page["Contents"]: yield index, part_index, key["Key"] def read_object(self, key: S3Path, bucket_name: str | None = None) -> list[str]: self.log.info(f"Reading '{key}'") body = self.conn.get_object(Bucket=self._get_bucket_name(bucket_name), Key=str(key))["Body"] return body.read().decode().splitlines() def write_object(self, key: S3Path, data: bytes, bucket_name: str | None = None): self.log.info(f"Writing '{key}'") self.conn.put_object(Bucket=self._get_bucket_name(bucket_name), Key=str(key), Body=data) def get_multipart_uploader(self, key: S3Path, bucket_name: str | None = None) -> S3MultiPartUploader: return S3MultiPartUploader(conn=self.conn, bucket_name=self._get_bucket_name(bucket_name), key=key) def wipe_by_prefix(self, prefix: S3Path, bucket_name: str | None = None, *, excluded_keys: Collection[S3Path] = ()): self.log.info(f"Wiping '{prefix}'") prefix = self._get_prefix(prefix) bucket_name = self._get_bucket_name(bucket_name) paginator = self.conn.get_paginator("list_objects") for result in paginator.paginate(Bucket=bucket_name, Prefix=prefix): if "Contents" not in result: break to_delete = [] for key in result["Contents"]: if key in excluded_keys: continue to_delete.append({"Key": key["Key"]}) if to_delete: # TODO: check response for errors self.conn.delete_objects(Bucket=bucket_name, Delete={"Objects": to_delete}) self.log.info(f"Wiping done '{prefix}'")