added integration tests and repository pattern for image, model and visual communication DTOs
Code Quality Pipeline / Check Code (pull_request) Failing after 3m21s

This commit is contained in:
brian
2025-02-28 22:24:52 +00:00
parent ac6a305da8
commit 090a1cd8bb
45 changed files with 2630 additions and 0 deletions
+83
View File
@@ -0,0 +1,83 @@
"""Definition of docstore interface."""
from __future__ import annotations
from abc import ABC, abstractmethod
from shared.docstore.src.classes import ModelData, VisualCommunication
class DocstoreInterface(ABC):
"""Docstore interface class."""
@abstractmethod
def connect(
self,
) -> None:
pass
@abstractmethod
def close(
self,
) -> None:
pass
@abstractmethod
def __enter__(
self,
) -> DocstoreInterface:
pass
@abstractmethod
def __exit__(
self,
exc_type,
exc_val,
exc_tb,
) -> None:
pass
@abstractmethod
def count_documents(
self,
only_with_annotation: bool,
) -> int:
pass
@abstractmethod
def list_names(
self,
only_with_annotation: bool,
) -> list[str]:
pass
@abstractmethod
def upsert_visual_comminucations(
self,
visual_communication_list: list[VisualCommunication],
) -> None:
pass
@abstractmethod
def upsert_annotations(
self,
visual_communication_name: str,
annotations: ModelData,
) -> None:
pass
@abstractmethod
def upsert_prediction(
self,
visual_communication_name: str,
predictions: ModelData,
) -> None:
pass
@abstractmethod
def get_visual_communication(
self,
with_annotation: bool,
name: str | None = None,
) -> VisualCommunication:
pass
+241
View File
@@ -0,0 +1,241 @@
"""Definition of docstore mongodb implementation."""
from __future__ import annotations
import logging
import os
from typing import Any
from pymongo import MongoClient
from pymongo.collection import Collection
from pymongo.database import Database
from shared.utils import check_env
from .classes import ModelData, VisualCommunication
from .docstore_interface import DocstoreInterface
from .exceptions import NoDocumentFoundException
class DocstoreMongo(DocstoreInterface):
"""Docstore interface."""
def __init__(self):
# ensure necessary env vars available
var_list = {
'MONGO_ENDPOINT',
'MONGO_DB',
'MONGO_COLLECTION',
}
check_env(var_list)
# prepare internal variables
self.client: MongoClient | None = None
self.db: Database | None = None
self.collection: Collection | None = None
def connect(self) -> None:
"""Connect to Mongo server."""
# prepare arguments
mongo_endpoint = str(os.getenv('MONGO_ENDPOINT'))
mongo_database = str(os.getenv('MONGO_DB'))
mongo_collection = str(os.getenv('MONGO_COLLECTION'))
# connect client
client = MongoClient(mongo_endpoint)
database = client[mongo_database]
collection = database[mongo_collection]
# set unique index on 'name'
collection.create_index(keys='name', unique=True)
# persist state
self._client = client
self._database = database
self._collection = collection
def close(self) -> None:
"""Close connection to Mongo server."""
self._client.close()
self._client = None
self._database = None
self._collection = None
def __enter__(self) -> DocstoreMongo:
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 count_documents(
self,
only_with_annotation: bool = False,
):
"""Get the total number of Visual Communication documents matching the
filter."""
assert isinstance(only_with_annotation, bool)
# build query
query = {}
if only_with_annotation:
query['annotation'] = {'$ne': None}
# execute query
num_docs = self._collection.count_documents(filter=query)
return num_docs
def list_names(
self,
only_with_annotation: bool,
) -> list[str]:
"""List names of Visual Communication documents that match the
filters."""
assert isinstance(only_with_annotation, bool)
assert isinstance(self._collection, Collection)
# build query
query = {}
if only_with_annotation:
query['annotation'] = {'$ne': None}
# execute query
res_list = self._collection.find(
filter=query,
projection={
'_id': False, # don't get id
'name': True, # include document name
},
)
# extract info
name_list = [elem['name'] for elem in res_list]
return name_list
def upsert_visual_comminucations(
self,
visual_communication_list: list[VisualCommunication],
) -> None:
"""Upsert Visual Communication document in the database."""
assert isinstance(visual_communication_list, list)
assert all(
isinstance(vis_com, VisualCommunication)
for vis_com in visual_communication_list
)
assert isinstance(self._collection, Collection)
# convert to dict
doc_list = [vis_com.model_dump() for vis_com in visual_communication_list]
# insert documents
response = self._collection.insert_many(documents=doc_list)
# check response
if not response.acknowledged:
raise OSError('failed inserting documents')
def upsert_annotations(
self,
visual_communication_name: str,
annotations: ModelData,
) -> None:
"""Upsert annotation for the document with matching name."""
assert isinstance(visual_communication_name, str)
assert len(visual_communication_name) > 0
assert isinstance(annotations, ModelData)
assert isinstance(self._collection, Collection)
# convert to dict
doc = annotations.model_dump()
# build query
query = {
'name': visual_communication_name,
}
update = {
'$set': {
'annotation': doc,
},
}
# execute query
res = self._collection.update_one(
filter=query,
update=update,
upsert=True,
)
# check response
if not res.acknowledged:
raise OSError(
f'failed upserting annotations for {visual_communication_name}',
)
def upsert_prediction(
self,
visual_communication_name: str,
predictions: ModelData,
) -> None:
"""Upsert prediction for the document with matching name."""
assert isinstance(visual_communication_name, str)
assert len(visual_communication_name) > 0
assert isinstance(annotations, ModelData)
assert isinstance(self._collection, Collection)
# convert to dict
doc = predictions.model_dump()
# build query
query = {
'name': visual_communication_name,
}
update = {
'$set': {
'prediction': doc,
},
}
# execute query
res = self._collection.update_one(
filter=query,
update=update,
upsert=True,
)
# check response
if not res.acknowledged:
raise OSError(
f'failed upserting predictions for {visual_communication_name}',
)
def get_visual_communication(
self,
with_annotation: bool = False,
name: str | None = None,
) -> VisualCommunication:
"""Get a random Visual Communication document that matches annotation
filter.
If name is not specified, a random document matching filter is
returned.
"""
assert isinstance(with_annotation, bool)
if name is not None:
assert isinstance(name, str)
assert len(name) > 0
# build query
query: dict[str, Any] = {}
if with_annotation:
query['annotation'] = {'$ne': None}
else:
query['annotation'] = {'$eq': None}
if name is not None:
query['name'] = {'$eq': name}
# execute query
res_list = self._collection.aggregate(
pipeline=[
{
'$match': query, # find using filters
},
{
'$sample': {
'size': 1, # get one random
},
},
],
)
doc_list = list(res_list)
# check result
if len(doc_list) == 0:
raise NoDocumentFoundException()
# convert
vis_com = VisualCommunication.model_validate(doc_list[0])
return vis_com
@@ -0,0 +1,45 @@
"""Definition of unittests for Docstore MongoDB implementation."""
import os
from unittest import TestCase
from unittest.mock import MagicMock
from pymongo import MongoClient
from pymongo.collection import Collection
from pymongo.database import Database
from shared.docstore.src.docstore_mongo import DocstoreMongo
class TestDocstoreMongoImplementation(TestCase):
def setUp(self):
# define relevant env vars
self.env_var_map = {
'MONGO_ENDPOINT': '192.168.1.2:27017',
'MONGO_DB': 'visual_critical_discourse_analysis',
'MONGO_COLLECTION': 'test-collection',
}
# set env vars
for key, val in self.env_var_map.items():
os.environ[key] = val
# set other variables
self.mongo_client_mock = MagicMock(MongoClient)
self.mongo_database_mock = MagicMock(Database)
self.mongo_collection_mock = MagicMock(Collection)
def tearDown(self):
# clear env vars
for key in self.env_var_map:
_ = os.environ.pop(key, default=None)
# reset reuseable mocks
self.mongo_client_mock.reset_mock()
self.mongo_database_mock.reset_mock()
self.mongo_collection_mock.reset_mock()
def test_instantiation_should_fail_when_env_not_set(self):
# ensure env not set
self.tearDown()
# run test
with self.assertRaises(OSError):
_ = DocstoreMongo()