diff --git a/shared/datastore/src/datastore_minio.py b/shared/datastore/src/datastore_minio.py new file mode 100644 index 0000000..8771740 --- /dev/null +++ b/shared/datastore/src/datastore_minio.py @@ -0,0 +1,241 @@ +"""Definition of datastore minio implementation.""" + +from __future__ import annotations + +import logging +import os +from collections import OrderedDict +from hashlib import md5 +from io import BytesIO +from traceback import format_exc + +import torch +from minio import Minio +from PIL import Image + +from shared.utils import check_env + +from .datastore_interface import DatastoreInterface + + +class DatastoreMinio(DatastoreInterface): + """Datastore interface.""" + + def __init__(self): + # ensure necessary env vars available + var_list = { + 'MINIO_ENDPOINT', + 'MINIO_ACCESS_KEY', + 'MINIO_SECRET_KEY', + 'MINIO_BUCKET_NAME', + } + check_env(var_list) + # prepare interval variables + self._client: Minio | None = None + self._bucket_name: str | None = None + + def connect(self) -> None: + """Connect to Minio server.""" + # prepare arguments + minio_endpoint = str(os.getenv('MINIO_ENDPOINT')) + minio_access_key = str(os.getenv('MINIO_ACCESS_KEY')) + minio_secret_key = str(os.getenv('MINIO_SECRET_KEY')) + minio_bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) + # connect client + client = Minio( + endpoint=minio_endpoint, + access_key=minio_access_key, + secret_key=minio_secret_key, + secure=False, + ) + # ensure bucket exists + if not client.bucket_exists(bucket_name=minio_bucket_name): + logging.debug('creating bucket: %s', minio_bucket_name) + client.make_bucket(bucket_name=minio_bucket_name) + logging.debug('finished') + # persist state + self._client = client + self._bucket_name = minio_bucket_name + + def close(self) -> None: + """Close connection to Minio server. + + N.B. Minio connection cannot be closed manually. + """ + self._client = None + self._bucket_name = None + + def __enter__(self) -> DatastoreMinio: + self.connect() + return self + + def __exit__(self, exc_type, exc_val, exc_tb) -> None: + if any( + ( + exc_type is not None, + exc_val is not None, + exc_tb is not None, + ), + ): + logging.error('error while exiting context') + self.close() + + def _put( + self, + object_name: str, + buffer: BytesIO, + ) -> None: + """Save in-memory buffer as object in Minio.""" + assert isinstance(object_name, str) + assert len(object_name) > 0 + assert isinstance(buffer, BytesIO) + assert isinstance(self._client, Minio) + assert isinstance(self._bucket_name, str) + # prepare for saving + num_bytes = len(buffer.getvalue()) + buffer.seek(0) + # send data to bucket + try: + self._client.put_object( + bucket_name=self._bucket_name, + object_name=object_name, + length=num_bytes, + data=buffer, + ) + logging.debug('saved data to %s', object_name) + except Exception as exc: + logging.error('failed saving data to MinIO') + raise exc + + def _get( + self, + object_name: str, + ) -> BytesIO: + """Get object from Minio as in-memory buffer.""" + assert isinstance(object_name, str) + assert len(object_name) > 0 + assert isinstance(self._client, Minio) + assert isinstance(self._bucket_name, str) + try: + # make request + response = self._client.get_object( + bucket_name=self._bucket_name, + object_name=object_name, + ) + assert response.status == 200 + # get buffer + buffer = BytesIO() + chunk_size = 2**14 + while chunk := response.read(chunk_size): + buffer.write(chunk) + buffer.seek(0) + logging.debug('got %s', object_name) + return buffer + except Exception as exc: + logging.error('failed getting data from MinIO') + logging.debug(format_exc()) + raise exc + finally: + # close connection if established + if 'response' in locals(): + response.close() + response.release_conn() + + def _delete( + self, + object_name: str, + ) -> None: + """Delete object from Minio.""" + assert isinstance(object_name, str) + assert len(object_name) > 0 + assert isinstance(self._client, Minio) + assert isinstance(self._bucket_name, str) + # remove object + try: + self._client.remove_object( + bucket_name=self._bucket_name, + object_name=object_name, + ) + logging.debug('deleted %s', object_name) + except Exception as exc: + logging.error('failed deleting %s', object_name) + logging.debug(format_exc()) + raise exc + + def put_image( + self, + image: Image.Image, + ) -> str: + """Put image in Minio.""" + assert isinstance(image, Image.Image) + # save data to buffer + buffer = BytesIO() + image.save(buffer, 'png') + # get md5 of buffer + checksum = md5(buffer.getbuffer()).hexdigest() + # build object path + object_path = f'images/{checksum}' + # send data to bucket + self._put( + object_name=object_path, + buffer=buffer, + ) + logging.debug('saved data to %s', object_path) + return checksum + + def get_image( + self, + object_name: str, + ) -> Image.Image: + """Get image from Minio.""" + assert isinstance(object_name, str) + assert len(object_name) > 0 + # build object path + object_path = f'images/{object_name}' + # get object from bucket + buffer = self._get( + object_name=object_path, + ) + # convert data to image + image = Image.open(buffer) + logging.debug('got data from %s', object_path) + return image + + def put_model( + self, + model: torch.nn.Module, + ) -> str: + """Put model in Minio.""" + assert isinstance(model, torch.nn.Module) + # save data to buffer + buffer = BytesIO() + torch.save(model.state_dict(), buffer) + # get md5 of image + checksum = md5(buffer.getbuffer()).hexdigest() + # build object path + object_path = f'models/{checksum}' + # send data to bucket + self._put( + object_name=object_path, + buffer=buffer, + ) + logging.debug('saved data to %s', object_path) + return checksum + + def get_model( + self, + object_name: str, + ) -> OrderedDict: + """Get model data from Minio.""" + assert isinstance(object_name, str) + assert len(object_name) > 0 + # build object path + object_path = f'models/{object_name}' + # get object from bucket + buffer = self._get( + object_name=object_path, + ) + # convert data to model checkpoint + model_content = torch.load(buffer) + logging.debug('got data from %s', object_path) + return model_content