import os
import numpy as np
import pandas as pd
from tqdm import tqdm
import matplotlib.pyplot as plt
import matplotlib as mpl
import yaml

from src.utils.vis import figure, subplots
from src.utils.signature import iterative_refinement, plot_group_sizes, save_group_lists_fasta

def get_signature(
        data_save_path, 
        diff_mat, 
        loci_list, 
        cap_list, 
        direction, 
        normalize = False, 
        cluster_loci = True, 
        cluster_cap = True,
        rerun_cluster = False
    ):

    # Load config
    with open('config.yaml', 'r') as f:
        config = yaml.safe_load(f)

    # Silhouette analysis
    min_k = config['signature_config']['kmeans']['min_k']
    max_k = config['signature_config']['kmeans']['max_k']
    stride = 1
    loci_mat = diff_mat
    if cluster_loci:
        loci_k, loci_assignment = silhouette_w_plot(data_save_path, loci_mat, min_k, max_k, stride, 'loci', normalize, rerun_cluster)
    else:
        loci_k = 1
        loci_assignment = np.zeros(len(loci_list))
    cap_mat = diff_mat.T
    if cluster_cap:
        cap_k, cap_assignment = silhouette_w_plot(data_save_path, cap_mat, min_k, max_k, stride, 'cap', normalize, rerun_cluster)
    else:
        cap_k = 1
        cap_assignment = np.zeros(len(cap_list))

    print(f'Running iterative refinement')
    loci_rank, cap_rank = iterative_refinement(diff_mat, loci_assignment, cap_assignment, loci_list, cap_list, top_group_percent = 0.5, top_item_percent = 0.5)

    # Group cap and loci according to the optimal k
    loci_group = pd.DataFrame({'loci': loci_list, 'group': loci_assignment, 'avg_diff': np.mean(diff_mat, axis=1), 'rank': loci_rank})
    cap_group = pd.DataFrame({'cap': cap_list, 'group': cap_assignment, 'avg_diff': np.mean(diff_mat, axis=0), 'rank': cap_rank})

    # Re-arrange the loci and caps
    loci_group = loci_group.sort_values(['group', 'rank'], ascending=[True, True])
    cap_group = cap_group.sort_values(['group', 'rank'], ascending=[True, True])
    # Re-arrange the diff matrix
    plot_mat = diff_mat[loci_group.index, :]
    plot_mat = plot_mat[:, cap_group.index]
    loci_group = loci_group.reset_index(drop=True)
    cap_group = cap_group.reset_index(drop=True)

    # Save loci_group and cap_group
    loci_group.to_csv(f'{data_save_path}/loci_group.csv')
    cap_group.to_csv(f'{data_save_path}/cap_group.csv')
    # Save plot_mat
    plot_mat_path = f'{data_save_path}/plot_matrix.npy'
    np.save(plot_mat_path, plot_mat)

    # Plot group sizes
    # plot_group_sizes(loci_group, cap_group, data_save_path)

    # Save group lists in FASTA-like format
    save_group_lists_fasta(loci_group, cap_group, data_save_path)

    # For plotting heatmap
    x_name = 'CAP'
    y_name = 'Loci'
    x_annotation = cap_group['cap']
    y_annotation = loci_group['loci'].values
    title_str = 'Differential CAP binding at loci'

    # Save dataframe
    plot_df = pd.DataFrame(plot_mat, index = loci_group['loci'], columns = cap_group['cap'])
    plot_df.to_csv(f'{data_save_path}/plot_matrix.csv')

    plot_dict = {
        'plot_mat': plot_mat,
        'loci_k': loci_k,
        'cap_k': cap_k,
        'loci_group': loci_group,
        'cap_group': cap_group,
        'data_save_path': data_save_path,
        'x_name': x_name,
        'y_name': y_name,
        'x_annotation': x_annotation,
        'y_annotation': y_annotation,
        'title_str': title_str,
        'direction': direction
    }

    return plot_dict

def silhouette_w_plot(data_save_path, data, min_k, max_k, stride, data_type, normalize = False, rerun_cluster = False):
    optimal_k_file = f'{data_save_path}/{data_type}_optimal_k.txt'
    optimal_assignment_file = f'{data_save_path}/{data_type}_optimal_assignment.npy'
    assignment_file = f'{data_save_path}/{data_type}_assignments.npy'
    score_file = f'{data_save_path}/{data_type}_scores.npy'
    if os.path.exists(optimal_k_file) and not rerun_cluster:
        optimal_k = int(open(optimal_k_file, 'r').read())
        optimal_assignment = np.load(optimal_assignment_file)
        assignments = np.load(assignment_file)
        scores = np.load(score_file)
    else:
        scores, assignments, optimal_k, optimal_assignment = silhouette_analysis(data, min_k, max_k, stride, normalize)
        # Save assignments
        np.save(assignment_file, assignments)
        np.save(score_file, scores)
        with open(optimal_k_file, 'w') as f:
            f.write(str(optimal_k))
        np.save(optimal_assignment_file, optimal_assignment)
    # Plot silhouette scores
    import matplotlib.pyplot as plt
    plt.plot(range(min_k, max_k, stride), scores)
    plt.xlabel('Number of clusters')
    plt.ylabel('Silhouette score')
    plt.title('Silhouette score vs number of clusters')
    plt.savefig(f'{data_save_path}/silhouette_score_{data_type}.png')
    plt.close()
    return optimal_k, optimal_assignment

def silhouette_analysis(data, min_k=2, max_k=20, stride=1, normalize = False):
    from sklearn.metrics import silhouette_score
    from sklearn.cluster import KMeans
    # Silhouette analysis
    silhouette_score_values = []
    assignments = []
    optimal_k = 0
    optimal_assignment = None
    for k in tqdm(range(min_k, max_k, stride)):
        print(f'Running k={k}')
        if normalize:
            # Standardize the data to unit length
            length = np.sqrt((data**2).sum(axis=1))[:,None]
            cluster_data = data / length
            # Set nan to 0
            cluster_data[np.isnan(cluster_data)] = 0
        else:
            cluster_data = data

        kmeans = KMeans(n_clusters=k, random_state=0, n_init=10).fit(cluster_data)
        labels = kmeans.labels_
        # Reorder labels based on the mean value of top 20% of each cluster
        mean_values = np.zeros(k)
        for i in range(k):
            selected_data = data[labels == i]
            mean_values[i] = np.mean(selected_data[selected_data >= np.percentile(selected_data, 80)])
        order = np.argsort(mean_values)[::-1]
        new_labels = np.zeros(len(labels))
        for i in range(k):
            new_labels[labels == order[i]] = i
        labels = new_labels

        score = silhouette_score(cluster_data, labels)
        silhouette_score_values.append(score)
        assignments.append(labels)
        if score == max(silhouette_score_values):
            optimal_k = k
            optimal_assignment = labels
    print(f'Optimal k: {optimal_k}')
    assignments = np.array(assignments)
    return silhouette_score_values, assignments, optimal_k, optimal_assignment
