renamed module

This commit is contained in:
brian
2024-10-20 20:00:13 +00:00
parent d1ee43e135
commit e2338ad710
26 changed files with 63 additions and 82 deletions
+3
View File
@@ -0,0 +1,3 @@
"""Database utils module content."""
from .connect_mongodb import connect_mongodb
+30
View File
@@ -0,0 +1,30 @@
"""Definition of function to connect to database using environment
variables."""
import logging
import os
from dotenv import load_dotenv
from pymongo import MongoClient
def connect_mongodb():
"""Connect to MongoDB using env vars."""
# load env vars
load_dotenv()
necessary_env_vars = [
'MONGO_HOST',
'MONGO_DB',
'MONGO_COLLECTION',
]
for env_var in necessary_env_vars:
assert env_var in os.environ, f"{env_var} not found"
# connect to database
client = MongoClient(os.getenv('MONGO_HOST'))
db = client[os.getenv('MONGO_DB')]
# extract collection
collection = db[os.getenv('MONGO_COLLECTION')]
# set unique index on "name"
collection.create_index('name', unique=True)
logging.debug('finished')
return collection, db, client
+20
View File
@@ -0,0 +1,20 @@
"""Definition of function to count documents in database."""
from __future__ import annotations
from pymongo.collection import Collection
def count_documents(
collection: Collection,
only_with_annotation: bool = False,
) -> int:
"""Get the total number of documents in database that matches the
filters."""
assert isinstance(collection, Collection)
assert isinstance(only_with_annotation, bool)
# build query
query = {}
if only_with_annotation:
query['annotation'] = {'$ne': None}
return collection.count_documents(filter=query)
+16
View File
@@ -0,0 +1,16 @@
"""Definition of get_dataset function."""
from typing import Literal
from pymongo.collection import Collection
def get_dataset(
collection: Collection,
type: Literal['train', 'test', 'validation'],
) -> list[str]:
"""Get list of data names for the corresponding type."""
assert isinstance(collection, Collection)
assert isinstance(type, str)
assert type in ['train', 'test', 'validation']
return []
+39
View File
@@ -0,0 +1,39 @@
"""Definition of function to get visual communication from database."""
from __future__ import annotations
import logging
from pymongo.collection import Collection
from shared.mongodb import NoDocumentFoundException, VisualCommunication
def get_visual_communication(
collection: Collection,
with_annotation: bool = False,
) -> VisualCommunication:
"""Get a random visual communication from the database."""
query = {}
if with_annotation:
query['annotation'] = {'$ne': None}
else:
query['annotation'] = {'$eq': None}
data = collection.aggregate(
pipeline=[
{
'$match': query, # find using filters
},
{
'$sample': {
'size': 1, # get one random
},
},
],
)
data_list = list(data) # read data from cursor object
if len(data_list) == 0:
raise NoDocumentFoundException()
vis_com = VisualCommunication.model_validate(data_list[0])
logging.debug('finished')
return vis_com
+30
View File
@@ -0,0 +1,30 @@
"""Definition of function to list names of all visual communication documents
in database."""
from __future__ import annotations
from pymongo.collection import Collection
def list_names(
collection: Collection,
only_with_annotation: bool = True,
) -> list[str]:
"""List the names of entries that match the filters."""
assert isinstance(collection, Collection)
assert isinstance(only_with_annotation, bool)
# build query
query = {}
if only_with_annotation:
query['annotation'] = {'$ne': None}
# execute query
res_list = collection.find(
filter=query,
projection={
'_id': False,
'name': True,
},
)
# extract information
name_list = [elem['name'] for elem in res_list]
return name_list
+38
View File
@@ -0,0 +1,38 @@
from __future__ import annotations
import logging
from pymongo.collection import Collection
from shared.mongodb import Dataset
def save_dataset(
collection: Collection,
dataset: Dataset,
) -> None:
"""Save dataset to database."""
res = collection.insert_one(
document=dataset.model_dump(),
)
logging.debug('inserted document: %s', res)
if __name__ == '__main__':
from dotenv import load_dotenv
load_dotenv('local.env')
from shared.mongodb import connect_mongodb, list_names
# connect to database
collection, db, client = connect_mongodb()
print(client.server_info())
name_list = list_names(collection=collection, only_with_annotation=True)
ds = Dataset.new_from_name_list(name_list=name_list)
print(ds)
# save_dataset(
# collection=collection,
# dataset=ds
# )
+30
View File
@@ -0,0 +1,30 @@
from __future__ import annotations
import logging
from pymongo.collection import Collection
from shared.dto import ModelData
def upsert_annotation(
collection: Collection,
vis_com_name: str,
annotations: ModelData,
) -> None:
"""Upserts annotation data in the database."""
query = {
'name': vis_com_name,
}
update = {
'$set': {
'annotation': annotations.model_dump(),
},
}
res = collection.update_one(
filter=query,
update=update,
upsert=True,
)
logging.info('upserted document: %s', res)
logging.info('finished')
+30
View File
@@ -0,0 +1,30 @@
from __future__ import annotations
import logging
from pymongo.collection import Collection
from shared.dto import ModelData
def upsert_prediction(
collection: Collection,
vis_com_name: str,
predictions: ModelData,
) -> None:
"""Upsert prediction data in the database."""
query = {
'name': vis_com_name,
}
update = {
'$set': {
'prediction': predictions.model_dump(),
},
}
res = collection.update_one(
filter=query,
update=update,
upsert=True,
)
logging.debug('upserted document: %s', res)
logging.info('finished')
+19
View File
@@ -0,0 +1,19 @@
from __future__ import annotations
from pymongo.collection import Collection
from shared.mongodb import VisualCommunication
def upsert_visual_communication(
collection: Collection,
visual_communication_list: list[VisualCommunication],
) -> bool:
"""Upsert VisualCommunication object in the database.
Returns bool stating success.
"""
response = collection.insert_many(
[vis_com.model_dump() for vis_com in visual_communication_list],
)
return response.acknowledged