Files
visual_critical_discourse_a…/shared/datastore/tests/integration/conftest.py
T
brian 16ad2b80ee
Code Quality Pipeline / Check Code (pull_request) Failing after 3m5s
updated tests to match new interface
2025-01-06 11:34:02 +00:00

197 lines
5.2 KiB
Python

"""Integration test configurations."""
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 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',
}
@pytest.fixture(scope='session', autouse=True)
def populate_env(
request: pytest.FixtureRequest,
) -> None:
"""Populate environment with variables used for testing."""
# read env-file for local testing
load_dotenv(
dotenv_path=Path(__file__).parent.parent.parent.parent.parent / 'server.env',
)
# update env
for key, val in env_var_map.items():
os.environ[key] = val
# ensure cleanup
def cleanup_env():
for key in env_var_map:
_ = os.environ.pop(key, default=None)
request.addfinalizer(cleanup_env)
@pytest.fixture(scope='session')
def datastore(
populate_env,
) -> Iterator[Datastore]:
# prepare arguments
minio_bucket_name = str(os.getenv('MINIO_BUCKET_NAME'))
# connect to minio
datastore_client = Datastore()
datastore_client.connect()
assert datastore_client._client is not None
# ensure bucket exists
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 datastore_client
# remove objects left behind by tests
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
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
def data() -> Iterator[bytes]:
# generate random data
num_bytes = 2**21 # 2 MB
data = random.randbytes(n=num_bytes)
# expose data
yield data
@pytest.fixture
def data_in_minio(
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'))
# convert data
buffer = BytesIO(data)
# 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=minio_object_name,
length=num_bytes,
data=buffer,
)
# expose data
yield buffer, minio_bucket_name, minio_object_name
# clean up
datastore._client.remove_object(
bucket_name=minio_bucket_name,
object_name=minio_object_name,
)
@pytest.fixture
def image() -> Iterator[Image.Image]:
# generate image
image = Image.new(mode='RGB', size=(480, 480))
# expose image
yield image
@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,
)
@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,
)