# -*- coding: utf-8 -*-
"""
Created on Mon May  8 14:41:09 2023

@author: gsun
"""
import numpy as np
import pandas as pd
from sklearn.metrics import matthews_corrcoef
from collections import defaultdict
from collections import Counter
from sklearn.metrics import roc_auc_score,accuracy_score
from sklearn.metrics import confusion_matrix
import math,time
from imblearn.combine import SMOTEENN
from imblearn.under_sampling import RandomUnderSampler
from imblearn.over_sampling import RandomOverSampler, SMOTE
from sklearn.feature_selection import SelectFromModel,SelectKBest,chi2,f_classif,mutual_info_classif,SelectFdr,SelectFpr
import os
from sklearn import preprocessing
from xgboost import XGBClassifier
from sklearn.metrics import f1_score,recall_score,precision_score
import random
from sklearn.preprocessing import StandardScaler
from sklearn.ensemble import RandomForestClassifier
import warnings
from tabulate import tabulate
import matplotlib.pyplot as plt

def f(x,y):
    if x == y: return 1
    else: return 0
def cal_metrics(y,pred,metric):
    if metric == 'acc':
        return accuracy_score(y,pred)
    if metric == 'recall f':
        return recall_score(y,pred,average=None)[0]
    if metric == 'recall p':
        p= recall_score(y,pred,average=None)
        if len(p) ==2:
            return p[1]
        else:
            return -1
    if metric == 'precision f':
        return precision_score(y,pred,average=None)[0]
    if metric == 'precision p':
        p = precision_score(y,pred,average=None)
        if len(p) ==2:
            return p[1]
        else:
            return -1
    if metric == 'f1 f':
        return f1_score(y,pred,average=None)[0]
    if metric == 'f1 p':
        return f1_score(y,pred,average=None)[1]
    
def compare(pr_df):    
    metrics = ['precision f','precision p','recall f','recall p']
    same = [[],[],[],[]]
    diff = [[],[],[],[]]
    overall = [[],[],[],[]]
    performance = [overall,same,diff]
    shapes = []
    for i,df in enumerate(pr_df):
        shapes.append(df.shape[0])
        pr_now = df.loc[df['last_label']!=df['now_label']].reset_index()
        pr_same = df.loc[df['last_label']==df['now_label']].reset_index()
        pr_status = [df,pr_same,pr_now]
        headers = ['Overall','|','PR_Same','|','PR!=NOW','|']
        print_data = []
        for metric in metrics:
            data_row =[metric]
            for j in range(0,len(pr_status)):
                df = pr_status[j]
                y = df['now_label']
                pred = df['predict']
                perf = cal_metrics(y,pred,metric)
                data_row+=[perf]
                data_row+="|"
                if metric == 'precision f':
                    performance[j][0].append(perf)
                elif metric == 'precision p':
                    performance[j][1].append(perf)
                elif metric == 'recall f':
                    performance[j][2].append(perf)
                elif metric == 'recall p':
                    performance[j][3].append(perf)
                  
            print_data.append(data_row)
        print(tabulate(print_data, headers=headers,floatfmt=".3f"))
    shapes.append(300) # hardcode project a
    shapes=list(map(lambda x:1000*x/max(shapes),shapes))
    return shapes,overall,same,diff
#%%
def add_project_a_sbs(pr_as_pred,overall,same,diff,pr_shapes):
    # the performance of SBS on project a is hardcoded since we cannot share the confidential dataset
    pr_as_pred_a = [[0.7092896174863388],[0.8209959623149394],[0.7092896174863388],[0.8209959623149394]]
    overall_a = [[0.19419237749546278],[0.5632432432432433],[0.11693989071038251],[0.7012113055181696]]
    same_a = [[0.1782178217821782],[0.5901759530791789],[0.1386748844375963],[0.6598360655737705]]
    diff_a = [[0.3695652173913043],[0.4876543209876543],[0.06390977443609022],[0.8909774436090225]]
    for i in range(0,4):
        pr_as_pred[i]+=pr_as_pred_a[i]
        overall[i]+=overall_a[i]
        same[i]+=same_a[i]
        diff[i]+=diff_a[i]
    pr_shapes[0] += [1724]
    pr_shapes[1] += [12003]
    return pr_as_pred,overall,same,diff,pr_shapes
