import dash
from dash import dcc, html
from dash.dependencies import Input, Output
import os
import re
import json
import pandas as pd
import numpy as np
import plotly.graph_objects as go
from flask import request
from urllib.parse import urlparse, parse_qs
import base64


app = dash.Dash(__name__)

meta_df = pd.read_json(
    "/home/georg/data/arxiv_ml_meta.json",
    dtype={"id": str},
    lines=True,
)
meta_df['submit_date'] = pd.to_datetime(meta_df['submit_date'], unit='ms')

DEFAULT_ORG_SET = set(["Facebook", "Google", "Microsoft", "Amazon", "Nvidia"])

def _get_paper_df(search_term, min_date="2015", min_citation=5, org_set=DEFAULT_ORG_SET):
    # format orgs
    meta_df['orgs_main'] = [
        " & ".join([o for o in s.split(" & ") if o in org_set]) 
        for s in meta_df['orgs'].values
    ]
    meta_df['orgs_main'] = meta_df['orgs_main'].replace("", "other")
    # filter papers
    plot_df = meta_df[
        (
            meta_df['title'].str.contains(r"\b{}\b".format(search_term), flags=re.IGNORECASE) |
            (meta_df['abstract'].str.count(r"\b{}\b".format(search_term), flags=re.IGNORECASE) >= 2)
        ) & (
            (
                (meta_df['submit_date'].dt.year >= int(min_date)) &
                (meta_df['n_citation'] >= min_citation)
            ) | (
                ((meta_df['submit_date'].max() - meta_df['submit_date']).dt.days <= 90) &
                (meta_df['orgs_main'] != "other")
            )
        )
    ].copy()
    plot_df.loc[(
        ((meta_df['submit_date'].max() - meta_df['submit_date']).dt.days <= 90) &
        (meta_df['orgs_main'] != "other")
    ), "n_citation"] += min_citation
    return plot_df

def get_data(search_term):
    plot_df = _get_paper_df(search_term)
    global dfs
    dfs = []
    data = []
    trace_names = list(plot_df['orgs_main'].unique())
    for name in trace_names:
        _df = plot_df[plot_df['orgs_main'] == name].copy()
        dfs.append(_df)
        customdata = _df['orgs'].values
        trace_data = {
            "x": _df["submit_date"],
            "y": _df["n_citation"],
            "name": name,
            "text": _df['title'],
            "customdata": customdata,
            "hovertemplate": "<b>%{text}</b><br>" +
                "%{x|%Y %b} - %{customdata}<br>" +
                "<extra></extra>",
            "mode": "markers",
            'marker': {
                'size': 10,
                'line': {
                    'width': 1,
                    'color': '#888',
                },
                'opacity': 0.75,
            },
        }
        if name == "other":
            trace_data["marker"]["color"] = "#ccc"
        data.append(trace_data)
    return data

loader_style = {
    'height': '100%',
    'width': '100%',
    'display': 'flex',
    'justify-content': 'center',
    'padding-top': '50px',
    'position': 'absolute',
    'z-index': '100',
    'background-color': 'white',
}

spinner_style = {
    'display': 'block',
    'width': '80px',
}

encoded_img = base64.b64encode(open('dash/spinner.gif', 'rb').read()).decode("utf8")

app.layout = html.Div(children=[
    dcc.Location(id='url', refresh=False),
    html.Div(
        id="sidebar-loader-div",
        style=loader_style,
        children=[
            html.Div(
                id="sidebar-loader",
                children=[
                    html.Img(
                        src='data:image/gif;base64,{}'.format(encoded_img),
                        style=spinner_style
                    )
                ]
            )
        ]
    ),
    dcc.Graph(id='paper-graph'),
    dcc.Store(id='clientside-figure-store'),
    html.Div(id="fig-output", children=""),
])

output_css={
    "font-family": "'Helvetica Neue', Helvetica, Arial",
}
output_css2={
    "font-family": "'Helvetica Neue', Helvetica, Arial",
    "font-size": "0.9em",
}

@app.callback(
    Output('clientside-figure-store', 'data'),
    Input('url', 'search'))
def store_data(url_query_str):
    global dfs
    parsed_url = urlparse(url_query_str)
    query_dict = parse_qs(parsed_url.query)
    search_term = query_dict["search_term"][0]
    figdata = get_data(search_term)
    data = {
        "figdata": figdata,
        "title": "Most cited papers for<br>`{}`".format(search_term),
        "npoints": [df.shape[0] for df in dfs]
    }
    return data

app.clientside_callback(
    """
    function(data, clickData) {
        fig_json = {
            'data': data.figdata,
            'layout': {
                'title': {
                    'text': data.title,
                    'y': 0.88,
                    'x': 0.5,
                    'xanchor': 'center',
                    'yanchor': 'top'
                },
                'xaxis': {'title': 'Publication date'},
                'yaxis': {'type': 'log', 'visible': false, 'showticklabels': false}
            },
        }
        document.getElementById("sidebar-loader-div").style['display'] = 'none';
        if (typeof clickData == 'undefined')
            return fig_json

        // update size of points
        trace_idx = clickData.points[0].curveNumber
        point_idx = clickData.points[0].pointNumber
        for (let n = 0; n < data.npoints.length; n++) {
          data.figdata[n].marker.size = 10;
        }
        const size_vec = [];
        for (let n = 0; n < data.npoints[trace_idx]; n++) {
            if (n == point_idx)
                size_vec[n] = 20;
            else
                size_vec[n] = 10;
        }
        data.figdata[trace_idx].marker.size = size_vec;

        return fig_json
    }
    """,
    Output('paper-graph', 'figure'),
    Input('clientside-figure-store', 'data'),
    Input('paper-graph', 'clickData'),
)

@app.callback(
    Output('fig-output', 'children'),
    Input('paper-graph', 'clickData'))
def display_click_data(clickData):
    if clickData is None:
        return None
    global dfs
    trace_idx = clickData["points"][0]["curveNumber"]
    point_idx = clickData["points"][0]["pointNumber"]

    row = dfs[trace_idx].iloc[point_idx]
    info_str = row["submit_date"].strftime("%Y %b")
    if len(row["orgs"]) > 0:
        info_str += " - " + row["orgs"]
    info_str += ", {} citations".format(row["n_citation"])
    output_layout = html.Div(children=[
        html.Div(children=html.A(
            href="https://arxiv.org/abs/{}".format(row['id']),
            target="_blank",
            children="paper link",
        ), style=output_css2),
        html.Div(children=html.B(children=row['title']), style=output_css),
        html.Div(children=html.Pre(children=info_str)),
        html.Div(children=row['abstract'], style=output_css2)
    ])
    return output_layout

if __name__ == '__main__':
    context = ('cert/383fff33778762e8.crt', 'cert/383fff33778762e8.key')
    app.run_server(debug=False, host='0.0.0.0', port=7880, ssl_context=context)
