import numpy as np
import pandas as pd
import math

import matplotlib.pyplot as plt
import seaborn as sns
import os

import time

from script_cfg import *

##################################################### GET BPP PATH

class BPP:
    USUAL_PAIRS = [ ('G', 'C'), ('C', 'G'), ('A', 'U'), ('U', 'A') ]

    
##################################################### GET BPP PATH

def get_bpp_path(seq):
    should_break = False
    for dir1 in os.listdir(CFG.BPP):
        for dir2 in os.listdir(f"{CFG.BPP}/{dir1}"):
            for dir3 in os.listdir(f"{CFG.BPP}/{dir1}/{dir2}"):
                for file in os.listdir(f"{CFG.BPP}/{dir1}/{dir2}/{dir3}"):
                    if file.strip('.txt') == seq:
                        should_break = True
                        break
                    if should_break:
                        break
                if should_break:
                    break
            if should_break:
                break
        if should_break:
            break
        
    path = f"{CFG.BPP}/{dir1}/{dir2}/{dir3}/{file}"
    return path

##################################################### GET STATS ON FEW BPP AND PLOT

def plot_proba_dist_seqs(train_partial, seqs_ids_partial):
    for seq_id in seqs_ids_partial:
        path = get_path(seq_id)
        wc_proba = pd.read_csv(path, header=None, sep=' ', names=['base_1', 'base_2', 'proba'])
        seq_len = len(train_partial[ train_partial.sequence_id == seq_id].sequence.values[0])
        plot_proba_dist(wc_proba, seq_id, seq_len)
    
    
