diff --git a/shared/datastore/src/get_image.py b/shared/datastore/src/get_image.py index 8f92f7d..ef81762 100644 --- a/shared/datastore/src/get_image.py +++ b/shared/datastore/src/get_image.py @@ -7,6 +7,8 @@ from io import BytesIO from minio import Minio from PIL import Image +from shared.utils import check_env + def get_image( client: Minio, @@ -17,11 +19,10 @@ def get_image( assert isinstance(object_name, str) assert len(object_name) > 0 # prepare arguments - assert 'MINIO_BUCKET_NAME' in os.environ - bucket_name = os.getenv('MINIO_BUCKET_NAME', default='') - subfolder = 'images' + check_env({'MINIO_BUCKET_NAME'}) + bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) + object_name = f'images/{object_name}' # get object from bucket - object_name = f'{subfolder}/{object_name}' try: response = client.get_object( bucket_name=bucket_name, diff --git a/shared/datastore/src/get_model.py b/shared/datastore/src/get_model.py index dd6f38d..2cc40e7 100644 --- a/shared/datastore/src/get_model.py +++ b/shared/datastore/src/get_model.py @@ -1,12 +1,14 @@ """Definition of get_model function.""" import logging +import os from collections import OrderedDict +from io import BytesIO import torch from minio import Minio -from .get import get +from shared.utils import check_env def get_model( @@ -16,13 +18,24 @@ def get_model( """Get model from model subfolder in bucket in Minio.""" assert isinstance(client, Minio) assert isinstance(object_name, str) - subfolder = 'models' - object_name = f'{subfolder}/{object_name}' - # get buffer - buffer = get( - client=client, - object_name=object_name, - ) + assert len(object_name) > 0 + # prepare arguments + check_env({'MINIO_BUCKET_NAME_MODELS'}) + bucket_name = str(os.getenv('MINIO_BUCKET_NAME_MODELS')) + object_name = f'models/{object_name}' + # get object from bucket + try: + response = client.get_object( + bucket_name=bucket_name, + object_name=object_name, + ) + buffer = BytesIO(response.data) + except Exception as exc: + logging.error('failed getting data from MinIO') + raise exc + finally: + response.close() + response.release_conn() # convert data to model checkpoint buffer.seek(0) model_content = torch.load(buffer)