moved functions into separate files

This commit is contained in:
Brian Bjarke Jensen
2024-03-20 19:18:57 +01:00
parent 657e9c4b1d
commit dc9f611e17
8 changed files with 174 additions and 161 deletions
-161
View File
@@ -1,161 +0,0 @@
from __future__ import annotations
import logging
from pymongo.collection import Collection
from .classes import NoDocumentFoundException
from .classes import VisualCommunication
from core.dto import ModelData
def count_documents(
collection: Collection,
has_annotation: bool = False,
) -> int:
"""
Get the total number of documents
in database that matches the filters.
"""
assert isinstance(collection, Collection)
assert isinstance(has_annotation, bool)
# build query
query = {}
if has_annotation:
query['annotation'] = {'$ne': None}
return collection.count_documents(filter=query)
def list_names(
collection: Collection,
has_annotation: bool = True,
) -> list[str]:
"""List the names of entries that match the filters."""
assert isinstance(collection, Collection)
assert isinstance(has_annotation, bool)
# build query
query = {}
if has_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
def total_documents(
collection: Collection,
) -> int:
"""Get total number of documents in database."""
return collection.count_documents(filter={})
def total_annotated(
collection: Collection,
) -> int:
"""Get total number of annotated documents in database."""
query = {
'annotation': {
'$ne': None,
},
}
return collection.count_documents(filter=query)
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([
{
'$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:
logging.error('failed getting visual communication')
raise NoDocumentFoundException()
logging.info('finished')
return VisualCommunication.model_validate(data_list[0])
def upsert_predictions(
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')
def upsert_annotations(
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')
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
+8
View File
@@ -0,0 +1,8 @@
from __future__ import annotations
from .count_documents import count_documents
from .get_visual_communication import get_visual_communication
from .list_names import list_names
from .upsert_annotation import upsert_annotations
from .upsert_prediction import upsert_predictions
from .upsert_visual_communication import upsert_visual_communication
+20
View File
@@ -0,0 +1,20 @@
from __future__ import annotations
from pymongo.collection import Collection
def count_documents(
collection: Collection,
has_annotation: bool = False,
) -> int:
"""
Get the total number of documents
in database that matches the filters.
"""
assert isinstance(collection, Collection)
assert isinstance(has_annotation, bool)
# build query
query = {}
if has_annotation:
query['annotation'] = {'$ne': None}
return collection.count_documents(filter=query)
@@ -0,0 +1,36 @@
from __future__ import annotations
import logging
from pymongo.collection import Collection
from core.database import NoDocumentFoundException
from core.database import 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([
{
'$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:
logging.error('failed getting visual communication')
raise NoDocumentFoundException()
logging.info('finished')
return VisualCommunication.model_validate(data_list[0])
+27
View File
@@ -0,0 +1,27 @@
from __future__ import annotations
from pymongo.collection import Collection
def list_names(
collection: Collection,
has_annotation: bool = True,
) -> list[str]:
"""List the names of entries that match the filters."""
assert isinstance(collection, Collection)
assert isinstance(has_annotation, bool)
# build query
query = {}
if has_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
+30
View File
@@ -0,0 +1,30 @@
from __future__ import annotations
import logging
from pymongo.collection import Collection
from core.dto import ModelData
def upsert_annotations(
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 core.dto import ModelData
def upsert_predictions(
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')
@@ -0,0 +1,23 @@
from __future__ import annotations
from pymongo.collection import Collection
from core.database 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