import pandas as pd
from collections import defaultdict, Counter
import matplotlib.pyplot as plt
import seaborn as sns
import numpy as np

sns.set_style("dark")


class Elo:
    def __init__(self, k, homefield=0):
        self.ratingDict = {}
        self.k = k
        self.homefield = homefield

    def addPlayer(self, name, rating=1500):
        self.ratingDict[name] = rating

    def gameOver(self, winner, loser, winnerHome):
        if winnerHome:
            result = self.expectResult(
                self.ratingDict[winner] + self.homefield, self.ratingDict[loser]
            )
        else:
            result = self.expectResult(
                self.ratingDict[winner], self.ratingDict[loser] + self.homefield
            )
        # print result
        self.ratingDict[winner] = self.ratingDict[winner] + self.k * (1 - result)
        self.ratingDict[loser] = self.ratingDict[loser] + self.k * (0 - (1 - result))

    def expectResult(self, p1, p2):
        exp = (p2 - p1) / 400.0
        return 1 / ((10.0 ** (exp)) + 1)


def get_preferfence_counts(input_df, title_name: str = ""):
    input_df = input_df.sort_values(by=["request_id", "preference"])
    preferences = input_df["preference"].values
    neg_prefs, pos_prefs = preferences[::2], preferences[1::2]
    assert set(neg_prefs) == {False}
    assert set(pos_prefs) == {True}
    paired_models = input_df["model_name"].values
    neg_models, pos_models = paired_models[::2], paired_models[1::2]
    win_counter = Counter()
    t_model = defaultdict(int)
    # init the elo league, oooh yeah
    eloLeague = Elo(k=0.01, homefield=0)
    for model in input_df["model_name"].unique():
        # if model == "chirp-v3-engine-d":
        #     eloLeague.addPlayer(model, 1050)
        # elif model == "chirp-v3-engine-d":
        #     eloLeague.addPlayer(model, 950)
        # else:
        #
        eloLeague.addPlayer(model, 1000)
    for i, (pos_model, neg_model) in enumerate(zip(pos_models, neg_models)):
        win_counter[f"{pos_model}_win_over_{neg_model}"] += 1
        sorted_models = sorted([pos_model, neg_model])
        t_model[f"{sorted_models[0]}_{sorted_models[1]}"] += 1
        if pos_model != neg_model:
            eloLeague.gameOver(pos_model, neg_model, 0)
        # if i == int(pos_models.shape[0] * 0.8):
        #     print("80% done")
        #     for model_name in eloLeague.ratingDict.keys():
        #         print(model_name, eloLeague.ratingDict[model_name])
        # if i == int(pos_models.shape[0] * 0.9):
        #     print("90% done")
        #     for model_name in eloLeague.ratingDict.keys():
        #         print(model_name, eloLeague.ratingDict[model_name])
    # for model_name in eloLeague.ratingDict.keys():
    #     print(model_name, eloLeague.ratingDict[model_name])

    models_performance = defaultdict(dict)
    for k in sorted(win_counter.keys()):
        v = win_counter[k]
        two_models = k.split("_win_over_")
        win_model, lose_model = two_models[0], two_models[1]
        sorted_tow_models = sorted([win_model, lose_model])
        total_model_count = t_model[f"{sorted_tow_models[0]}_{sorted_tow_models[1]}"]
        win_ratio = v / total_model_count
        win_err = np.sqrt(win_ratio * (1 - win_ratio) / total_model_count)
        models_performance[win_model][lose_model] = win_ratio
        models_performance[win_model][lose_model + "_err"] = win_err
        print(f"{k}, win ratio {win_ratio:.3f}, counts {v}")

    # for model_name in eloLeague.ratingDict.keys():
    #     print(model_name, eloLeague.ratingDict[model_name])
    # print(models_performance)
    count_matrix = []
    count_err_matrix = []
    kept_win_models = []
    for win_model in sorted(models_performance.keys()):
        count_row = []
        count_err_row = []
        for lose_model in sorted(models_performance.keys()):
            win_ratio = models_performance[win_model].get(lose_model, 0.0)
            win_err = models_performance[win_model].get(lose_model + "_err", 0.0)
            if win_ratio == 1 or win_ratio == 0:
                win_ratio = 0
                win_err = 0
            else:
                # normalize it
                win_ratio -= 0.5
            # don't show the negative
            if win_ratio < 0:
                win_ratio = 0
                win_err = 0
            # convert to percentage
            count_row.append(win_ratio * 100)
            count_err_row.append(
                f" {win_ratio * 100:.3f} \n±{win_err * 100:.3f}"
                if win_ratio != 0
                else ""
            )
        if all(i == 0 for i in count_row):
            continue
        count_matrix.append(count_row)
        count_err_matrix.append(count_err_row)
        kept_win_models.append(win_model)

    # print(count_matrix)
    count_matrix = np.array(count_matrix)
    count_err_matrix = np.array(count_err_matrix)
    # print(count_matrix, count_matrix.shape)
    # rare case that there is nothing to compare...
    if count_matrix.shape[0] == 0:
        return
    # print(count_err_matrix, count_err_matrix.shape)
    # mask = np.zeros_like(count_matrix, dtype='bool')
    # mask[np.triu_indices_from(mask)] = True

    # clear the empty columns
    kept_lose_models = []
    empty_indeces = []
    for index, lose_model in enumerate(sorted(models_performance.keys())):
        if all(i == 0 for i in count_matrix[:, index]):
            empty_indeces.append(index)
        else:
            kept_lose_models.append(lose_model)
    count_matrix = np.delete(
        count_matrix,
        empty_indeces,
        axis=1,
    )
    count_err_matrix = np.delete(
        count_err_matrix,
        empty_indeces,
        axis=1,
    )

    model_names = sorted(models_performance.keys())
    plt.clf()
    plt.figure(figsize=(10, 8))
    sns.heatmap(
        count_matrix,
        annot=count_err_matrix,
        fmt="",
        cmap="viridis",
        xticklabels=kept_lose_models,
        yticklabels=kept_win_models,  # win models
        annot_kws={"fontsize": 12},
    )
    plt.xlabel("Lose", fontsize=14)
    plt.ylabel("Win", fontsize=14)
    plt.text(
        4.5,
        0.0,
        "ELOs \n"
        + " \n".join(
            [
                f"{model_name.replace('chirp-v3-engine-', '')}: {eloLeague.ratingDict[model_name]:.1f}"
                for model_name in model_names
            ]
        ),
        fontsize=10,
    )
    plt.title(
        f'Model Preference \n {"" + title_name if title_name else ""}',
        fontsize=24,
        loc="left",
    )
    plt.show()
