152 lines
4.5 KiB
Python
152 lines
4.5 KiB
Python
from dash import Dash, Input, Output, State, ALL
|
|
import dash_bootstrap_components as dbc
|
|
from dash_auth import BasicAuth
|
|
import logging
|
|
from typing import List
|
|
from pydantic import ValidationError
|
|
import os
|
|
|
|
from .layout import app_layout
|
|
from src.database import (
|
|
connect,
|
|
get_visual_communication,
|
|
NoDocumentFoundException,
|
|
upsert_annotations,
|
|
ModelOutputs
|
|
)
|
|
|
|
# setup app
|
|
app = Dash(__name__, external_stylesheets=[dbc.themes.BOOTSTRAP])
|
|
app.title = "visual critical discourse analysis".title()
|
|
app.layout = app_layout
|
|
server = app.server
|
|
|
|
# setup authentication
|
|
AUTH_DICT = {
|
|
os.getenv("DASH_AUTH_USERNAME"): os.getenv("DASH_AUTH_PASSWORD")
|
|
}
|
|
BasicAuth(app, AUTH_DICT)
|
|
|
|
# connect to database
|
|
collection, db, client = connect()
|
|
|
|
|
|
# define callbacks
|
|
@app.callback(
|
|
Output("alert-element", "is_open"),
|
|
Output("alert-element", "children"),
|
|
Input("alert-message", "data")
|
|
)
|
|
def show_alert(
|
|
msg: str | None
|
|
):
|
|
if msg is None or msg == "":
|
|
return False, ""
|
|
logging.info(f"updated alert message: {msg}")
|
|
return True, msg
|
|
|
|
|
|
@app.callback(
|
|
Output("alert-message", "data"),
|
|
Output("vis-com-name", "data"),
|
|
Output("image-container", "src"),
|
|
Output({"type": "annotation", "index": ALL}, "value"),
|
|
Input("next-button", "n_clicks"),
|
|
State("vis-com-name", "data"),
|
|
State("image-container", "src"),
|
|
State({"type": "annotation", "index": ALL}, "id"),
|
|
State({"type": "annotation", "index": ALL}, "value"),
|
|
prevent_initial_call=True,
|
|
)
|
|
def cycle_visual_communication_data(
|
|
n_clicks: int,
|
|
vis_com_name: str,
|
|
image_src: str,
|
|
annotation_keys: List,
|
|
annotation_values: List,
|
|
):
|
|
logging.info("began cycling visual communication data")
|
|
global collection
|
|
# prepare default response
|
|
response = [
|
|
"",
|
|
vis_com_name,
|
|
image_src,
|
|
annotation_values
|
|
]
|
|
# check if next-button clicked
|
|
if n_clicks == 0:
|
|
logging.info("stopping early: next-button has not yet been clicked")
|
|
return response
|
|
# check if visual communication name is set
|
|
if len(vis_com_name) > 0:
|
|
logging.info("saving annotations to database: %s", vis_com_name)
|
|
try:
|
|
# extract option keys
|
|
annotation_keys = [
|
|
elem["index"]
|
|
for elem in annotation_keys
|
|
]
|
|
# ensure all options are set
|
|
logging.info(annotation_keys)
|
|
for option, value in zip(annotation_keys, annotation_values):
|
|
if value is None:
|
|
raise ValueError(f"{option} is not set")
|
|
# prepare data to save
|
|
annotation_keys = [
|
|
elem.replace(' ', '_')
|
|
for elem
|
|
in annotation_keys
|
|
]
|
|
annotation_values = [
|
|
elem.replace(' ', '_').lower()
|
|
for elem
|
|
in annotation_values
|
|
]
|
|
annotation_map = {
|
|
key: value
|
|
for key, value
|
|
in zip(annotation_keys, annotation_values)
|
|
}
|
|
# instantiate ModelOutputs object
|
|
annotations = ModelOutputs.from_annotations(**annotation_map)
|
|
# save data to database
|
|
upsert_annotations(
|
|
collection=collection,
|
|
vis_com_name=vis_com_name,
|
|
annotations=annotations
|
|
)
|
|
except (ValueError, ValidationError) as exc:
|
|
msg = f"failed saving annotation: {exc}"
|
|
logging.warning(msg)
|
|
response[0] = msg
|
|
return tuple(response)
|
|
# get new visual communication
|
|
logging.info("trying to get new visual communication")
|
|
try:
|
|
# get data
|
|
vis_com = get_visual_communication(
|
|
collection=collection,
|
|
with_annotation=False
|
|
)
|
|
# set variables
|
|
vis_com_name = vis_com.name
|
|
image_src = vis_com.webencoded_image()
|
|
if vis_com.prediction is not None:
|
|
# TODO: update to use optional predictions
|
|
pass
|
|
else:
|
|
# reset annotations
|
|
annotation_values = [None for elem in annotation_values]
|
|
except NoDocumentFoundException:
|
|
msg = "no unannotated data in database"
|
|
logging.warning(msg)
|
|
response[0] = msg
|
|
return tuple(response)
|
|
else:
|
|
response[1] = vis_com_name
|
|
response[2] = image_src
|
|
response[3] = annotation_values
|
|
logging.info("finished getting visual communication: %s", vis_com_name)
|
|
return tuple(response)
|