def plot_proba_dist(df, seq_id, seq_len):
    fig = plt.figure(figsize=(10, 1))
    fig.suptitle(f"{seq_id} - Len: {seq_len} - Pairs: {len(df)} ",color='#3f3f3f', y=1.05, fontweight='bold', fontsize='small')
    mosaic = '''AB
                AB
                '''
    mosaix_ = ['A', 'B']
    ax = fig.subplot_mosaic(mosaic)
    sns.histplot(data=df, x='proba', bins=40, edgecolor=None, ax=ax['A']);
    sns.boxplot(data=df, x='proba', orient='h', saturation=0.9, 
                showmeans=True, meanprops=dict(markerfacecolor='w', markeredgecolor='w', marker='o'), medianprops=dict(color='w', alpha=1, linewidth=2),
                boxprops=dict(linewidth=False), whiskerprops=dict(color='#151515', linewidth=1), capprops=dict(color='#151515', linewidth=1), showfliers=True, fliersize=3, flierprops=dict(markerfacecolor='#6a6a6a', markeredgecolor='#6a6a6a', marker='d'),
                ax=ax['B'])
    
    for i in mosaix_:
        ax[i].grid(axis="both", lw=0.5, ls=':')
        ax[i].spines[['right', 'top']].set_visible(False)
        ax[i].spines['left'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
        ax[i].spines['bottom'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
        ax[i].tick_params(axis='both', color='#3f3f3f', colors='#3f3f3f', labelsize='x-small')
        
        
##################################################### GET STATS ON MORE BPP AND PLOT
        
def get_wc_proba_stats(train_partial, start, end):
    
    if (start + 1 > len(train_partial)) | (end + 1 > len(train_partial)) | (start >= end):
        return f"start and / or end args are not compatibles with train_partial lenght"
    
    else:
    
        wc_proba_stats = pd.DataFrame(columns = ['sequence_id', 'sequence', 'pairs_comb', 
                                              'pairs', 'pairs_pct', 'mean_proba', 'median_proba',
                                              'pairs_unusual', 'pairs_unusual_pct_overpair', 
                                              'pairs_high_proba', 'pairs_high_proba_pct', 
                                              'pairs_low_proba', 'pairs_low_proba_pct'],
                                   index = range(end - start))

        start_time = time.perf_counter()

        for j, seq_id in enumerate(train_partial.sequence_id.values[start:end]):

            path = get_bpp_path(seq_id)
            seq = train_partial[ train_partial.sequence_id == seq_id].sequence.values[0]
            wc_proba = pd.read_csv(path, header=None, sep=' ', names=['base_1', 'base_2', 'proba'])
            wc_proba.sort_values(by='proba', ascending=False, inplace=True)

            pairs_high_proba = 0
            pairs_low_proba = 0
            pairs_unusual_high_proba = 0
            pairs_unusual_low_proba = 0

            for i in range(len(wc_proba)):
                pairs_bases = (seq[wc_proba.base_1.values[i] - 1], seq[wc_proba.base_2.values[i] - 1])

                if wc_proba.proba.values[i] >= 0.005:
                    pairs_high_proba += 1
                    if pairs_bases not in BPP.USUAL_PAIRS:
                        pairs_unusual_high_proba += 1
                else:
                    pairs_low_proba += 1
                    if pairs_bases not in BPP.USUAL_PAIRS:
                        pairs_unusual_low_proba += 1

            wc_proba_stats.loc[j, 'sequence_id'] = seq_id
            wc_proba_stats.loc[j, 'sequence'] = seq
            wc_proba_stats.loc[j, 'pairs_high_proba'] = pairs_high_proba
            wc_proba_stats.loc[j, 'pairs_low_proba'] = pairs_low_proba
            wc_proba_stats.loc[j, 'pairs_unusual_high_proba'] = pairs_unusual_high_proba
            wc_proba_stats.loc[j, 'pairs_unusual_low_proba'] = pairs_unusual_low_proba

            wc_proba_stats.loc[j, 'mean_proba'] = np.mean(wc_proba.proba) * 100
            sorted_ = sorted(list(wc_proba.proba.values))
            if len(sorted_)%2 == 0:
                wc_proba_stats.loc[j, 'median_proba'] = ( sorted_[ int(len(sorted_) / 2 - 1) ] + sorted_[ int(len(sorted_) / 2) ] ) / 2 * 100
            else:
                wc_proba_stats.loc[j, 'median_proba'] = sorted_[ int(len(sorted_) / 2 - 1) ] * 100
            
            if j%100 == 0:
                print(f"{j}, {time.perf_counter() - start_time:0.1f} sec")
                
        wc_proba_stats.pairs_comb = wc_proba_stats.sequence.apply(lambda row: math.comb(len(row), 2))

        wc_proba_stats.pairs = wc_proba_stats.pairs_high_proba + wc_proba_stats.pairs_low_proba
        wc_proba_stats.pairs_pct = wc_proba_stats.pairs / wc_proba_stats.pairs_comb * 100

        wc_proba_stats.pairs_unusual = wc_proba_stats.pairs_unusual_high_proba + wc_proba_stats.pairs_unusual_low_proba
        wc_proba_stats.pairs_unusual_pct_overpair = wc_proba_stats.pairs_unusual / wc_proba_stats.pairs * 100

        wc_proba_stats.pairs_high_proba_pct = wc_proba_stats.pairs_high_proba / wc_proba_stats.pairs_comb
        wc_proba_stats.pairs_low_proba_pct = wc_proba_stats.pairs_low_proba / wc_proba_stats.pairs_comb * 100

        wc_proba_stats.pairs = wc_proba_stats.pairs.astype("float32")
        wc_proba_stats.pairs_pct = wc_proba_stats.pairs_pct.astype("float32")
        wc_proba_stats.mean_proba = wc_proba_stats.mean_proba.astype("float32")
        wc_proba_stats.median_proba = wc_proba_stats.median_proba.astype("float32")
        wc_proba_stats.pairs_unusual = wc_proba_stats.pairs_unusual.astype("float32")
        wc_proba_stats.pairs_unusual_pct_overpair = wc_proba_stats.pairs_unusual_pct_overpair.astype("float32")
        wc_proba_stats.pairs_high_proba = wc_proba_stats.pairs_high_proba.astype("float32")
        wc_proba_stats.pairs_high_proba_pct = wc_proba_stats.pairs_high_proba_pct.astype("float32")
        wc_proba_stats.pairs_low_proba = wc_proba_stats.pairs_low_proba.astype("float32")
        wc_proba_stats.pairs_low_proba_pct = wc_proba_stats.pairs_low_proba_pct.astype("float32")
        wc_proba_stats.pairs_unusual_high_proba = wc_proba_stats.pairs_unusual_high_proba.astype("float32")
        wc_proba_stats.pairs_unusual_low_proba = wc_proba_stats.pairs_unusual_low_proba.astype("float32")

        print(f"Timer_end, {time.perf_counter() - start_time:0.1f} sec")
    
        return wc_proba_stats

def plot_wc_proba_stats(wc_proba_stats, cols_):
        
    ##### Figure
    mosaic = '''ABCDEF
                ABCDEF
                IJKLMN
                '''

    mosaiq_hist = ['A', 'B', 'C', 'D', 'E', 'F']
    mosaiq_box = ['I', 'J', 'K', 'L', 'M', 'N']

    fig = plt.figure(constrained_layout=False, figsize=(20, 3))
    fig.subplots_adjust(left=0.05, bottom=0.05, right=0.95, top=0.95, wspace=0.2, hspace=1)
    ax = fig.subplot_mosaic(mosaic)
    fig.suptitle('Distribution of WC pairs probabilities (partial train)', color='#3f3f3f', y=1.05, fontweight='bold', fontsize='large')

    ##### Data
    for i, metric in zip(mosaiq_hist, cols_):
        sns.histplot(data=wc_proba_stats[cols_], x=metric, bins=20, label=metric, shrink=1, edgecolor=None, kde=True, ax=ax[i]);

    for i, metric in zip(mosaiq_box, cols_):
        sns.boxplot(data=wc_proba_stats[cols_], x=metric, orient='h', saturation=0.9, showmeans=True, meanprops=dict(markerfacecolor='w', markeredgecolor='w', marker='o'), medianprops=dict(color='w', alpha=1, linewidth=2), boxprops=dict(linewidth=False), 
                    whiskerprops=dict(color='#151515', linewidth=1), capprops=dict(color='#151515', linewidth=1), showfliers=True, fliersize=3, flierprops=dict(markerfacecolor='#6a6a6a', markeredgecolor='#6a6a6a', marker='d'),
                    ax=ax[i])

    ##### Style
    for i in mosaiq_hist + mosaiq_box:
        ax[i].set_xlabel(ax[i].get_xlabel(), loc='center', color='#3f3f3f', fontsize='small')
        ax[i].set_ylabel('', loc='center', color='#3f3f3f', fontsize='small', rotation='vertical', labelpad=4)
        ax[i].grid(axis="both", lw=0.5, ls=':')
        ax[i].spines[['right', 'top']].set_visible(False)
        ax[i].spines['left'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
        ax[i].spines['bottom'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
        ax[i].tick_params(axis='both', color='#3f3f3f', colors='#3f3f3f', labelsize='x-small')
        

##################################################### PLOT BASE-PAIR AND BASE-POSITION PAIRING PROBA
# (ON SEVERAL SEQ)

def plot_pairing_proba(bpps_, pps_, figsize=(20, 4)):
    fig, ax = plt.subplots(nrows=2, ncols=len(bpps_), figsize=(20, 4))
    fig.suptitle('Base-pair and base-position proba of pairing',color='#3f3f3f', y=1.05, fontweight='bold', fontsize='large')
    fig.subplots_adjust(left=0.05, bottom=0.05, right=0.95, top=0.95, wspace=0.2, hspace=0.2)

    for i in range(len(bpps_)):
        ax[0, i].imshow(bpps_[i], origin='lower', cmap='gist_heat_r')
        sns.lineplot(pps_[i], ls='-', lw='0.5', marker='', color='#dc0073ff', ax=ax[1, i]);

    for j in range(len(bpps_)):
        for i in range(2):
            ax[i, j].axis('tight')
            ax[i, j].set_xlabel(ax[i, j].get_xlabel(), loc='center', color='#3f3f3f', fontsize='small')
            ax[i, j].set_ylabel(ax[i, j].get_ylabel(), loc='center', color='#3f3f3f', fontsize='small', rotation='vertical', labelpad=4)
            ax[i, j].grid(axis="both", lw=0.5, ls=':')
            ax[i, j].spines[['right', 'top']].set_visible(False)
            ax[i, j].spines['left'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
            ax[i, j].spines['bottom'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
            ax[i, j].tick_params(axis='both', color='#3f3f3f', colors='#3f3f3f', labelsize='x-small')
