from functools import wraps from typing import Collection, Iterator, List from smelog.factory import BoundLogger from app_types import S3Path from config import S3_ARTIFACTS_BUCKET from utils.key import Key from utils.perf_counter import timeit from .aws import AWSMixin __all__ = ["S3Mixin"] def prefixed(f): @wraps(f) def wrapper(self, path: S3Path = "", *args, **kwargs): if kwargs.pop("prefix", True) and not str(path).startswith(str(self.key)): path = f"{self.key}/{path}" return f(self, path, *args, **kwargs) return wrapper class S3Mixin(AWSMixin): logger: BoundLogger timestamp: int key: Key s3_artifacts_bucket: str = S3_ARTIFACTS_BUCKET def __init__(self, **kwargs): super().__init__(**kwargs) self.s3_client = None @staticmethod def __get_sorting_key(x: str) -> List[int]: x = x.replace(".csv", "").split("_") sort_key = [] for i in range(3, 0, -1): try: sort_key.append(int(x[-i])) except (IndexError, ValueError): continue return sort_key def _pre_run(self): super()._pre_run() self.s3_client = self._boto_session.client("s3") @prefixed def get_prefixed_s3_path(self, path: S3Path) -> str: return path @prefixed def get_full_s3_path(self, path: S3Path) -> str: return f"s3://{self.s3_artifacts_bucket}/{path}" @prefixed def get_batched_s3_keys(self, path: S3Path = "") -> Iterator[Iterator[str]]: paginator = self.s3_client.get_paginator("list_objects") for result in paginator.paginate(Bucket=self.s3_artifacts_bucket, Prefix=path, Delimiter="/"): yield (prefix.get("Prefix") for prefix in result.get("CommonPrefixes", [])) @prefixed def get_s3_keys(self, path: S3Path = "") -> Iterator[str]: keys = [] paginator = self.s3_client.get_paginator("list_objects") for result in paginator.paginate(Bucket=self.s3_artifacts_bucket, Prefix=path): if "Contents" not in result: continue for key in result["Contents"]: keys.append(key["Key"]) keys.sort(key=self.__get_sorting_key) return iter(keys) @prefixed def get_s3_object(self, path: S3Path) -> List[str]: self.logger.debug(f"Reading '{path}'") result = self.get_raw_s3_object(path).decode().splitlines() self.logger.debug(f"Reading done '{path}'") return result @prefixed @timeit() def get_raw_s3_object(self, path: S3Path) -> bytes: return self.s3_client.get_object(Bucket=self.s3_artifacts_bucket, Key=path)["Body"].read() @prefixed @timeit() def put_s3_object(self, path: S3Path, data: bytes): self.logger.debug(f"Writing '{path}'") self.s3_client.put_object(Bucket=self.s3_artifacts_bucket, Key=path, Body=data) self.logger.debug(f"Writing done '{path}'") @prefixed def is_data_available(self, path: S3Path) -> bool: objects = self.s3_client.list_objects(Bucket=self.s3_artifacts_bucket, Prefix=path) return "Contents" in objects @prefixed @timeit() def wipe_folder(self, path: S3Path = "", *, excluded_keys: Collection[S3Path] = ()) -> List[S3Path]: self.logger.info(f"Wiping '{path}'") deleted_keys: List[str] = [] paginator = self.s3_client.get_paginator("list_objects") for result in paginator.paginate(Bucket=self.s3_artifacts_bucket, Prefix=path): if "Contents" not in result: break to_delete = [] for key in result["Contents"]: if key in excluded_keys: continue deleted_keys.append(key["Key"]) to_delete.append({"Key": key["Key"]}) if to_delete: # TODO: check response for errors self.s3_client.delete_objects(Bucket=self.s3_artifacts_bucket, Delete={"Objects": to_delete}) self.logger.info(f"Wiping done '{path}'") return deleted_keys