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, plot_heatmap
from src.utils.signature import iterative_refinement, plot_group_sizes, save_group_lists_fasta, compute_jaccard_matrix

def get_signature(
        data_save_path, 
        data_mat, 
        loci_list, 
        cap_list, 
        direction, 
        cluster_loci = True, 
        cluster_cap = True,
        rerun_cluster = False
    ):
    
    # Load config
    with open('config.yaml', 'r') as f:
        config = yaml.safe_load(f)

    # NMF parameters from config
    random_state = config['signature_config']['seed']
    n_components = config['signature_config']['nmf']['n_components']
    init = config['signature_config']['nmf']['init']
    max_iter = config['signature_config']['nmf']['max_iter']
    n_components_range = range(config['signature_config']['nmf']['min_components'], config['signature_config']['nmf']['max_components'] + 1)

    W_npy_path = os.path.join(data_save_path, 'W_loci_by_k.npy')
    H_npy_path = os.path.join(data_save_path, 'H_k_by_cap.npy')
    W_df_path = os.path.join(data_save_path, 'W_loci_by_k.csv')
    H_df_path = os.path.join(data_save_path, 'H_k_by_cap.csv')

    if os.path.exists(W_npy_path) and os.path.exists(H_npy_path) and not rerun_cluster:
        W = np.load(W_npy_path)
        H = np.load(H_npy_path)
        W_df = pd.read_csv(W_df_path, index_col=0)
        H_df = pd.read_csv(H_df_path, index_col=0)
    else:
        W, H = run_nmf(data=data_mat, n_components=n_components, init=init, random_state=random_state, max_iter=max_iter)

        # Save W and H matrices
        np.save(W_npy_path, W)
        np.save(H_npy_path, H)
        W_df = pd.DataFrame(W, index=loci_list)
        W_df.to_csv(W_df_path, index=True)
        H_df = pd.DataFrame(H, columns=cap_list)
        H_df.to_csv(H_df_path, index=True)

    # Group CAPs and loci
    loci_group = group_loci(W_df, data_mat)
    cap_group = group_cap(H_df, data_mat)
    loci_group['group'] = loci_group['group'].astype(int)
    cap_group['group'] = cap_group['group'].astype(int)

    # Refine ranks of loci and CAPs within each group
    loci_rank, cap_rank = iterative_refinement(data_mat, loci_group['group'].values, cap_group['group'].values, loci_list, cap_list, top_group_percent = 0.5, top_item_percent = 0.5)
    loci_group['rank'] = loci_rank
    cap_group['rank'] = cap_rank

    # Build loci_group and cap_group dataframes
    loci_assignment = loci_group['group'].values.astype(int)
    cap_assignment = cap_group['group'].values.astype(int)
    loci_group = pd.DataFrame({'loci': loci_list, 'group': loci_assignment, 'avg_binding': np.mean(data_mat, axis=1), 'rank': loci_rank})
    cap_group = pd.DataFrame({'cap': cap_list, 'group': cap_assignment, 'avg_binding': np.mean(data_mat, axis=0), 'rank': cap_rank})

    # Track original positions
    loci_group['original_index'] = range(len(loci_group))
    cap_group['original_index'] = range(len(cap_group))

    # 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])
    
    # Build the plot matrix
    plot_mat = data_mat[np.ix_(loci_group['original_index'], cap_group['original_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.drop(columns=['original_index']).to_csv(os.path.join(data_save_path, 'loci_group.csv'))
    cap_group.drop(columns=['original_index']).to_csv(os.path.join(data_save_path, 'cap_group.csv'))
    # Save plot_mat
    plot_mat_path = os.path.join(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 = 'Average max CAP binding at loci'

    loci_k = n_components
    cap_k = n_components

    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 run_nmf(data, n_components, init='random', random_state=9, max_iter=200):
    print("Running vanilla NMF with n_components =", n_components)
    from sklearn.decomposition import NMF
    model = NMF(n_components=n_components, init=init, random_state=random_state, max_iter=max_iter)
    W = model.fit_transform(data)
    H = model.components_
    return W, H

def group_loci(W_df, data_mat):
    return group_entities(W_df, data_mat, entity_name='loci')

def group_cap(H_df, data_mat):
    return group_entities(H_df.T, data_mat.T, entity_name='cap')

def group_entities(W_df, data_mat, entity_name='loci'):
    # W_df: (n_entities, n_components)
    # For each entity, find component with highest value and assign entity to that component
    entity_assignments = W_df.idxmax(axis=1)  # Series: entity_name -> group
    # Create ordered list that matches W_df.index order
    entity_list = W_df.index.tolist() 
    group_assignments = [entity_assignments[entity] for entity in entity_list]
    # Return DataFrame in original order
    entity_group = pd.DataFrame({
        entity_name: entity_list,
        'group': group_assignments
    })
    return entity_group

def compute_group_similarity(plot_dict1, plot_dict2, celltypes_list, data_save_path):
    loci_group_1 = plot_dict1['loci_group']
    loci_group_2 = plot_dict2['loci_group']
    cap_group_1 = plot_dict1['cap_group']
    cap_group_2 = plot_dict2['cap_group']
    celltype_1, celltype_2 = celltypes_list

    # Compute similarity matrices for loci and CAP groups
    loci_sim, loci_groups_1, loci_groups_2 = compute_jaccard_matrix(loci_group_1, loci_group_2, 'loci', 'group')
    cap_sim, cap_groups_1, cap_groups_2 = compute_jaccard_matrix(cap_group_1, cap_group_2, 'cap', 'group')
    
    # Check if shapes match for combined similarity
    if loci_sim.shape == cap_sim.shape:
        combined_sim = (loci_sim + cap_sim) / 2
        combined_groups_1 = loci_groups_1  # Use loci groups for labeling
        combined_groups_2 = loci_groups_2
    else:
        print(f"Warning: Loci similarity shape {loci_sim.shape} != CAP similarity shape {cap_sim.shape}")
        print("Computing combined similarity using intersection of groups...")
        
        # Find common groups for both datasets
        common_groups_1 = sorted(set(loci_groups_1) & set(cap_groups_1))
        common_groups_2 = sorted(set(loci_groups_2) & set(cap_groups_2))
        
        # Extract submatrices for common groups
        loci_indices_1 = [loci_groups_1.index(g) for g in common_groups_1 if g in loci_groups_1]
        loci_indices_2 = [loci_groups_2.index(g) for g in common_groups_2 if g in loci_groups_2]
        cap_indices_1 = [cap_groups_1.index(g) for g in common_groups_1 if g in cap_groups_1]
        cap_indices_2 = [cap_groups_2.index(g) for g in common_groups_2 if g in cap_groups_2]
        
        if len(loci_indices_1) > 0 and len(loci_indices_2) > 0 and len(cap_indices_1) > 0 and len(cap_indices_2) > 0:
            loci_sub = loci_sim[np.ix_(loci_indices_1, loci_indices_2)]
            cap_sub = cap_sim[np.ix_(cap_indices_1, cap_indices_2)]
            combined_sim = (loci_sub + cap_sub) / 2
            combined_groups_1 = common_groups_1
            combined_groups_2 = common_groups_2
        else:
            print("Warning: No common groups found. Skipping combined similarity.")
            combined_sim = None
            combined_groups_1 = []
            combined_groups_2 = []

    # Create and save dataframes
    similarity_data = [
        (loci_sim, 'loci', loci_groups_1, loci_groups_2),
        (cap_sim, 'cap', cap_groups_1, cap_groups_2)
    ]
    
    if combined_sim is not None:
        similarity_data.append((combined_sim, 'combined', combined_groups_1, combined_groups_2))
    
    for sim_matrix, name, groups_1, groups_2 in similarity_data:
        sim_df = pd.DataFrame(
            sim_matrix, 
            index=[f'{celltype_1}_group_{g}' for g in groups_1], 
            columns=[f'{celltype_2}_group_{g}' for g in groups_2]
        )
        sim_df.to_csv(os.path.join(data_save_path, f'{name}_jaccard_similarity_matrix.csv'))

    # Print best matches
    if combined_sim is not None:
        print("\nBest matches (Group1 -> Group2):")
        for i, g1 in enumerate(combined_groups_1):
            best_match_idx = np.argmax(combined_sim[i, :])
            best_score = combined_sim[i, best_match_idx]
            best_g2 = combined_groups_2[best_match_idx]
            print(f"Group {g1} -> Group {best_g2} (score: {best_score:.3f})")
    else:
        print("\nSkipping best matches due to shape mismatch.")

    # Plot heatmaps
    heatmap_configs = [
        (loci_sim, 'Loci group Similarity Matrix', 'Loci group Jaccard Similarity', 'loci_similarity_heatmap'),
        (cap_sim, 'CAP group Similarity Matrix', 'CAP group Jaccard Similarity', 'cap_similarity_heatmap')
    ]
    
    # Only add combined heatmap if combined_sim is not None
    if combined_sim is not None:
        heatmap_configs.append((combined_sim, 'Combined group Similarity Matrix', 'Combined group Jaccard Similarity', 'combined_similarity_heatmap'))
    
    for sim_matrix, title, cbar_label, filename in heatmap_configs:
        plot_heatmap(
            plot_mat=sim_matrix,
            plot_title=title,
            cbar_label=cbar_label,
            color_map='viridis',
            x_name=f'{celltype_2} Groups',
            y_name=f'{celltype_1} Groups',
            filename=filename,
            data_save_path=data_save_path
        )

if __name__ == "__main__":
    main()