From f52720685e1f51686b0a0beb7ff51efd76422e29 Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 11:21:08 +0000 Subject: [PATCH 01/10] fixed relative import --- model/src/models/visual_communication.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/model/src/models/visual_communication.py b/model/src/models/visual_communication.py index 0a841d2..256203b 100644 --- a/model/src/models/visual_communication.py +++ b/model/src/models/visual_communication.py @@ -4,7 +4,7 @@ from __future__ import annotations from torch import nn -from shared.mongodb.classes import ModelData +from shared.mongodb.src.classes import ModelData from .angle import AngleTail from .contact import ContactTail From 7a1e17c29f6cdafb4eaaec8ba0329dd21b3bbc08 Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 11:21:36 +0000 Subject: [PATCH 02/10] updated to raise proper error --- shared/utils/check_env.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/shared/utils/check_env.py b/shared/utils/check_env.py index f9a0e89..d35d23c 100644 --- a/shared/utils/check_env.py +++ b/shared/utils/check_env.py @@ -11,4 +11,5 @@ def check_env( assert all(isinstance(elem, str) for elem in var_list) # check that env vars are set for env_var in var_list: - assert env_var in os.environ, f"environment variable not set: {env_var}" + if env_var not in os.environ: + raise OSError(f"environment variable not set: {env_var}") From 226fb1a206d991eef38ba6925edd58c2e4e1d32f Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 11:22:34 +0000 Subject: [PATCH 03/10] defined interface --- shared/datastore/src/datastore_interface.py | 66 +++++++++++++++++++++ 1 file changed, 66 insertions(+) create mode 100644 shared/datastore/src/datastore_interface.py diff --git a/shared/datastore/src/datastore_interface.py b/shared/datastore/src/datastore_interface.py new file mode 100644 index 0000000..97aa365 --- /dev/null +++ b/shared/datastore/src/datastore_interface.py @@ -0,0 +1,66 @@ +"""Definition of datastore interface.""" + +from abc import ABC, abstractmethod +from collections import OrderedDict + +from PIL import Image +from torch.nn import Module + + +class DatastoreInterface(ABC): + """Datastore interface class.""" + + @abstractmethod + def connect( + self, + ) -> None: + pass + + @abstractmethod + def close( + self, + ) -> None: + pass + + @abstractmethod + def __enter__( + self, + ) -> None: + self.connect() + + @abstractmethod + def __exit__( + self, + exc_type, + exc_val, + exc_tb, + ) -> None: + self.close() + + @abstractmethod + def put_image( + self, + image: Image.Image, + ) -> str: + pass + + @abstractmethod + def get_image( + self, + object_name: str, + ) -> Image.Image: + pass + + @abstractmethod + def put_model( + self, + model: Module, + ) -> str: + pass + + @abstractmethod + def get_model( + self, + object_name: str, + ) -> OrderedDict: + pass From 3b677e5e978fb95f00db5a83ad1b3c1f4af3a1be Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 11:26:23 +0000 Subject: [PATCH 04/10] fixed bug in context manager definition --- shared/datastore/src/datastore_interface.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/shared/datastore/src/datastore_interface.py b/shared/datastore/src/datastore_interface.py index 97aa365..7cd5445 100644 --- a/shared/datastore/src/datastore_interface.py +++ b/shared/datastore/src/datastore_interface.py @@ -1,5 +1,7 @@ """Definition of datastore interface.""" +from __future__ import annotations + from abc import ABC, abstractmethod from collections import OrderedDict @@ -25,8 +27,9 @@ class DatastoreInterface(ABC): @abstractmethod def __enter__( self, - ) -> None: + ) -> DatastoreInterface: self.connect() + return self @abstractmethod def __exit__( From de4f984110bcec6e43f2b779daa839b5f1199032 Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 11:26:47 +0000 Subject: [PATCH 05/10] added minio-specific interface --- shared/datastore/src/datastore_minio.py | 241 ++++++++++++++++++++++++ 1 file changed, 241 insertions(+) create mode 100644 shared/datastore/src/datastore_minio.py 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 From 9ee27b65a84b44f48927d0a616f734fbcc092553 Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 11:27:26 +0000 Subject: [PATCH 06/10] removed unused files --- shared/datastore/__init__.py | 11 +------ shared/datastore/src/__init__.py | 9 +----- shared/datastore/src/connect_minio.py | 38 ----------------------- shared/datastore/src/delete.py | 31 ------------------- shared/datastore/src/get.py | 44 --------------------------- shared/datastore/src/get_image.py | 35 --------------------- shared/datastore/src/get_model.py | 36 ---------------------- shared/datastore/src/put.py | 36 ---------------------- shared/datastore/src/put_image.py | 40 ------------------------ shared/datastore/src/put_model.py | 43 -------------------------- 10 files changed, 2 insertions(+), 321 deletions(-) delete mode 100644 shared/datastore/src/connect_minio.py delete mode 100644 shared/datastore/src/delete.py delete mode 100644 shared/datastore/src/get.py delete mode 100644 shared/datastore/src/get_image.py delete mode 100644 shared/datastore/src/get_model.py delete mode 100644 shared/datastore/src/put.py delete mode 100644 shared/datastore/src/put_image.py delete mode 100644 shared/datastore/src/put_model.py diff --git a/shared/datastore/__init__.py b/shared/datastore/__init__.py index 63e7a47..aec85ae 100644 --- a/shared/datastore/__init__.py +++ b/shared/datastore/__init__.py @@ -1,10 +1 @@ -from .src import ( - connect_minio, - delete, - get, - get_image, - get_model, - put, - put_image, - put_model, -) +from .src import Datastore diff --git a/shared/datastore/src/__init__.py b/shared/datastore/src/__init__.py index f4e0653..2d860b3 100644 --- a/shared/datastore/src/__init__.py +++ b/shared/datastore/src/__init__.py @@ -1,8 +1 @@ -from .connect_minio import connect_minio -from .delete import delete -from .get import get -from .get_image import get_image -from .get_model import get_model -from .put import put -from .put_image import put_image -from .put_model import put_model +from .datastore_minio import DatastoreMinio as Datastore diff --git a/shared/datastore/src/connect_minio.py b/shared/datastore/src/connect_minio.py deleted file mode 100644 index 7dadb4e..0000000 --- a/shared/datastore/src/connect_minio.py +++ /dev/null @@ -1,38 +0,0 @@ -"""Definition of connect function.""" - -import logging -import os - -from minio import Minio - -from shared.utils import check_env - - -def connect_minio() -> Minio: - """Connect to MinIO server.""" - # ensure necessary env vars available - var_list = { - 'MINIO_ENDPOINT', - 'MINIO_ACCESS_KEY', - 'MINIO_SECRET_KEY', - 'MINIO_BUCKET_NAME', - } - check_env(var_list) - # 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.info('creating bucket: %s', minio_bucket_name) - client.make_bucket(bucket_name=minio_bucket_name) - logging.debug('finished') - return client diff --git a/shared/datastore/src/delete.py b/shared/datastore/src/delete.py deleted file mode 100644 index 6bd926a..0000000 --- a/shared/datastore/src/delete.py +++ /dev/null @@ -1,31 +0,0 @@ -"""Definition of delete function.""" - -import logging -from traceback import print_exc - -from minio import Minio - - -def delete( - client: Minio, - bucket_name: str, - object_name: str, -) -> None: - """Delete object from MinIO.""" - assert isinstance(client, Minio) - assert isinstance(bucket_name, str) - assert len(bucket_name) > 0 - assert isinstance(object_name, str) - assert len(object_name) > 0 - # remove object - try: - client.remove_object( - bucket_name=bucket_name, - object_name=object_name, - ) - except Exception as exc: - logging.error('failed deleting %s', object_name) - print_exc() - raise exc - else: - logging.debug('deleted %s', object_name) diff --git a/shared/datastore/src/get.py b/shared/datastore/src/get.py deleted file mode 100644 index 09ce2c6..0000000 --- a/shared/datastore/src/get.py +++ /dev/null @@ -1,44 +0,0 @@ -"""Definition of get function.""" - -import logging -from io import BytesIO -from traceback import print_exc - -from minio import Minio - - -def get( - client: Minio, - bucket_name: str, - object_name: str, -) -> BytesIO: - """Get buffer from bucket in MinIO.""" - assert isinstance(client, Minio) - assert isinstance(bucket_name, str) - assert len(bucket_name) > 0 - assert isinstance(object_name, str) - assert len(object_name) > 0 - try: - # make request - response = client.get_object( - bucket_name=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') - print_exc() - raise exc - finally: - # close connection if established - if 'response' in locals(): - response.close() - response.release_conn() diff --git a/shared/datastore/src/get_image.py b/shared/datastore/src/get_image.py deleted file mode 100644 index 596226b..0000000 --- a/shared/datastore/src/get_image.py +++ /dev/null @@ -1,35 +0,0 @@ -"""Definition of get_image function.""" - -import logging -import os - -from minio import Minio -from PIL import Image - -from shared.utils import check_env - -from .get import get - - -def get_image( - client: Minio, - object_name: str, -) -> Image.Image: - """Get image from image subfolder in bucket in Minio.""" - assert isinstance(client, Minio) - assert isinstance(object_name, str) - assert len(object_name) > 0 - # prepare arguments - check_env({'MINIO_BUCKET_NAME'}) - bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) - object_name = f'images/{object_name}' - # get object from bucket - buffer = get( - client=client, - bucket_name=bucket_name, - object_name=object_name, - ) - # convert data to image - image = Image.open(buffer) - logging.debug('got data from %s', object_name) - return image diff --git a/shared/datastore/src/get_model.py b/shared/datastore/src/get_model.py deleted file mode 100644 index 59c349a..0000000 --- a/shared/datastore/src/get_model.py +++ /dev/null @@ -1,36 +0,0 @@ -"""Definition of get_model function.""" - -import logging -import os -from collections import OrderedDict - -import torch -from minio import Minio - -from shared.utils import check_env - -from .get import get - - -def get_model( - client: Minio, - object_name: str, -) -> OrderedDict: - """Get model from model subfolder in bucket in Minio.""" - assert isinstance(client, Minio) - assert isinstance(object_name, str) - assert len(object_name) > 0 - # prepare arguments - check_env({'MINIO_BUCKET_NAME'}) - bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) - object_name = f'models/{object_name}' - # get object from bucket - buffer = get( - client=client, - bucket_name=bucket_name, - object_name=object_name, - ) - # convert data to model checkpoint - model_content = torch.load(buffer) - logging.debug('finished') - return model_content diff --git a/shared/datastore/src/put.py b/shared/datastore/src/put.py deleted file mode 100644 index 372a1d7..0000000 --- a/shared/datastore/src/put.py +++ /dev/null @@ -1,36 +0,0 @@ -"""Definition of put function.""" - -import logging -from io import BytesIO - -from minio import Minio - - -def put( - client: Minio, - buffer: BytesIO, - bucket_name: str, - object_name: str, -) -> None: - """Put buffer in bucket in MinIO and return MD5 checksum as object name.""" - assert isinstance(client, Minio) - assert isinstance(buffer, BytesIO) - assert isinstance(bucket_name, str) - assert len(bucket_name) > 0 - assert isinstance(object_name, str) - assert len(object_name) > 0 - # prepare for saving - num_bytes = len(buffer.getvalue()) - buffer.seek(0) - # send data to bucket - try: - client.put_object( - bucket_name=bucket_name, - object_name=object_name, - length=num_bytes, - data=buffer, - ) - except Exception as exc: - logging.error('failed saving data to MinIO') - raise exc - logging.debug('saved data to %s', object_name) diff --git a/shared/datastore/src/put_image.py b/shared/datastore/src/put_image.py deleted file mode 100644 index bea9dd6..0000000 --- a/shared/datastore/src/put_image.py +++ /dev/null @@ -1,40 +0,0 @@ -"""Definition of put_image function.""" - -import logging -import os -from hashlib import md5 -from io import BytesIO - -from minio import Minio -from PIL import Image - -from .put import put - - -def put_image( - client: Minio, - image: Image.Image, -) -> str: - """Put image in image subfolder in bucket in Minio and return MD5 checksum - used as object name.""" - assert isinstance(client, Minio) - assert isinstance(image, Image.Image) - # get bucket name from env - bucket_name = os.getenv('MINIO_BUCKET_NAME', default='') - assert len(bucket_name) > 0 - # save data to buffer - buffer = BytesIO() - image.save(buffer, 'png') - # get md5 of buffer - checksum = md5(buffer.getbuffer()).hexdigest() - # set object name - object_name = f'images/{checksum}' - # send data to bucket - put( - client=client, - buffer=buffer, - bucket_name=bucket_name, - object_name=object_name, - ) - logging.debug('finished') - return checksum diff --git a/shared/datastore/src/put_model.py b/shared/datastore/src/put_model.py deleted file mode 100644 index 69436a9..0000000 --- a/shared/datastore/src/put_model.py +++ /dev/null @@ -1,43 +0,0 @@ -"""Definition of put_model function.""" - -import logging -import os -from hashlib import md5 -from io import BytesIO - -import torch -from minio import Minio -from torch.nn import Module - -from shared.utils import check_env - -from .put import put - - -def put_model( - client: Minio, - model: Module, -) -> str: - """Put model in model subfolder in bucket in Minio and return MD5 checksum - used as object name.""" - assert isinstance(client, Minio) - assert isinstance(model, Module) - # get bucket name from env - check_env({'MINIO_BUCKET_NAME'}) - bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) - # save data to buffer - buffer = BytesIO() - torch.save(model.state_dict(), buffer) - # get md5 of image - checksum = md5(buffer.getbuffer()).hexdigest() - # set object name - object_name = f'models/{checksum}' - # send data to bucket - put( - client=client, - buffer=buffer, - bucket_name=bucket_name, - object_name=object_name, - ) - logging.debug('finished') - return checksum From f28bff7e4759f0668b759035834d040b7be7f3d9 Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 11:29:31 +0000 Subject: [PATCH 07/10] implemented new interface --- .../src/classes/visual_communication.py | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/shared/mongodb/src/classes/visual_communication.py b/shared/mongodb/src/classes/visual_communication.py index a21d0b4..2291e6b 100755 --- a/shared/mongodb/src/classes/visual_communication.py +++ b/shared/mongodb/src/classes/visual_communication.py @@ -12,7 +12,7 @@ from PIL import Image from pydantic import BaseModel, ConfigDict from pymongo.collection import Collection -from shared.datastore import get_image, put_image +from shared.datastore import Datastore from shared.mongodb.src.classes import ModelData @@ -39,10 +39,10 @@ class VisualCommunication(BaseModel): """Upload image to MinIO and return MD5 checksum of hashed image.""" assert isinstance(image, Image.Image) assert isinstance(minio_client, Minio) - object_name = put_image( - client=minio_client, - image=image, - ) + with Datastore() as ds: + object_name = ds.put_image( + image=image, + ) return object_name @classmethod @@ -91,10 +91,10 @@ class VisualCommunication(BaseModel): """Load image data from minio.""" assert isinstance(minio_client, Minio) # get image from minio - image = get_image( - client=minio_client, - object_name=self.object_name, - ) + with Datastore() as ds: + image = ds.get_image( + object_name=self.object_name, + ) return image def save_to_mongo(self, collection: Collection) -> None: From 88698dd73505a10c0021e13323e68db90e186342 Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 11:29:54 +0000 Subject: [PATCH 08/10] fixed relative import paths --- shared/mongodb/src/get_visual_communication.py | 4 ++-- shared/mongodb/src/upsert_annotation.py | 2 +- shared/mongodb/src/upsert_prediction.py | 2 +- shared/mongodb/src/upsert_visual_communication.py | 2 +- 4 files changed, 5 insertions(+), 5 deletions(-) diff --git a/shared/mongodb/src/get_visual_communication.py b/shared/mongodb/src/get_visual_communication.py index b9c0dbb..57a4071 100755 --- a/shared/mongodb/src/get_visual_communication.py +++ b/shared/mongodb/src/get_visual_communication.py @@ -6,8 +6,8 @@ import logging from pymongo.collection import Collection -from shared.mongodb.classes import VisualCommunication -from shared.mongodb.exceptions import NoDocumentFoundException +from shared.mongodb.src.classes import VisualCommunication +from shared.mongodb.src.exceptions import NoDocumentFoundException def get_visual_communication( diff --git a/shared/mongodb/src/upsert_annotation.py b/shared/mongodb/src/upsert_annotation.py index 9e61bcc..d7c8c4d 100755 --- a/shared/mongodb/src/upsert_annotation.py +++ b/shared/mongodb/src/upsert_annotation.py @@ -4,7 +4,7 @@ import logging from pymongo.collection import Collection -from shared.mongodb.classes import ModelData +from shared.mongodb.src.classes import ModelData def upsert_annotation( diff --git a/shared/mongodb/src/upsert_prediction.py b/shared/mongodb/src/upsert_prediction.py index 0f4d628..7f1ca8f 100755 --- a/shared/mongodb/src/upsert_prediction.py +++ b/shared/mongodb/src/upsert_prediction.py @@ -4,7 +4,7 @@ import logging from pymongo.collection import Collection -from shared.mongodb.classes import ModelData +from shared.mongodb.src.classes import ModelData def upsert_prediction( diff --git a/shared/mongodb/src/upsert_visual_communication.py b/shared/mongodb/src/upsert_visual_communication.py index a69b9fa..8ea39af 100755 --- a/shared/mongodb/src/upsert_visual_communication.py +++ b/shared/mongodb/src/upsert_visual_communication.py @@ -2,7 +2,7 @@ from __future__ import annotations from pymongo.collection import Collection -from shared.mongodb.classes import VisualCommunication +from shared.mongodb.src.classes import VisualCommunication def upsert_visual_communication( From 16ad2b80ee6bd4ea9e53247726574f1749134447 Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 11:34:02 +0000 Subject: [PATCH 09/10] updated tests to match new interface --- .../tests/integration/base_crud_test.py | 63 +++-- .../datastore/tests/integration/conftest.py | 130 +++++++--- .../tests/integration/connect_test.py | 6 +- .../tests/integration/image_crud_test.py | 95 ++++---- .../tests/integration/model_crud_test.py | 126 ++++++---- .../tests/unit/connect_minio_test.py | 33 --- .../tests/unit/datastore_minio_test.py | 230 ++++++++++++++++++ shared/datastore/tests/unit/delete_test.py | 61 ----- shared/datastore/tests/unit/get_image_test.py | 63 ----- shared/datastore/tests/unit/get_model_test.py | 63 ----- shared/datastore/tests/unit/get_test.py | 60 ----- shared/datastore/tests/unit/put_image_test.py | 41 ---- shared/datastore/tests/unit/put_model_test.py | 40 --- shared/datastore/tests/unit/put_test.py | 80 ------ 14 files changed, 495 insertions(+), 596 deletions(-) delete mode 100644 shared/datastore/tests/unit/connect_minio_test.py create mode 100644 shared/datastore/tests/unit/datastore_minio_test.py delete mode 100644 shared/datastore/tests/unit/delete_test.py delete mode 100644 shared/datastore/tests/unit/get_image_test.py delete mode 100644 shared/datastore/tests/unit/get_model_test.py delete mode 100644 shared/datastore/tests/unit/get_test.py delete mode 100644 shared/datastore/tests/unit/put_image_test.py delete mode 100644 shared/datastore/tests/unit/put_model_test.py delete mode 100644 shared/datastore/tests/unit/put_test.py diff --git a/shared/datastore/tests/integration/base_crud_test.py b/shared/datastore/tests/integration/base_crud_test.py index dda9ee2..723ffc7 100644 --- a/shared/datastore/tests/integration/base_crud_test.py +++ b/shared/datastore/tests/integration/base_crud_test.py @@ -7,7 +7,7 @@ from io import BytesIO import minio import pytest -from shared.datastore import delete, get, put +from shared.datastore import Datastore def same_data( @@ -39,77 +39,72 @@ def same_data( def test_should_get_data( - minio_client, + datastore: Datastore, data_in_minio, ): - data, bucket_name, object_name = data_in_minio - received_data = get( - client=minio_client, - bucket_name=bucket_name, + # ARRANGE + data, _, object_name = data_in_minio + # ACT + received_data = datastore._get( object_name=object_name, ) - + # ASSERT assert isinstance(data, BytesIO) assert same_data(data, received_data) def test_should_delete_data( - minio_client, + datastore: Datastore, data_in_minio, ): - _, bucket_name, object_name = data_in_minio - delete( - client=minio_client, - bucket_name=bucket_name, + # ARRANGE + _, _, object_name = data_in_minio + # ACT + datastore._delete( object_name=object_name, ) + # ASSERT with pytest.raises(minio.error.S3Error): - _ = get( - client=minio_client, - bucket_name=bucket_name, + _ = datastore._get( object_name=object_name, ) def test_should_put_data( - minio_client, + datastore: Datastore, data, ): - # prepare variables - minio_bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) + # ARRANGE minio_object_name = str(os.getenv('MINIO_OBJECT_NAME')) buffer = BytesIO(data) - put( - client=minio_client, + # ACT + datastore._put( + object_name=minio_object_name, buffer=buffer, - bucket_name=minio_bucket_name, - object_name=minio_object_name, ) - received_data = get( - client=minio_client, - bucket_name=minio_bucket_name, + received_data = datastore._get( object_name=minio_object_name, ) + # ASSERT assert isinstance(received_data, BytesIO) assert same_data(received_data, buffer) def test_should_update_data( - minio_client, + datastore: Datastore, data_in_minio, ): - buffer, bucket_name, object_name = data_in_minio - put( - client=minio_client, + # ARRANGE + buffer, _, object_name = data_in_minio + # ACT + datastore._put( + object_name=object_name, buffer=buffer, - bucket_name=bucket_name, - object_name=object_name, ) - received_data = get( - client=minio_client, - bucket_name=bucket_name, + received_data = datastore._get( object_name=object_name, ) + # ASSERT assert same_data(received_data, buffer) diff --git a/shared/datastore/tests/integration/conftest.py b/shared/datastore/tests/integration/conftest.py index 342ebae..741fb25 100644 --- a/shared/datastore/tests/integration/conftest.py +++ b/shared/datastore/tests/integration/conftest.py @@ -3,14 +3,18 @@ import os import random from collections.abc import Iterator +from hashlib import md5 from io import BytesIO from pathlib import Path import pytest +import torch from dotenv import load_dotenv -from minio import Minio from PIL import Image +from model.src.models import VisualCommunicationModel +from shared.datastore import Datastore + env_var_map = { 'MINIO_BUCKET_NAME': 'test-bucket', 'MINIO_OBJECT_NAME': '46KXJMFIAPVLM0TKRFZR5YPPTVJ6PJNX', @@ -39,35 +43,34 @@ def populate_env( @pytest.fixture(scope='session') -def minio_client( +def datastore( populate_env, -) -> Iterator[Minio]: +) -> Iterator[Datastore]: # 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 to minio - client = Minio( - endpoint=minio_endpoint, - access_key=minio_access_key, - secret_key=minio_secret_key, - secure=False, - ) + datastore_client = Datastore() + datastore_client.connect() + assert datastore_client._client is not None # ensure bucket exists - if not client.bucket_exists(bucket_name=minio_bucket_name): - client.make_bucket(bucket_name=minio_bucket_name) + if not datastore_client._client.bucket_exists(bucket_name=minio_bucket_name): + datastore_client._client.make_bucket(bucket_name=minio_bucket_name) # expose client - yield client + yield datastore_client # remove objects left behind by tests - for obj in client.list_objects(bucket_name=minio_bucket_name, recursive=True): - client.remove_object( + for obj in datastore_client._client.list_objects( + bucket_name=minio_bucket_name, + recursive=True, + ): + datastore_client._client.remove_object( bucket_name=obj.bucket_name, object_name=obj.object_name, ) # remove bucket - client.remove_bucket(bucket_name=minio_bucket_name) - assert not client.bucket_exists(bucket_name=minio_bucket_name) + datastore_client._client.remove_bucket(bucket_name=minio_bucket_name) + assert not datastore_client._client.bucket_exists(bucket_name=minio_bucket_name) + # disconnect from minio + datastore_client.close() @pytest.fixture @@ -81,9 +84,10 @@ def data() -> Iterator[bytes]: @pytest.fixture def data_in_minio( - minio_client, - data, + datastore: Datastore, + data: bytes, ) -> Iterator[tuple[BytesIO, str, str]]: + assert datastore._client is not None # prepare arguments minio_bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) minio_object_name = str(os.getenv('MINIO_OBJECT_NAME')) @@ -93,7 +97,7 @@ def data_in_minio( num_bytes = len(buffer.getvalue()) buffer.seek(0) # send data to bucket - minio_client.put_object( + datastore._client.put_object( bucket_name=minio_bucket_name, object_name=minio_object_name, length=num_bytes, @@ -102,7 +106,7 @@ def data_in_minio( # expose data yield buffer, minio_bucket_name, minio_object_name # clean up - minio_client.remove_object( + datastore._client.remove_object( bucket_name=minio_bucket_name, object_name=minio_object_name, ) @@ -116,13 +120,77 @@ def image() -> Iterator[Image.Image]: yield image -# @pytest.fixture -# def image_in_minio( -# image: Image.Image, -# ) -> tuple(Image.Image, str): -# # +@pytest.fixture +def image_in_minio( + datastore: Datastore, + image: Image.Image, +) -> Iterator[tuple[Image.Image, str]]: + assert datastore._client is not None + # prepare arguments + minio_bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) + # 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}' + # prepare for saving + num_bytes = len(buffer.getvalue()) + buffer.seek(0) + # send data to bucket + datastore._client.put_object( + bucket_name=minio_bucket_name, + object_name=object_path, + length=num_bytes, + data=buffer, + ) + # expose image and object name + yield image, checksum + # cleanup + datastore._client.remove_object( + bucket_name=minio_bucket_name, + object_name=object_path, + ) -# # expose image and object name -# yield image, object_name -# # cleanup +@pytest.fixture +def model() -> Iterator[torch.nn.Module]: + # generate model + model = VisualCommunicationModel().to('cpu') + # expose model + yield model + + +@pytest.fixture +def model_in_minio( + datastore: Datastore, + model: torch.nn.Module, +) -> Iterator[tuple[torch.nn.Module, str]]: + assert datastore._client is not None + # prepare arguments + minio_bucket_name = str(os.getenv('MINIO_BUCKET_NAME')) + # 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}' + # prepare for saving + num_bytes = len(buffer.getvalue()) + buffer.seek(0) + # send data to bucket + datastore._client.put_object( + bucket_name=minio_bucket_name, + object_name=object_path, + length=num_bytes, + data=buffer, + ) + # expose model and object name + yield model, checksum + # cleanup + datastore._client.remove_object( + bucket_name=minio_bucket_name, + object_name=object_path, + ) diff --git a/shared/datastore/tests/integration/connect_test.py b/shared/datastore/tests/integration/connect_test.py index 25b5aaf..ea5868e 100644 --- a/shared/datastore/tests/integration/connect_test.py +++ b/shared/datastore/tests/integration/connect_test.py @@ -2,9 +2,9 @@ from minio import Minio -from shared.datastore import connect_minio +from shared.datastore import Datastore def test_should_return_correct_type(): - client = connect_minio() - assert isinstance(client, Minio) + with Datastore() as ds: + assert isinstance(ds._client, Minio) diff --git a/shared/datastore/tests/integration/image_crud_test.py b/shared/datastore/tests/integration/image_crud_test.py index e91c9a0..c3b0dbb 100644 --- a/shared/datastore/tests/integration/image_crud_test.py +++ b/shared/datastore/tests/integration/image_crud_test.py @@ -1,9 +1,10 @@ """Integration tests related to image CRUD.""" import numpy as np +import pytest from PIL import Image -# from shared.datastore import connect_minio, get_image, put_image +from shared.datastore import Datastore def same_image( @@ -25,49 +26,57 @@ def same_image( return True -# def test_should_get_image( -# image_in_minio, -# ): -# image, object_name = image_in_minio -# client = connect_minio() -# received_image = get_image( -# client=client, -# object_name=object_name, -# ) -# assert isinstance(image, Image.Image) -# assert received_image == image +def test_should_get_image( + datastore: Datastore, + image_in_minio: tuple[Image.Image, str], +): + # ARRANGE + image, object_name = image_in_minio + # ACT + received_image = datastore.get_image( + object_name=object_name, + ) + # ASSERT + assert isinstance(image, Image.Image) + assert same_image(image, received_image) -# def test_should_put_image( -# image, -# ): -# client = connect_minio() -# object_name = put_image( -# client=client, -# image=image, -# ) -# assert isinstance(object_name, str) -# assert len(object_name) > 0 -# received_image = get_image( -# client=client, -# object_name=object_name, -# ) -# assert same_image(received_image, image) +def test_should_put_image( + datastore: Datastore, + image: Image.Image, +): + # ARRANGE + object_name = datastore.put_image( + image=image, + ) + assert isinstance(object_name, str) + assert len(object_name) > 0 + # ACT + received_image = datastore.get_image( + object_name=object_name, + ) + # ASSERT + assert same_image(image, received_image) -# def test_should_update_image( -# image_in_minio, -# ): -# image, object_name = image_in_minio -# client = connect_minio() -# object_name = put_image( -# client=client, -# image=image, -# ) -# assert isinstance(object_name, str) -# assert len(object_name) > 0 -# received_image = get_image( -# client=client, -# object_name=object_name, -# ) -# assert received_image == image +def test_should_update_image( + datastore: Datastore, + image_in_minio: tuple[Image.Image, str], +): + # ARRANGE + image, object_name = image_in_minio + object_name = datastore.put_image( + image=image, + ) + assert isinstance(object_name, str) + assert len(object_name) > 0 + # ACT + received_image = datastore.get_image( + object_name=object_name, + ) + # ASSERT + assert same_image(image, received_image) + + +if __name__ == '__main__': + pytest.main() diff --git a/shared/datastore/tests/integration/model_crud_test.py b/shared/datastore/tests/integration/model_crud_test.py index a5e126a..2f41472 100644 --- a/shared/datastore/tests/integration/model_crud_test.py +++ b/shared/datastore/tests/integration/model_crud_test.py @@ -1,52 +1,90 @@ """Integration tests related to model CRUD.""" -# from torch.nn import Module +from collections import OrderedDict -# from shared.datastore import connect_minio, get_model, put_model +from torch.nn import Module -# def test_should_get_model( -# model_in_minio, -# ): -# model, object_name = model_in_minio -# client = connect_minio() -# received_model = get_model( -# client=client, -# object_name=object_name, -# ) -# assert isinstance(model, Module) -# assert received_model == model +from model.src.models import VisualCommunicationModel +from shared.datastore import Datastore -# def test_should_put_model( -# model, -# ): -# client = connect_minio() -# object_name = put_model( -# client=client, -# model=model, -# ) -# assert isinstance(object_name, str) -# assert len(object_name) > 0 -# received_model = get_model( -# client=client, -# object_name=object_name, -# ) -# assert received_model == model +def same_model( + model_a: Module, + model_b: Module, +) -> bool: + """Check if two models are the same class, have the same number of + parameters and contain the same weights.""" + assert isinstance(model_a, Module) + assert isinstance(model_b, Module) + # compare model classes + assert type(model_a) is type(model_b) + # compare number of parameters + params_a = list(model_a.parameters()) + params_b = list(model_b.parameters()) + if len(params_a) != len(params_b): + return False + # compare model weights + for p_a, p_b in zip(params_a, params_b): + if p_a.data.ne(p_b.data).sum() > 0: + return False + return True -# def test_should_update_model( -# model_in_minio, -# ): -# model, object_name = model_in_minio -# client = connect_minio() -# object_name = put_model( -# client=client, -# model=model, -# ) -# assert isinstance(object_name, str) -# assert len(object_name) > 0 -# received_model = get_model( -# client=client, -# object_name=object_name, -# ) -# assert received_model == model +def test_should_get_model( + datastore: Datastore, + model_in_minio: tuple[Module, str], +): + # ARRANGE + model, object_name = model_in_minio + # ACT + model_data = datastore.get_model( + object_name=object_name, + ) + assert isinstance(model_data, OrderedDict) + received_model = VisualCommunicationModel().to('cpu') + received_model.load_state_dict(model_data) + # ASSERT + assert same_model(model, received_model) + + +def test_should_put_model( + datastore: Datastore, + model: Module, +): + # ARRANGE + object_name = datastore.put_model( + model=model, + ) + assert isinstance(object_name, str) + assert len(object_name) > 0 + # ACT + model_data = datastore.get_model( + object_name=object_name, + ) + assert isinstance(model_data, OrderedDict) + received_model = VisualCommunicationModel().to('cpu') + received_model.load_state_dict(model_data) + # ASSERT + assert same_model(model, received_model) + + +def test_should_update_model( + datastore: Datastore, + model_in_minio: tuple[Module, str], +): + # ARRANGE + model, object_name = model_in_minio + object_name = datastore.put_model( + model=model, + ) + assert isinstance(object_name, str) + assert len(object_name) > 0 + # ACT + model_data = datastore.get_model( + object_name=object_name, + ) + assert isinstance(model_data, OrderedDict) + received_model = VisualCommunicationModel().to('cpu') + received_model.load_state_dict(model_data) + # ASSERT + assert same_model(model, received_model) diff --git a/shared/datastore/tests/unit/connect_minio_test.py b/shared/datastore/tests/unit/connect_minio_test.py deleted file mode 100644 index 508ac60..0000000 --- a/shared/datastore/tests/unit/connect_minio_test.py +++ /dev/null @@ -1,33 +0,0 @@ -"""Definition of unittests for connect_minio function.""" - -import os -import unittest - -from shared.datastore import connect_minio - - -class TestConnectMinio(unittest.TestCase): - - def setUp(self): - # define relevant env vars - self.env_var_map = { - 'MINIO_ENDPOINT': '192.168.1.2', - 'MINIO_ACCESS_KEY': 'randomAccess_key', - 'MINIO_SECRET_KEY': 'randomSecret_key', - 'MINIO_BUCKET_NAME': 'test-bucket-name', - } - # set env vars - for key, val in self.env_var_map.items(): - os.environ[key] = val - - def tearDown(self): - # clear env vars - for key in self.env_var_map: - _ = os.environ.pop(key, default=None) - - def test_should_fail_when_env_not_set(self): - # ensure env not set - self.tearDown() - # run test - with self.assertRaises(AssertionError): - _ = connect_minio() diff --git a/shared/datastore/tests/unit/datastore_minio_test.py b/shared/datastore/tests/unit/datastore_minio_test.py new file mode 100644 index 0000000..2ce55c3 --- /dev/null +++ b/shared/datastore/tests/unit/datastore_minio_test.py @@ -0,0 +1,230 @@ +"""Definition of unittests for Datastore instantiation.""" + +import os +from hashlib import md5 +from io import BytesIO +from unittest import TestCase +from unittest.mock import ANY, MagicMock, patch + +from minio import Minio +from PIL import Image +from urllib3 import BaseHTTPResponse + +from shared.datastore.src.datastore_minio import DatastoreMinio + + +class TestDatastoreMinioInstantiation(TestCase): + + def setUp(self): + # define relevant env vars + self.env_var_map = { + 'MINIO_ENDPOINT': '192.168.1.2', + 'MINIO_ACCESS_KEY': 'randomAccess_key', + 'MINIO_SECRET_KEY': 'randomSecret_key', + 'MINIO_BUCKET_NAME': 'test-bucket-name', + } + # set env vars + for key, val in self.env_var_map.items(): + os.environ[key] = val + # set other variables + self.object_name = '46KXJMFIAPVLM0TKRFZR5YPPTVJ6PJNX' + self.image = Image.new(mode='RGB', size=(480, 480)) + self.buffer = BytesIO() + self.image.save(self.buffer, 'png') + self.num_bytes = len(self.buffer.getvalue()) + self.checksum = md5(self.buffer.getbuffer()).hexdigest() + + def tearDown(self): + # clear env vars + for key in self.env_var_map: + _ = os.environ.pop(key, default=None) + + def test_instantiation_should_fail_when_env_not_set(self): + # ensure env not set + self.tearDown() + # run test + with self.assertRaises(OSError): + _ = DatastoreMinio() + + @patch('shared.datastore.src.datastore_minio.Minio') + def test_connect_should_call_Minio_with_env_vars(self, minio_mock): + # ARRANGE + datastore = DatastoreMinio() + minio_mock().bucket_exists.return_value = False + # ACT + datastore.connect() + # ASSERT + minio_mock.assert_called_with( + endpoint=self.env_var_map['MINIO_ENDPOINT'], + access_key=self.env_var_map['MINIO_ACCESS_KEY'], + secret_key=self.env_var_map['MINIO_SECRET_KEY'], + secure=False, + ) + minio_mock().bucket_exists.assert_called_with( + bucket_name=self.env_var_map['MINIO_BUCKET_NAME'], + ) + minio_mock().make_bucket.assert_called_with( + bucket_name=self.env_var_map['MINIO_BUCKET_NAME'], + ) + + def test_close_should_overwrite_private_variables(self): + # ARRANGE + datastore = DatastoreMinio() + datastore._client = MagicMock() + datastore._bucket_name = MagicMock() + assert isinstance(datastore._client, MagicMock) + assert isinstance(datastore._bucket_name, MagicMock) + # ACT + datastore.close() + # ASSERT + assert datastore._client is None + assert datastore._bucket_name is None + + @patch('shared.datastore.src.datastore_minio.DatastoreMinio.connect') + @patch('shared.datastore.src.datastore_minio.DatastoreMinio.close') + def test_context_management_implemented( + self, + mocked_close_method, + mocked_connect_method, + ): + # ARRANGE, ACT and ASSERT + with DatastoreMinio() as ds: + mocked_connect_method.assert_called_once() + assert isinstance(ds, DatastoreMinio) + mocked_close_method.assert_called_once() + + def test_should_call_put_object_with_arguments(self): + # ARRANGE + datastore = DatastoreMinio() + datastore._client = MagicMock(Minio) + datastore._bucket_name = self.env_var_map['MINIO_BUCKET_NAME'] + self.buffer.seek(0) + # ACT + datastore._put( + object_name=self.object_name, + buffer=self.buffer, + ) + # ASSERT + datastore._client.put_object.assert_called_with( + bucket_name=self.env_var_map['MINIO_BUCKET_NAME'], + object_name=self.object_name, + length=self.num_bytes, + data=self.buffer, + ) + + def test_should_call_get_object_with_arguments(self): + # ARRANGE + datastore = DatastoreMinio() + datastore._client = MagicMock(Minio) + datastore._client.get_object.return_value = MagicMock( + BaseHTTPResponse, + status=200, + ) + # datastore._client.get_object.read.side_effect = [b'random ', b'test', b'text'] + datastore._bucket_name = self.env_var_map['MINIO_BUCKET_NAME'] + # ACT + with self.assertRaises(TypeError): # dont care to mock even more... + datastore._get( + object_name=self.object_name, + ) + # ASSERT + datastore._client.get_object.assert_called_with( + bucket_name=self.env_var_map['MINIO_BUCKET_NAME'], + object_name=self.object_name, + ) + + def test_should_call_remove_object_with_arguments(self): + # ARRANGE + datastore = DatastoreMinio() + datastore._client = MagicMock(Minio) + datastore._bucket_name = self.env_var_map['MINIO_BUCKET_NAME'] + # ACT + datastore._delete( + object_name=self.object_name, + ) + # ASSERT + datastore._client.remove_object.assert_called_with( + bucket_name=self.env_var_map['MINIO_BUCKET_NAME'], + object_name=self.object_name, + ) + + def test_put_image_should_call_put_object_with_arguments(self): + # ARRANGE + datastore = DatastoreMinio() + datastore._client = MagicMock(Minio) + datastore._bucket_name = self.env_var_map['MINIO_BUCKET_NAME'] + object_path = f'images/{self.checksum}' + # ACT + datastore.put_image( + image=self.image, + ) + # ASSERT + datastore._client.put_object.assert_called_with( + bucket_name=self.env_var_map['MINIO_BUCKET_NAME'], + object_name=object_path, + length=self.num_bytes, + data=ANY, # saved to different buffer when converting image + ) + + def test_get_image_should_call_get_object_with_arguments(self): + # ARRANGE + datastore = DatastoreMinio() + datastore._client = MagicMock(Minio) + datastore._client.get_object.return_value = MagicMock( + BaseHTTPResponse, + status=200, + ) + datastore._bucket_name = self.env_var_map['MINIO_BUCKET_NAME'] + object_path = f'images/{self.checksum}' + # ACT + with self.assertRaises(TypeError): # dont care to mock even more... + datastore.get_image( + object_name=self.checksum, + ) + # ASSERT + datastore._client.get_object.assert_called_with( + bucket_name=self.env_var_map['MINIO_BUCKET_NAME'], + object_name=object_path, + ) + + # @patch('shared.datastore.src.datastore_minio.torch.serialization.save') + # def test_put_model_should_call_put_object_with_arguments(self, mocked_torch_fn): + # # ARRANGE + # datastore = DatastoreMinio() + # datastore._client = MagicMock(Minio) + # datastore._bucket_name = self.env_var_map['MINIO_BUCKET_NAME'] + # object_path = f'images/{self.checksum}' + # mocked_torch_fn.return_value = None + # model = MagicMock(torch.nn.Module) + # # ACT + # datastore.put_model( + # model=model, + # ) + # # ASSERT + # datastore._client.put_object.assert_called_with( + # bucket_name=self.env_var_map['MINIO_BUCKET_NAME'], + # object_name=object_path, + # length=self.num_bytes, + # data=ANY, # saved to different buffer when converting data + # ) + + def test_get_model_should_call_get_object_with_arguments(self): + # ARRANGE + datastore = DatastoreMinio() + datastore._client = MagicMock(Minio) + datastore._client.get_object.return_value = MagicMock( + BaseHTTPResponse, + status=200, + ) + datastore._bucket_name = self.env_var_map['MINIO_BUCKET_NAME'] + object_path = f'models/{self.checksum}' + # ACT + with self.assertRaises(TypeError): # dont care to mock even more... + datastore.get_model( + object_name=self.checksum, + ) + # ASSERT + datastore._client.get_object.assert_called_with( + bucket_name=self.env_var_map['MINIO_BUCKET_NAME'], + object_name=object_path, + ) diff --git a/shared/datastore/tests/unit/delete_test.py b/shared/datastore/tests/unit/delete_test.py deleted file mode 100644 index 06c902d..0000000 --- a/shared/datastore/tests/unit/delete_test.py +++ /dev/null @@ -1,61 +0,0 @@ -"""Definition of unittests for delete function.""" - -import unittest -from unittest.mock import Mock - -from minio import Minio - -from shared.datastore import delete - - -class TestDelete(unittest.TestCase): - - def setUp(self): - # set relevant variables - self.client = Mock(spec=Minio) - self.bucket_name = 'test-bucket' - self.object_name = '46KXJMFIAPVLM0TKRFZR5YPPTVJ6PJNX' - # set bad arguments - self.bad_client = 'not-minio-type' - self.bad_string = float(0.0) - self.len_0_string = '' - - def test_should_fail_on_wrong_input_type_client(self): - with self.assertRaises(AssertionError): - delete( - client=self.bad_client, - bucket_name=self.bucket_name, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_bucket_name(self): - with self.assertRaises(AssertionError): - delete( - client=self.client, - bucket_name=self.bad_string, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_length_bucket_name(self): - with self.assertRaises(AssertionError): - delete( - client=self.client, - bucket_name=self.len_0_string, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_object_name(self): - with self.assertRaises(AssertionError): - delete( - client=self.client, - bucket_name=self.bucket_name, - object_name=self.bad_string, - ) - - def test_should_fail_on_wrong_input_length_object_name(self): - with self.assertRaises(AssertionError): - delete( - client=self.client, - bucket_name=self.bucket_name, - object_name=self.len_0_string, - ) diff --git a/shared/datastore/tests/unit/get_image_test.py b/shared/datastore/tests/unit/get_image_test.py deleted file mode 100644 index ab72324..0000000 --- a/shared/datastore/tests/unit/get_image_test.py +++ /dev/null @@ -1,63 +0,0 @@ -"""Definition of unittests for get_image function.""" - -import os -import unittest -from unittest.mock import MagicMock - -from minio import Minio - -from shared.datastore import get_image - - -class TestGetImage(unittest.TestCase): - - def setUp(self): - # set relevant variables - self.client = MagicMock(spec=Minio) - self.object_name = '46KXJMFIAPVLM0TKRFZR5YPPTVJ6PJNX' - self.env_var_map = { - 'MINIO_BUCKET_NAME': 'test-bucket', - } - # set bad arguments - self.bad_client = 'not-minio-type' - self.bad_string = float(0.0) - self.len_0_string = '' - # populate env - for key, val in self.env_var_map.items(): - os.environ[key] = val - - def tearDown(self): - # clean env - for key in self.env_var_map: - _ = os.environ.pop(key, default=None) - - def test_should_fail_when_env_not_set(self): - # ensure env not set - self.tearDown() - # run test - with self.assertRaises(AssertionError): - get_image( - client=self.client, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_client(self): - with self.assertRaises(AssertionError): - get_image( - client=self.bad_client, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_object_name(self): - with self.assertRaises(AssertionError): - get_image( - client=self.client, - object_name=self.bad_string, - ) - - def test_should_fail_on_wrong_input_length_object_name(self): - with self.assertRaises(AssertionError): - get_image( - client=self.client, - object_name=self.len_0_string, - ) diff --git a/shared/datastore/tests/unit/get_model_test.py b/shared/datastore/tests/unit/get_model_test.py deleted file mode 100644 index 0ba1f76..0000000 --- a/shared/datastore/tests/unit/get_model_test.py +++ /dev/null @@ -1,63 +0,0 @@ -"""Definition of unittest for get_model function.""" - -import os -import unittest -from unittest.mock import MagicMock - -from minio import Minio - -from shared.datastore import get_model - - -class TestGetModel(unittest.TestCase): - - def setUp(self): - # set relevant variables - self.client = MagicMock(spec=Minio) - self.object_name = '46KXJMFIAPVLM0TKRFZR5YPPTVJ6PJNX' - self.env_var_map = { - 'MINIO_BUCKET_NAME': 'test-bucket', - } - # set bad arguments - self.bad_client = 'not-minio-type' - self.bad_string = float(0.0) - self.len_0_string = '' - # populate env - for key, val in self.env_var_map.items(): - os.environ[key] = val - - def tearDown(self): - # clean env - for key in self.env_var_map: - _ = os.environ.pop(key, default=None) - - def test_should_fail_when_env_not_set(self): - # ensure env not set - self.tearDown() - # run test - with self.assertRaises(AssertionError): - get_model( - client=self.client, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_client(self): - with self.assertRaises(AssertionError): - get_model( - client=self.bad_client, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_object_name(self): - with self.assertRaises(AssertionError): - get_model( - client=self.client, - object_name=self.bad_string, - ) - - def test_should_fail_on_wrong_input_length_object_name(self): - with self.assertRaises(AssertionError): - get_model( - client=self.client, - object_name=self.len_0_string, - ) diff --git a/shared/datastore/tests/unit/get_test.py b/shared/datastore/tests/unit/get_test.py deleted file mode 100644 index 056f2da..0000000 --- a/shared/datastore/tests/unit/get_test.py +++ /dev/null @@ -1,60 +0,0 @@ -"""Definition of unittests for get function.""" - -import unittest -from unittest.mock import MagicMock - -from minio import Minio - -from shared.datastore import get - - -class TestGet(unittest.TestCase): - def setUp(self): - # set relevant variables - self.client = MagicMock(spec=Minio) - self.bucket_name = 'test-bucket' - self.object_name = '46KXJMFIAPVLM0TKRFZR5YPPTVJ6PJNX' - # set bad arguments - self.bad_client = 'not-minio-type' - self.bad_string = float(0.0) - self.len_0_string = '' - - def test_should_fail_on_wrong_input_type_client(self): - with self.assertRaises(AssertionError): - get( - client=self.bad_client, - bucket_name=self.bucket_name, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_bucket_name(self): - with self.assertRaises(AssertionError): - get( - client=self.client, - bucket_name=self.bad_string, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_length_bucket_name(self): - with self.assertRaises(AssertionError): - get( - client=self.client, - bucket_name=self.len_0_string, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_object_name(self): - with self.assertRaises(AssertionError): - get( - client=self.client, - bucket_name=self.bucket_name, - object_name=self.bad_string, - ) - - def test_should_fail_on_wrong_input_length_object_name(self): - with self.assertRaises(AssertionError): - get( - client=self.client, - bucket_name=self.bucket_name, - object_name=self.len_0_string, - ) diff --git a/shared/datastore/tests/unit/put_image_test.py b/shared/datastore/tests/unit/put_image_test.py deleted file mode 100644 index e971478..0000000 --- a/shared/datastore/tests/unit/put_image_test.py +++ /dev/null @@ -1,41 +0,0 @@ -"""Definition of unittests for put_image function.""" - -import os -import unittest -from unittest.mock import MagicMock - -from minio import Minio -from PIL import Image - -from shared.datastore import put_image - - -class TestPutImage(unittest.TestCase): - - def setUp(self): - # set relevant variables - self.client = MagicMock(spec=Minio) - self.image = Image.new(mode='RGB', size=(480, 480)) - self.env_var_map = { - 'MINIO_BUCKET_NAME': 'test-bucket', - } - # set bad arguments - self.bad_client = 'not-minio-type' - self.bad_image = 'not-image-type' - # populate env - for key, val in self.env_var_map.items(): - os.environ[key] = val - - def test_should_fail_on_wrong_input_type_client(self): - with self.assertRaises(AssertionError): - put_image( - client=self.bad_client, - image=self.image, - ) - - def test_should_fail_on_wrong_input_type_image(self): - with self.assertRaises(AssertionError): - put_image( - client=self.client, - image=self.bad_image, - ) diff --git a/shared/datastore/tests/unit/put_model_test.py b/shared/datastore/tests/unit/put_model_test.py deleted file mode 100644 index 790e839..0000000 --- a/shared/datastore/tests/unit/put_model_test.py +++ /dev/null @@ -1,40 +0,0 @@ -"""Definition of unittests for put_model function.""" - -import os -import unittest -from unittest.mock import MagicMock - -from minio import Minio -from torch.nn import Module - -from shared.datastore import put_model - - -class TestPutModel(unittest.TestCase): - def setUp(self): - # set relevant variables - self.client = MagicMock(spec=Minio) - self.model = MagicMock(spec=Module) - self.env_var_map = { - 'MINIO_BUCKET_NAME': 'test-bucket', - } - # set bad arguments - self.bad_client = 'not-minio-type' - self.bad_model = 'not-image-type' - # populate env - for key, val in self.env_var_map.items(): - os.environ[key] = val - - def test_should_fail_on_wrong_input_type_client(self): - with self.assertRaises(AssertionError): - put_model( - client=self.bad_client, - model=self.model, - ) - - def test_should_fail_on_wrong_input_type_model(self): - with self.assertRaises(AssertionError): - put_model( - client=self.client, - model=self.bad_model, - ) diff --git a/shared/datastore/tests/unit/put_test.py b/shared/datastore/tests/unit/put_test.py deleted file mode 100644 index b0c8501..0000000 --- a/shared/datastore/tests/unit/put_test.py +++ /dev/null @@ -1,80 +0,0 @@ -"""Definition of unittests for put function.""" - -import unittest -from io import BytesIO -from unittest.mock import MagicMock - -from minio import Minio -from PIL import Image - -from shared.datastore import put - - -class TestPut(unittest.TestCase): - def setUp(self): - # set relevant variables - self.client = MagicMock(spec=Minio) - self.bucket_name = 'test-bucket' - self.object_name = '46KXJMFIAPVLM0TKRFZR5YPPTVJ6PJNX' - self.image = Image.new(mode='RGB', size=(480, 480)) - self.buffer = BytesIO() - self.image.save(self.buffer, 'png') - # set bad arguments - self.bad_client = 'not-minio-type' - self.bad_string = float(0.0) - self.bad_buffer = 'not-buffer-type' - self.len_0_string = '' - - def test_should_fail_on_wrong_input_type_client(self): - with self.assertRaises(AssertionError): - put( - client=self.bad_client, - buffer=self.buffer, - bucket_name=self.bucket_name, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_buffer(self): - with self.assertRaises(AssertionError): - put( - client=self.client, - buffer=self.bad_buffer, - bucket_name=self.bucket_name, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_bucket_name(self): - with self.assertRaises(AssertionError): - put( - client=self.client, - buffer=self.buffer, - bucket_name=self.bad_string, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_length_bucket_name(self): - with self.assertRaises(AssertionError): - put( - client=self.client, - buffer=self.buffer, - bucket_name=self.len_0_string, - object_name=self.object_name, - ) - - def test_should_fail_on_wrong_input_type_object_name(self): - with self.assertRaises(AssertionError): - put( - client=self.client, - buffer=self.buffer, - bucket_name=self.bucket_name, - object_name=self.bad_string, - ) - - def test_should_fail_on_wrong_input_length_object_name(self): - with self.assertRaises(AssertionError): - put( - client=self.client, - buffer=self.buffer, - bucket_name=self.bucket_name, - object_name=self.len_0_string, - ) From aff0ae8fc66a8826ac1408d018b0e50045651ffe Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 6 Jan 2025 13:42:21 +0000 Subject: [PATCH 10/10] fixed types --- misc/generate_random_prediction.py | 10 ++++++---- misc/get_visual_communication.py | 8 +++++--- misc/image_upload.py | 10 ++++++---- misc/image_upload_to_server.py | 10 ++++++---- misc/model_io.py | 8 ++++---- misc/new_model_to_minio.py | 8 ++++---- misc/prediction_upload.py | 10 ++++++---- model/src/main.py | 9 +++++---- model/src/utils/load_model.py | 10 ++++------ model/src/utils/vcda_dataset.py | 10 ++++------ model/train.py | 9 +++++---- other/transfer_images_minio_subfolder.py | 21 ++++++++++++--------- other/transfer_images_mongo_minio.py | 11 ++++++----- shared/utils/tests/unit/check_env_test.py | 2 +- web_ui/src/app/init_app.py | 19 ++++++++++--------- web_ui/src/main.py | 7 ++++--- 16 files changed, 88 insertions(+), 74 deletions(-) diff --git a/misc/generate_random_prediction.py b/misc/generate_random_prediction.py index aea921a..3670fe9 100644 --- a/misc/generate_random_prediction.py +++ b/misc/generate_random_prediction.py @@ -4,9 +4,9 @@ from pathlib import Path from dotenv import load_dotenv -from shared.datastore import connect_minio +from shared.datastore import Datastore from shared.mongodb import connect_mongodb -from shared.mongodb.classes import VisualCommunication +from shared.mongodb.src.classes import VisualCommunication from shared.utils import check_env, setup_logging from web_ui.src.main import NECESSARY_ENV_VAR_LIST @@ -20,7 +20,9 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minIO - minio_client = connect_minio() + datastore = Datastore() + datastore.connect() + assert datastore._client is not None # connect to MongoDB collection, db, client = connect_mongodb() # get list of image paths @@ -30,7 +32,7 @@ if __name__ == '__main__': print(img_path_list) # instantiate data object vis_com_list = [ - VisualCommunication.from_file(path, minio_client=minio_client) + VisualCommunication.from_file(path, minio_client=datastore._client) for path in img_path_list ] # generate random predictions diff --git a/misc/get_visual_communication.py b/misc/get_visual_communication.py index b8885c5..9643a6c 100644 --- a/misc/get_visual_communication.py +++ b/misc/get_visual_communication.py @@ -4,7 +4,7 @@ from pathlib import Path from dotenv import load_dotenv -from shared.datastore import connect_minio +from shared.datastore import Datastore from shared.mongodb import connect_mongodb, get_visual_communication from shared.utils import check_env, setup_logging from web_ui.src.main import NECESSARY_ENV_VAR_LIST @@ -19,11 +19,13 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minIO - minio_client = connect_minio() + datastore = Datastore() + datastore.connect() + assert datastore._client is not None # connect to MongoDB collection, db, client = connect_mongodb() # get visual communication vis_com = get_visual_communication(collection) print(repr(vis_com)) - image = vis_com.get_image(minio_client=minio_client) + image = vis_com.get_image(minio_client=datastore._client) image.show() diff --git a/misc/image_upload.py b/misc/image_upload.py index 2a85e40..9db6652 100644 --- a/misc/image_upload.py +++ b/misc/image_upload.py @@ -5,9 +5,9 @@ from pathlib import Path from dotenv import load_dotenv from pymongo.errors import DuplicateKeyError -from shared.datastore import connect_minio +from shared.datastore import Datastore from shared.mongodb import connect_mongodb -from shared.mongodb.classes import VisualCommunication +from shared.mongodb.src.classes import VisualCommunication from shared.utils import check_env, setup_logging from web_ui.src.main import NECESSARY_ENV_VAR_LIST @@ -21,7 +21,9 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minIO - minio_client = connect_minio() + datastore = Datastore() + datastore.connect() + assert datastore._client is not None # connect to MongoDB collection, db, client = connect_mongodb() # get list of image paths @@ -31,7 +33,7 @@ if __name__ == '__main__': print(img_path_list) # instantiate data object vis_com_list = [ - VisualCommunication.from_file(path, minio_client=minio_client) + VisualCommunication.from_file(path, minio_client=datastore._client) for path in img_path_list ] for vis_com in vis_com_list: diff --git a/misc/image_upload_to_server.py b/misc/image_upload_to_server.py index 4c8f02e..2333c27 100644 --- a/misc/image_upload_to_server.py +++ b/misc/image_upload_to_server.py @@ -5,9 +5,9 @@ from pathlib import Path from dotenv import load_dotenv from pymongo.errors import DuplicateKeyError -from shared.datastore import connect_minio +from shared.datastore import Datastore from shared.mongodb import connect_mongodb -from shared.mongodb.classes import VisualCommunication +from shared.mongodb.src.classes import VisualCommunication from shared.utils import check_env, setup_logging from web_ui.src.main import NECESSARY_ENV_VAR_LIST @@ -21,7 +21,9 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minIO - minio_client = connect_minio() + datastore = Datastore() + datastore.connect() + assert datastore._client is not None # connect to MongoDB collection, db, client = connect_mongodb() # get list of image paths @@ -37,7 +39,7 @@ if __name__ == '__main__': print(f"found {len(img_path_list)} images") # create visual communication objects vis_com_list = [ - VisualCommunication.from_file(path, minio_client=minio_client) + VisualCommunication.from_file(path, minio_client=datastore._client) for path in img_path_list ] print(f"created {len(vis_com_list)} visual communication objects") diff --git a/misc/model_io.py b/misc/model_io.py index d6d7635..b5bec6d 100644 --- a/misc/model_io.py +++ b/misc/model_io.py @@ -7,7 +7,7 @@ from dotenv import load_dotenv from torchinfo import summary from model.src.models import VisualCommunicationModel -from shared.datastore import connect_minio, put_model +from shared.datastore import Datastore from shared.utils import setup_logging DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') @@ -20,14 +20,14 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minio - client = connect_minio() + datastore = Datastore() + datastore.connect() # instantiate model model = VisualCommunicationModel().to(DEVICE) # show model weights summary(model) # put buffer in minio bucket - hash_str = put_model( - client=client, + hash_str = datastore.put_model( model=model, ) print(f"hash string: {hash_str}") diff --git a/misc/new_model_to_minio.py b/misc/new_model_to_minio.py index 2f516a0..416dd25 100644 --- a/misc/new_model_to_minio.py +++ b/misc/new_model_to_minio.py @@ -7,7 +7,7 @@ from dotenv import load_dotenv from torchinfo import summary from model.src.models import VisualCommunicationModel -from shared.datastore import connect_minio, put_model +from shared.datastore import Datastore from shared.utils import setup_logging if __name__ == '__main__': @@ -18,14 +18,14 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minio - client = connect_minio() + datastore = Datastore() + datastore.connect() # instantiate model model = VisualCommunicationModel(download_resnet_weights=True) # show model weights summary(model) # put buffer in minio bucket - hash_str = put_model( - client=client, + hash_str = datastore.put_model( model=model, ) print(f"hash string: {hash_str}") diff --git a/misc/prediction_upload.py b/misc/prediction_upload.py index 6c27211..2506fe7 100644 --- a/misc/prediction_upload.py +++ b/misc/prediction_upload.py @@ -4,9 +4,9 @@ from pathlib import Path from dotenv import load_dotenv -from shared.datastore import connect_minio +from shared.datastore import Datastore from shared.mongodb import connect_mongodb, upsert_prediction -from shared.mongodb.classes import VisualCommunication +from shared.mongodb.src.classes import VisualCommunication from shared.utils import check_env, setup_logging from web_ui.src.main import NECESSARY_ENV_VAR_LIST @@ -20,7 +20,9 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minIO - minio_client = connect_minio() + datastore = Datastore() + datastore.connect() + assert datastore._client is not None # connect to MongoDB collection, db, client = connect_mongodb() # get list of image paths @@ -29,7 +31,7 @@ if __name__ == '__main__': img_path_list = [path for path in img_dir.glob('*.jpeg') if path.is_file()] # instantiate data object vis_com_list = [ - VisualCommunication.from_file(path, minio_client=minio_client) + VisualCommunication.from_file(path, minio_client=datastore._client) for path in img_path_list ] # generate random predictions diff --git a/model/src/main.py b/model/src/main.py index 45cd7eb..91d22c2 100644 --- a/model/src/main.py +++ b/model/src/main.py @@ -9,7 +9,7 @@ from models import VisualCommunicationModel from tqdm import tqdm from utils import DEVICE, VCDADataset, load_model -from shared.datastore import connect_minio +from shared.datastore import Datastore from shared.mongodb.classes import ModelData from shared.utils import setup_logging @@ -17,13 +17,14 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minio - minio_client = connect_minio() + datastore = Datastore() + datastore.connect() # instantiate model - model: VisualCommunicationModel = load_model(client=minio_client) + model: VisualCommunicationModel = load_model(client=datastore._client) model.eval() # setup dataset dataset = VCDADataset( - minio_client=minio_client, + minio_client=datastore._client, data_name_list=[ '02dbaf48d713e4e6d3a6b98fd2dc866e', ], diff --git a/model/src/utils/load_model.py b/model/src/utils/load_model.py index 311539f..992f498 100644 --- a/model/src/utils/load_model.py +++ b/model/src/utils/load_model.py @@ -4,10 +4,9 @@ import logging from pathlib import Path import torch -from minio import Minio from model.src.models import VisualCommunicationModel -from shared.datastore import get_model +from shared.datastore import Datastore from .get_model_name import get_model_name @@ -15,11 +14,11 @@ DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') def load_model( - client: Minio, + datastore: Datastore, ) -> VisualCommunicationModel: """Instantiate model with weights loaded from latest model saved in MinIO.""" - assert isinstance(client, Minio) + assert isinstance(datastore, Datastore) # instantiate model model = VisualCommunicationModel() # get model object name @@ -27,8 +26,7 @@ def load_model( model_object_name = get_model_name(path=model_name_path) logging.info('using model: %s', model_object_name) # load model from minio - model_checkpoint = get_model( - client=client, + model_checkpoint = datastore.get_model( object_name=model_object_name, ) model.load_state_dict(model_checkpoint) diff --git a/model/src/utils/vcda_dataset.py b/model/src/utils/vcda_dataset.py index 8f97a4b..43f57e7 100644 --- a/model/src/utils/vcda_dataset.py +++ b/model/src/utils/vcda_dataset.py @@ -2,7 +2,6 @@ import random -from minio import Minio from PIL import Image from torch import Tensor from torch.utils.data import Dataset @@ -16,7 +15,7 @@ from torchvision.transforms.functional import ( to_tensor, ) -from shared.datastore import get_image +from shared.datastore import Datastore # resnet18 original normalization values RESNET_NORMALIZE_MEAN = [0.485, 0.456, 0.406] @@ -28,13 +27,13 @@ class VCDADataset(Dataset): def __init__( self, - minio_client: Minio, + datastore: Datastore, data_name_list: list[str], do_augment: bool = False, random_annotations: bool = False, ): super().__init__() - self.minio_client = minio_client + self.datastore = datastore self.data_name_list = data_name_list self.do_augment = do_augment self.random_annotations = random_annotations @@ -57,8 +56,7 @@ class VCDADataset(Dataset): def __getitem__(self, idx): # get image from database object_name = self.data_name_list[idx] - image = get_image( - client=self.minio_client, + image = self.datastore.get_image( object_name=object_name, ) tensor = self.image_to_tensor(image) diff --git a/model/train.py b/model/train.py index d75904e..0446285 100644 --- a/model/train.py +++ b/model/train.py @@ -18,7 +18,7 @@ from torch.optim.lr_scheduler import ExponentialLR from torch.utils.data import DataLoader from model.src.utils import VCDADataset, get_class -from shared.datastore import connect_minio +from shared.datastore import Datastore def parse_arguments(): @@ -57,18 +57,19 @@ optimizer = torch.optim.Adam(model.parameters(), lr=args.lr) loss_fn = nn.CrossEntropyLoss() # create datasets and loaders -minio_client = connect_minio() +datastore = Datastore() +datastore.connect() with open('model/src/dataset/train.csv', encoding='utf-8') as fh: train_data_name_list = fh.read().split('\n') train_dataset = VCDADataset( - minio_client=minio_client, + datastore=datastore, data_name_list=train_data_name_list, ) train_loader = DataLoader(dataset=train_dataset, num_workers=args.loader_workers) with open('model/src/dataset/val.csv', encoding='utf-8') as fh: val_data_name_list = fh.read().split('\n') -val_dataset = VCDADataset(minio_client=minio_client, data_name_list=val_data_name_list) +val_dataset = VCDADataset(datastore=datastore, data_name_list=val_data_name_list) val_loader = DataLoader(dataset=val_dataset, num_workers=args.loader_workers) # create trainer and evaluator diff --git a/other/transfer_images_minio_subfolder.py b/other/transfer_images_minio_subfolder.py index cb5305e..4c960a8 100644 --- a/other/transfer_images_minio_subfolder.py +++ b/other/transfer_images_minio_subfolder.py @@ -1,11 +1,12 @@ """Script to move minio images to subfolder.""" +import os from pathlib import Path from dotenv import load_dotenv from PIL import Image -from shared.datastore import connect_minio, get, put_image +from shared.datastore import Datastore from shared.utils import setup_logging if __name__ == '__main__': @@ -16,24 +17,26 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minio - minio_client = connect_minio() + datastore = Datastore() + datastore.connect() + assert datastore._client is not None # list images in bucket - BUCKET_NAME = 'visual-critical-discourse-analysis' - obj_list = minio_client.list_objects( + BUCKET_NAME = os.getenv( + 'MINIO_BUCKET_NAME', + default='visual-critical-discourse-analysis', + ) + obj_list = datastore._client.list_objects( bucket_name=BUCKET_NAME, ) # begin moving images for obj in obj_list: # get image from minio - buffer = get( - client=minio_client, - bucket_name=BUCKET_NAME, + buffer = datastore._get( object_name=obj.object_name, ) # convert data to image image = Image.open(buffer) # put image into minio subfolder - put_image( - client=minio_client, + _ = datastore.put_image( image=image, ) diff --git a/other/transfer_images_mongo_minio.py b/other/transfer_images_mongo_minio.py index 7e8f5cc..8b8033c 100644 --- a/other/transfer_images_mongo_minio.py +++ b/other/transfer_images_mongo_minio.py @@ -9,7 +9,7 @@ from bson import ObjectId from dotenv import load_dotenv from pymongo.collection import Collection -from shared.datastore import connect_minio, put_image +from shared.datastore import Datastore from shared.mongodb import connect_mongodb from shared.mongodb.src.classes import VisualCommunication from shared.utils import check_env, setup_logging @@ -81,7 +81,9 @@ if __name__ == '__main__': # setup logging setup_logging() # connect to minIO - minio_client = connect_minio() + datastore = Datastore() + datastore.connect() + assert datastore._client is not None # connect to MongoDB collection, db, client = connect_mongodb() # list documents in mongoDB @@ -97,11 +99,10 @@ if __name__ == '__main__': try: # get image image = vis_com.get_image( - minio_client=minio_client, + minio_client=datastore._client, ) # put buffer in minio - object_name = put_image( - client=minio_client, + object_name = datastore.put_image( image=image, ) except Exception as exc: diff --git a/shared/utils/tests/unit/check_env_test.py b/shared/utils/tests/unit/check_env_test.py index 32ec322..7bc4bc3 100644 --- a/shared/utils/tests/unit/check_env_test.py +++ b/shared/utils/tests/unit/check_env_test.py @@ -39,7 +39,7 @@ class TestFunctionCheckEnv(unittest.TestCase): variable that is not set.""" var_list = {self.not_set_env_var} msg = f'environment variable not set: {self.not_set_env_var}' - with self.assertRaises(AssertionError, msg=msg): + with self.assertRaises(OSError, msg=msg): check_env(var_list) def test_env_vars_set(self): diff --git a/web_ui/src/app/init_app.py b/web_ui/src/app/init_app.py index 897aa69..e479e74 100644 --- a/web_ui/src/app/init_app.py +++ b/web_ui/src/app/init_app.py @@ -8,11 +8,10 @@ import os import dash_bootstrap_components as dbc from dash import ALL, Dash, Input, Output, State from dash_auth import BasicAuth -from minio import Minio from pydantic import ValidationError from pymongo.collection import Collection -from shared.datastore import delete as delete_from_minio +from shared.datastore import Datastore from shared.mongodb import count_documents, get_visual_communication, upsert_annotation from shared.mongodb.src.classes import ModelData, VisualCommunication from shared.mongodb.src.exceptions import NoDocumentFoundException @@ -22,9 +21,12 @@ from .layout import app_layout def init_app( mongo_collection: Collection, - minio_client: Minio, + datastore: Datastore, ) -> Dash: """Initialise web UI application.""" + assert isinstance(mongo_collection, Collection) + assert isinstance(datastore, Datastore) + assert datastore._client is not None # setup app app = Dash( name='visual_critical_discourse_analysis_web_ui', @@ -129,11 +131,12 @@ def init_app( failed_filename_list.append(filename) continue try: + assert datastore._client is not None # instantiate to upload image to minio vis_com = VisualCommunication.from_name_and_image( name=filename, image=image, - minio_client=minio_client, + minio_client=datastore._client, ) except Exception as exc: logging.debug(exc) @@ -146,10 +149,7 @@ def init_app( logging.debug(exc) failed_filename_list.append(filename) # remove document from minio - bucket_name = os.getenv('MINIO_BUCKET_NAME', default='') - delete_from_minio( - client=minio_client, - bucket_name=bucket_name, + datastore._delete( object_name=vis_com.object_name, ) assert ( @@ -240,7 +240,8 @@ def init_app( ) # set variables vis_com_name = vis_com.name - image_src = vis_com.webencoded_image(minio_client=minio_client) + assert datastore._client is not None + image_src = vis_com.webencoded_image(minio_client=datastore._client) if vis_com.prediction is not None: # TODO: update to use optional predictions pass diff --git a/web_ui/src/main.py b/web_ui/src/main.py index d5a6de7..5546b47 100644 --- a/web_ui/src/main.py +++ b/web_ui/src/main.py @@ -4,7 +4,7 @@ from __future__ import annotations import os -from shared.datastore import connect_minio +from shared.datastore import Datastore from shared.mongodb import connect_mongodb from shared.utils import check_env, setup_logging @@ -34,12 +34,13 @@ setup_logging() collection, db, client = connect_mongodb() # connect to minio -minio_client = connect_minio() +datastore = Datastore() +datastore.connect() # initialise application app = init_app( mongo_collection=collection, - minio_client=minio_client, + datastore=datastore, ) server = app.server server.config.update(SECRET_KEY=os.urandom(24))