252 lines
7.6 KiB
Python
252 lines
7.6 KiB
Python
from __future__ import annotations
|
|
|
|
import logging
|
|
import os
|
|
|
|
import dash_bootstrap_components as dbc
|
|
from dash import ALL
|
|
from dash import Dash
|
|
from dash import Input
|
|
from dash import Output
|
|
from dash import State
|
|
from dash_auth import BasicAuth
|
|
from pydantic import ValidationError
|
|
|
|
from .layout import app_layout
|
|
from core.database import connect
|
|
from core.database import count_documents
|
|
from core.database import get_visual_communication
|
|
from core.database import NoDocumentFoundException
|
|
from core.database import upsert_annotation
|
|
from core.database import upsert_visual_communication
|
|
from core.database import VisualCommunication
|
|
from core.dto import ModelData
|
|
|
|
# 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('success-element', 'is_open'),
|
|
Output('success-element', 'children'),
|
|
Input('success-message', 'data'),
|
|
)
|
|
def show_success(
|
|
msg: str | None,
|
|
):
|
|
if msg is None or msg == '':
|
|
return False, ''
|
|
logging.info(f"updated success message: {msg}")
|
|
return True, msg
|
|
|
|
|
|
@app.callback(
|
|
Output('progress-bar', 'value'),
|
|
Output('progress-bar', 'max'),
|
|
Output('progress-bar', 'label'),
|
|
Input('next-button', 'n_clicks'),
|
|
Input('upload-data', 'filename'),
|
|
)
|
|
def update_progress_bar(
|
|
n_clicks: int,
|
|
filename_list: list[str] | None,
|
|
):
|
|
logging.info('began updating progress bar')
|
|
global collection
|
|
# get values from database
|
|
num_total = count_documents(
|
|
collection=collection,
|
|
only_with_annotation=False,
|
|
)
|
|
num_handled = count_documents(
|
|
collection=collection,
|
|
only_with_annotation=True,
|
|
)
|
|
limit = int(num_total/20)
|
|
label_str = f"{num_handled}/{num_total}" if num_handled >= limit else ''
|
|
return num_handled, num_total, label_str
|
|
|
|
|
|
@app.callback(
|
|
Output('success-message', 'data', allow_duplicate=True),
|
|
Output('alert-message', 'data', allow_duplicate=True),
|
|
Input('upload-data', 'contents'),
|
|
State('upload-data', 'filename'),
|
|
prevent_initial_call=True,
|
|
)
|
|
def upload_images(
|
|
content_list: list[str] | None,
|
|
filename_list: list[str] | None,
|
|
):
|
|
logging.info('began uploading images')
|
|
global collection
|
|
# stop early if possible
|
|
if content_list is None or filename_list is None:
|
|
logging.warning('content_list is None -> nothing to upload')
|
|
return None, 'nothing selected to upload'.title()
|
|
# build list of visual communication
|
|
visual_communication_list = []
|
|
for content, filename in zip(content_list, filename_list):
|
|
try:
|
|
image = VisualCommunication.decode_image(content)
|
|
vis_com = VisualCommunication(
|
|
name=filename,
|
|
image=image,
|
|
)
|
|
except Exception as exc:
|
|
logging.warning(
|
|
(
|
|
'failed creating VisualCommunication object '
|
|
f"from file {filename}"
|
|
),
|
|
exc,
|
|
)
|
|
else:
|
|
visual_communication_list.append(vis_com)
|
|
# upsert documents
|
|
success = upsert_visual_communication(
|
|
collection=collection,
|
|
visual_communication_list=visual_communication_list,
|
|
)
|
|
if success:
|
|
logging.info(f"succesfully uploaded {len(filename_list)} images")
|
|
return 'successfully uploaded images'.title(), None
|
|
logging.warning('failed inserting images into database')
|
|
return None, 'failed uploading images'.title()
|
|
|
|
|
|
@app.callback(
|
|
Output('alert-message', 'data', allow_duplicate=True),
|
|
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 = ModelData.from_annotations(**annotation_map)
|
|
# save data to database
|
|
upsert_annotation(
|
|
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)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
app.run(debug=True)
|