def add_project_a_bf(pr_as_pred,overall,same,diff,pr_shapes):
    # the performance of BF on project a is hardcoded since we cannot share the confidential dataset
    pr_as_pred_a = [[0.7092896174863388],[0.8209959623149394],[0.7092896174863388],[0.8209959623149394]]
    overall_a = [[0.88], [0.6405910473707084], [0.09617486338797815], [0.9919246298788694]]
    same_a = [[1.0], [0.6850084222346996], [0.13559322033898305], [1.0]]
    diff_a = [[0.0], [0.48846153846153845], [0.0], [0.9548872180451128]]
    for i in range(0,4):
        pr_as_pred[i]+=pr_as_pred_a[i]
        overall[i]+=overall_a[i]
        same[i]+=same_a[i]
        diff[i]+=diff_a[i]
    pr_shapes[0] += [1724]
    pr_shapes[1] += [12003]
    return pr_as_pred,overall,same,diff,pr_shapes
    
#%%
def plot_scatter_plot(x,y,fsize,shapes,x_label = "",y_label = "",title =""):
    plt.rcParams.update({'font.size': 22})

    titles = ['Precision Fail','Precision Pass', 'Recall Fail', 'Recall Pass']
    x = np.array(x)
    y = np.array(y)
    colors = np.random.rand(len(x[0]))
    colors = list(range(5,26))
    colors = list(map(lambda x:x/20,colors))
    i,j = 0,0
    plt.figure()
    r_x = np.arange(0.0,1.1 , 0.1)
    fig, ax = plt.subplots(2, 2, sharex=True, sharey=True,figsize=(fsize,fsize))
    for i in range(0,2):
        for j in range(0,2):
            ax[i,j].plot(x[i*2+j][-1],y[i*2+j][-1],'yD',markersize=shapes[-1]/10)
            ax[i,j].set(xlim=(0, 1), ylim=(0, 1))
            ax[i,j].plot([0, 1], [0, 1], ls="--", c=".3",alpha=0.5)
            ax[i,j].plot([0, 0.95], [0.05, 1], ls="-.", c=".5",linewidth=1,alpha=0.5)
            ax[i,j].plot([0.05, 1], [0, 0.95], ls="-.", c=".5",linewidth=1,alpha=0.5)
            ax[i,j].scatter(x[i*2+j], y[i*2+j], s=shapes, edgecolors='gray', c=colors, alpha=0.5,cmap='gray')
            ax[i,j].set_title(titles[i*2+j],fontsize=22)
            ax[i,j].fill_between(r_x,r_x, color='grey', alpha=0.1,
                     interpolate=True)
            ax[i,j].set_xlabel(x_label,fontsize=22)
            ax[i,j].set_ylabel(y_label,fontsize=22)
    fig.suptitle(title, fontsize=16)
    
    fig.tight_layout()
    plt.show()
#%%
def bar_projects(pr_shapes):
    combined = list(zip(pr_shapes[1],pr_shapes[0]))
    sorted_combined = sorted(combined, key=lambda x: x[0])
    x_sorted, y_sorted = zip(*sorted_combined)
    percentages = [100*a/b for a,b in zip(y_sorted, x_sorted)]
    
    textures = ['x' if i ==20 else '//' for i in range(len(percentages))]
    plt.figure(figsize=(15,8))
    # Define the data
    plt.grid(which='both',axis='y',linewidth=0.5)
    colors=['#BFBFBF']*20+['#FFFF14']
    plt.bar(range(len(percentages))[:-1], percentages[:-1],edgecolor='gray',alpha=0.8,color=colors[:-1],hatch=textures[:-1],label="Open-Source")
    plt.bar(range(len(percentages))[-1], percentages[-1],edgecolor='gray',alpha=0.8,color=colors[-1],hatch=textures[-1],label="Project A")
    tick_labels = list(range(0,21))
    plt.xticks(range(len(percentages)), tick_labels)
    plt.xlabel('Projects Ordered by Size', fontsize = 32)
    plt.ylabel('Percentage (%)', fontsize = 32)
    plt.legend()
    plt.show()