diff --git a/src/database/classes.py b/src/database/classes.py index 2fcda2a..90ffab9 100644 --- a/src/database/classes.py +++ b/src/database/classes.py @@ -52,7 +52,7 @@ class ModelOutputs(BaseModel): def from_random(cls) -> ModelOutputs: """Instantiate with random numbers.""" kwargs = { - field: field_info.annotation.from_random() + field: field_info.annotation.from_random() # type: ignore for field, field_info in cls.model_fields.items() } @@ -124,7 +124,7 @@ class VisualCommunication(BaseModel): return VisualCommunication(name=name, image=image) @field_serializer("image") - def serialize_image(image: Image.Image) -> bytes: + def serialize_image(image: Image.Image) -> bytes: # type: ignore buffer = BytesIO() image.save(buffer, format="JPEG") return buffer.getvalue() diff --git a/src/database/utils.py b/src/database/utils.py index 22bf085..75420c1 100644 --- a/src/database/utils.py +++ b/src/database/utils.py @@ -36,7 +36,7 @@ def get_visual_communication( if with_annotation: query["annotation"] = {"$ne": None} else: - query["annotation"] = None + query["annotation"] = {"$eq": None} data = collection.aggregate([ { "$match": query # find using filters @@ -47,13 +47,12 @@ def get_visual_communication( } } ]) - data = list(data) # read data from cursor object - if len(data) == 0: + data_list = list(data) # read data from cursor object + if len(data_list) == 0: logging.error("failed getting visual communication") raise NoDocumentFoundException() - data = data[0] logging.info("finished") - return VisualCommunication.model_validate(data) + return VisualCommunication.model_validate(data_list[0]) def upsert_predictions( diff --git a/src/web/app.py b/src/web/app.py index 4e6549c..fe465f6 100644 --- a/src/web/app.py +++ b/src/web/app.py @@ -103,13 +103,13 @@ def cycle_visual_communication_data( for elem in annotation_values ] - annotations = { + annotation_map = { key: value for key, value in zip(annotation_keys, annotation_values) } # instantiate ModelOutputs object - annotations = ModelOutputs.from_annotations(**annotations) + annotations = ModelOutputs.from_annotations(**annotation_map) # save data to database upsert_annotations( collection=collection,