import yaml
import numpy as np
import os
import pandas as pd
from tqdm import tqdm
import argparse
import re

def main():
    parser = argparse.ArgumentParser() #hua added for differet configs
    parser.add_argument(
        "-c", "--config",
        help="Path(s) to YAML config."
    )
    args = parser.parse_args()

    # resolution order: CLI --config (in order) > env APP_CONFIG > ./config.yaml
    config_path = args.config
    config_name=re.sub(r'^(?:config_)?|\.ya?ml$', '',config_path)
    # Load config
    with open(config_path, 'r') as f:
        config = yaml.safe_load(f)

    # Get differential vs conserved gene lists
    # RNA-seq mode
    if config['loci_source_config']['mode'] == 'RNA_seq':
        rna_config = config['loci_source_config']['rna']
        diff_exp_df = read_rna_data(rna_config)
        diff_exp_df = add_loci(diff_exp_df, config)
        diff_exp_df = get_promoter_accessibility(diff_exp_df, config)
        gene_type_dict = split_gene_type(diff_exp_df, rna_config)
        # Filter by promoter accessibility
        gene_type_dict = filter_by_rna_promoter_accessibility(gene_type_dict, config)
    # ATAC-seq mode
    elif config['loci_source_config']['mode'] == 'ATAC_seq':
        atac_seq_specific_config = config['loci_source_config']['atac_seq_specific_config']
        if atac_seq_specific_config['mode'] == 'Promoter':
            atac_config = config['loci_source_config']['atac']
            diff_acc_df = read_filter_gff_data(config)
            diff_acc_df = get_promoter_accessibility(diff_acc_df, config, lazy=False)
            gene_type_dict = atac_split_filter_gene_type(diff_acc_df, config)
        elif atac_seq_specific_config['mode'] == 'All_peaks':
            raise NotImplementedError('All_peaks mode not implemented')
        else:
            raise ValueError(f'Invalid mode: {config["loci_source_config"]["mode"]}')
    else:
        raise ValueError(f'Invalid mode: {config["loci_source_config"]["mode"]}')

    # Save gene list and merged bed file for inference
    save_gene_list(gene_type_dict, config,config_name) #hua changed

    # Copy merged_loci.bed to inputs folder
    copy_merged_loci_to_inputs(config,config_name) #hua changed

def read_filter_gff_data(config):
    gff_path = os.path.join(config['input_resource']['root'], config['input_resource']['sequence'], config['preprocessing_data_config']['gff_file'])
    print(f'Reading gff file from {gff_path}')
    gff_df = read_genes(gff_path)

    # Remove genes on chrX, chrY, and chrM
    gff_df = gff_df[~gff_df['seqname'].isin(['chrX', 'chrY', 'chrM'])]

    # Keep a few transcript_types
    transcript_types = ['protein_coding', 'lincRNA', 'lncRNA', 'antisense', 'processed_transcript',
                        'retained_intron', 'sense_intronic', 'sense_overlapping',
                        'macro_lncRNA', 'bidirectional_promoter_lncRNA', 'non_coding',
                        'miRNA', 'snRNA', 'snoRNA', 'rRNA']
    gff_df = gff_df[gff_df['transcript_type'].isin(transcript_types)]

    na_start = gff_df['start'].isna()
    na_end = gff_df['end'].isna()
    na_gene = gff_df[na_start | na_end]
    if len(na_gene) > 0:
        print(f'Warning: {len(na_gene)} genes with NA start or end')
        print(na_gene['gene_name'].to_string(index=False))

    gff_df = gff_df[~(na_start | na_end)]

    # Add TSS column
    gff_df['tss'] = np.where(gff_df['strand'] == '+', gff_df['start'], gff_df['end']).astype(int)

    tss_radius = config['loci_source_config']['atac_seq_specific_config']['promoter']['tss_radius']
    gff_df = add_inference_loci(gff_df, tss_radius, config)
    return gff_df

def save_gene_list(gene_type_dict, config,config_name): #hua added
    save_root = config["preprocessing_data_config"]["output"]["path"] + '/'+config_name+'/gene_list'
    os.makedirs(save_root, exist_ok=True)
    # Save gene list and merged bed file for inference
    for gene_list_type, gene_list_df in gene_type_dict.items():
        # Save gene list
        gene_list_df.to_csv(f'{save_root}/{gene_list_type}_gene_list.csv', index=False)
        # Save bed file
        bed_df = gene_list_df[['seqname', 'tss_inference_start', 'tss_inference_end', 'transcript_name']].copy()
        bed_df['tss_inference_start'] = bed_df['tss_inference_start'].astype(int)
        bed_df['tss_inference_end'] = bed_df['tss_inference_end'].astype(int)
        bed_df['score'] = 0
        bed_df['strand'] = gene_list_df['strand']
        bed_df.to_csv(f'{save_root}/{gene_list_type}_gene_list.bed', sep='\t', index=False, header=False)

    # Merge loci for inference
    merge_loci(gene_type_dict, save_root)

def merge_loci(gene_type_dict, save_root):
    # Merge overlapping loci from all gene lists
    merged_df = pd.concat(gene_type_dict.values())
    
    # Sort by chromosome and start position
    merged_df = merged_df.sort_values(['seqname', 'tss_inference_start'])
    
    # Merge overlapping regions
    merged_regions = []
    current_chrom = None
    current_start = None 
    current_end = None
    current_names = []
    current_strands = []
    
    for _, row in merged_df.iterrows():
        if current_chrom != row['seqname'] or row['tss_inference_start'] > current_end:
            # Save previous merged region
            if current_chrom is not None:
                merged_regions.append({
                    'seqname': current_chrom,
                    'start': current_start,
                    'end': current_end,
                    'transcript_name': '|'.join(current_names),
                    'score': 0,
                    'strand': '|'.join(current_strands)
                })
            # Start new region
            current_chrom = row['seqname']
            current_start = row['tss_inference_start']
            current_end = row['tss_inference_end']
            current_names = [row['transcript_name']]
            current_strands = [row['strand']]
        else:
            # Extend current region
            current_end = max(current_end, row['tss_inference_end'])
            current_names.append(row['transcript_name'])
            current_strands.append(row['strand'])
    
    # Add final region
    if current_chrom is not None:
        merged_regions.append({
            'seqname': current_chrom,
            'start': current_start,
            'end': current_end,
            'transcript_name': '|'.join(current_names),
            'score': 0,
            'strand': '|'.join(current_strands)
        })
    
    # Convert to dataframe and save
    merged_df = pd.DataFrame(merged_regions)
    merged_df.to_csv(f'{save_root}/merged_loci.csv', index=False)
    merged_df.to_csv(f'{save_root}/merged_loci.bed', sep='\t', index=False, header=False)

def get_filter_location(diff_df, top_n):
    # Make sure not too many loci and from the same gene
    g_idx = 0
    o_idx = 0
    gene_set = set()
    if 'gene_name' in diff_df.columns:
        gene_list = diff_df['gene_name'].values
        for g_idx, gene in enumerate(gene_list):
            gene_set.add(gene)
            if len(gene_set) == top_n:
                break
    # Make sure not too many loci and from the same region
    overlap_threshold = 0.5
    region_df = diff_df[['seqname', 'tss_inference_start', 'tss_inference_end']]
    unique_region_length = 0
    for o_idx, row in region_df.iterrows():
        if o_idx == 0:
            continue
        prior_df = region_df.iloc[:o_idx]
        same_chrom_df = prior_df[prior_df['seqname'] == row['seqname']]
        # Calculate if there are overlaps over 50%
        starts = same_chrom_df['tss_inference_start'].values
        ends = same_chrom_df['tss_inference_end'].values
        current_start = row['tss_inference_start']
        current_end = row['tss_inference_end']
        overlaps = np.maximum(0, np.minimum(current_end, ends) - np.maximum(current_start, starts))
        half_length = (current_end - current_start) / 2
        if np.sum(overlaps) / half_length > overlap_threshold:
            continue
        unique_region_length += 1
        if unique_region_length > top_n:
            break
    return max(g_idx, o_idx), max(len(gene_set), unique_region_length)

def filter_by_rna_promoter_accessibility(gene_list_dict, config):

    for gene_list_type, diff_df in gene_list_dict.items():
        # Remove genes with small accessibility 
        accessibility_threshold = config['loci_source_config']['atac']['active_locus_threshold']
        cross_condition_max_accessibility = diff_df[['experimental_max_accessibility', 'control_max_accessibility']].max(axis=1)
        select_idx = cross_condition_max_accessibility > accessibility_threshold
        gene_list_dict[gene_list_type] = diff_df[select_idx]

    upregulated_df, downregulated_df, conserved_df = gene_list_dict['upregulated'], gene_list_dict['downregulated'], gene_list_dict['conserved']

    acc_log2fc_threshold = config['loci_source_config']['atac']['log2fc_threshold']

    num_diff = config['loci_source_config']['rna']['top_diff_genes']
    upregulated_df = add_log2fc(upregulated_df)
    upregulated_df = upregulated_df[upregulated_df['accessibility_log2fc'] > acc_log2fc_threshold]
    upregulated_df = upregulated_df.sort_values(by='abs_rna_acc_log2fc_sum', ascending=False).reset_index(drop=True)
    filter_loc, num_unique = get_filter_location(upregulated_df, num_diff) # Get the location of the top num_diff unique genes
    upregulated_df = upregulated_df.head(filter_loc).reset_index(drop=True)
    if num_unique < num_diff:
        print(f'Warning: Only {num_unique} unique upregulated genes found after filtering by promoter accessibility')

    downregulated_df = add_log2fc(downregulated_df)
    downregulated_df = downregulated_df[downregulated_df['accessibility_log2fc'] < -acc_log2fc_threshold]
    downregulated_df = downregulated_df.sort_values(by='abs_rna_acc_log2fc_sum', ascending=False).reset_index(drop=True)
    filter_loc, num_unique = get_filter_location(downregulated_df, num_diff) # Get the location of the top num_diff unique genes
    downregulated_df = downregulated_df.head(filter_loc).reset_index(drop=True)
    if num_unique < num_diff:
        print(f'Warning: Only {num_unique} unique downregulated genes found after filtering by promoter accessibility')

    num_conserved = config['loci_source_config']['rna']['top_conserved_genes']
    conserved_df = add_log2fc(conserved_df)
    conserved_df = conserved_df.sort_values(by='abs_rna_acc_log2fc_sum', ascending=True).reset_index(drop=True)
    filter_loc, num_unique = get_filter_location(conserved_df, num_conserved) # Get the location of the top num_conserved unique genes
    conserved_df = conserved_df.head(filter_loc).reset_index(drop=True)
    if num_unique < num_conserved:
        print(f'Warning: Only {num_unique} unique conserved genes found after filtering by promoter accessibility')

    gene_list_dict = {'upregulated': upregulated_df, 'downregulated': downregulated_df, 'conserved': conserved_df}

    return gene_list_dict
    
def add_log2fc(diff_df):
    diff_df['accessibility_log2fc'] = np.log2(diff_df['experimental_max_accessibility'] / diff_df['control_max_accessibility'])
    diff_df['rna_acc_log2fc_sum'] = diff_df['log2fc'] + diff_df['accessibility_log2fc']
    diff_df['abs_rna_acc_log2fc_sum'] = diff_df['rna_acc_log2fc_sum'].abs()
    return diff_df

def get_promoter_accessibility(diff_df, config, lazy=True):
    import zarr
    # Load experimental and control zarr files
    zarr_path_dict = {condition: [os.path.join(config['input_resource']['root'], config['input_resource']['atac'], f"{rep}.zarr") for rep in config['loci_source_config']['atac']['data_path'][condition].split(',')] for condition in ['experimental', 'control']} #hua changed
    if lazy:
        zarr_dict = {condition: [zarr.open(zarr_pa, mode='r') for zarr_pa in zarr_path] for condition, zarr_path in zarr_path_dict.items()} #hua changed
    else:
        print('Loading zarr files to memory')
        zarr_dict = {condition: [load_zarr_to_memory(zarr_pa) for zarr_pa in zarr_path] for condition, zarr_path in zarr_path_dict.items()} #hua changed

    for zarr_condition, zarr_data in zarr_dict.items():
        # Filter by promoter accessibility
        max_accessibility_list = []
        mean_accessibility_list = []
        # Load promoter accessibility data by chromosome
        print(f'Loading {zarr_condition} promoter accessibility data')
        import tqdm
        for row_idx, row in tqdm.tqdm(diff_df.iterrows(), total=len(diff_df)):
            chrom = row['seqname']
            start = row['tss_inference_start']
            end = row['tss_inference_end']
            max_accessibility_list.append(np.max(np.stack([zarr_da['chrs'][chrom][start:end] for zarr_da in zarr_data],axis=0))) #hua changed
            mean_accessibility_list.append(np.mean(np.stack([zarr_da['chrs'][chrom][start:end] for zarr_da in zarr_data],axis=0))) #hua changed
        diff_df[f'{zarr_condition}_max_accessibility'] = max_accessibility_list
        diff_df[f'{zarr_condition}_mean_accessibility'] = mean_accessibility_list
    return diff_df
def load_zarr_to_memory(zarr_path):
    import zarr
    zarr_data = zarr.open(zarr_path, mode='r')
    zarr_data_dict = {'chrs': {}}
    print(f'Loading to memory: {zarr_path}')
    for chrom in tqdm(zarr_data['chrs'].keys()):
        zarr_data_dict['chrs'][chrom] = zarr_data['chrs'][chrom][:]
    return zarr_data_dict

def read_rna_data(rna_config):
    diff_exp_df = pd.read_csv(rna_config['log2fc_csv'])
    assert 'gene_name' in diff_exp_df.columns and 'experimental_exp' in diff_exp_df.columns and 'control_exp' in diff_exp_df.columns and 'log2fc' in diff_exp_df.columns
    return diff_exp_df

def add_loci(diff_exp_df, config):
    # Add loci to the gene list
    gff_path = os.path.join(config['input_resource']['root'], config['input_resource']['sequence'], config['preprocessing_data_config']['gff_file'])
    diff_exp_df = get_gene_loci(diff_exp_df, gff_path)
    tss_radius = config['loci_source_config']['rna']['tss_radius']
    diff_exp_df = add_inference_loci(diff_exp_df, tss_radius, config)
    return diff_exp_df

def add_inference_loci(diff_exp_df, tss_radius, config):
    diff_exp_df['tss_inference_start'] = diff_exp_df['tss'] - tss_radius
    diff_exp_df['tss_inference_end'] = diff_exp_df['tss'] + tss_radius

    # Clip to chromosome length
    chr_sizes = load_chr_sizes(os.path.join(config['input_resource']['root'], config['input_resource']['sequence'], config['preprocessing_data_config']['chr_sizes_file']))

    diff_exp_df['tss_inference_start'] = np.maximum(diff_exp_df['tss_inference_start'], 0)
    diff_exp_df['tss_inference_end'] = np.minimum(diff_exp_df['tss_inference_end'], diff_exp_df['seqname'].map(chr_sizes))

    return diff_exp_df

def atac_split_filter_gene_type(diff_df, config):
    # Add log2fc
    diff_df['accessibility_log2fc'] = np.log2(diff_df['experimental_max_accessibility'] / diff_df['control_max_accessibility'])

    # Filter by accessiblity
    accessibility_threshold = config['loci_source_config']['atac']['active_locus_threshold']
    cross_condition_max_accessibility = diff_df[['experimental_max_accessibility', 'control_max_accessibility']].max(axis=1)
    select_idx = cross_condition_max_accessibility > accessibility_threshold
    diff_df = diff_df[select_idx]

    # Split by absolute log2fc
    diff_df['abs_accessibility_log2fc'] = diff_df['accessibility_log2fc'].abs()
    diff_df = diff_df.sort_values(by='abs_accessibility_log2fc', ascending=False)

    # Get conserved genes
    atac_config = config['loci_source_config']['atac']
    num_conserved = atac_config['top_conserved_genes']
    conserved_df = diff_df.copy().sort_values(by='abs_accessibility_log2fc', ascending=True).reset_index(drop=True)
    filter_loc, num_unique = get_filter_location(conserved_df, num_conserved) # Get the location of the top num_conserved unique genes
    conserved_df = conserved_df.head(filter_loc).reset_index(drop=True)
    if num_unique < num_conserved:
        print(f'Warning: Only {num_unique} unique conserved genes found after filtering by promoter accessibility')

    # Get upregulated and downregulated genes
    atac_threshold = atac_config['log2fc_threshold']
    num_diff = atac_config['top_diff_genes']
    upregulated_df = diff_df[diff_df['accessibility_log2fc'] > atac_threshold].copy().sort_values(by='accessibility_log2fc', ascending=False).reset_index(drop=True)
    filter_loc, num_unique = get_filter_location(upregulated_df, num_diff) # Get the location of the top num_diff unique genes
    upregulated_df = upregulated_df.head(filter_loc).reset_index(drop=True)
    if num_unique < num_diff:
        print(f'Warning: Only {num_unique} unique upregulated genes found after filtering by promoter accessibility')

    downregulated_df = diff_df[diff_df['accessibility_log2fc'] < -atac_threshold].copy().sort_values(by='accessibility_log2fc', ascending=True).reset_index(drop=True)
    filter_loc, num_unique = get_filter_location(downregulated_df, num_diff) # Get the location of the top num_diff unique genes
    downregulated_df = downregulated_df.head(filter_loc).reset_index(drop=True)
    if num_unique < num_diff:
        print(f'Warning: Only {num_unique} unique downregulated genes found after filtering by promoter accessibility')

    return {'upregulated': upregulated_df, 'downregulated': downregulated_df, 'conserved': conserved_df}

def split_gene_type(diff_exp_df, rna_config):
    # Split by absolute log2fc
    diff_exp_df['abs_log2fc'] = diff_exp_df['log2fc'].abs()
    diff_exp_df = diff_exp_df.sort_values(by='abs_log2fc', ascending=False)

    # Get conserved genes
    num_conserved = rna_config['top_conserved_genes'] * 4 # Leave room for atac-seq filtering
    conserved_genes = diff_exp_df.sort_values(by='abs_log2fc', ascending=True).head(num_conserved).reset_index(drop=True)

    # Get upregulated and downregulated genes
    rna_threshold = rna_config['log2fc_threshold']
    num_diff = rna_config['top_diff_genes'] * 2 # Leave room for atac-seq filtering
    upregulated_genes = diff_exp_df[diff_exp_df['log2fc'] > rna_threshold].sort_values(by='log2fc', ascending=False).head(num_diff).reset_index(drop=True)
    downregulated_genes = diff_exp_df[diff_exp_df['log2fc'] < -rna_threshold].sort_values(by='log2fc', ascending=True).head(num_diff).reset_index(drop=True)

    return {'upregulated': upregulated_genes, 'downregulated': downregulated_genes, 'conserved': conserved_genes}

def get_gene_loci(gene_df, gff_path):
    # Find location of genes in the genome with a annotation file
    print(f'Reading gff file from {gff_path}')
    gff_df = read_genes(gff_path)
    # Find the gene names in the annotation file
    gene_names = gene_df['gene_name'].values
    selected_gene_gff_df = []
    print(f'Finding gene loci for {len(gene_names)} genes')
    for gene_name in gene_names:
        gene_gff_df = gff_df[gff_df['gene_name'] == gene_name]
        selected_gene_gff_df.append(gene_gff_df)
    selected_gene_gff_df = pd.concat(selected_gene_gff_df).reset_index(drop=True)
    # Join the split_df with the gff_df
    merged_df = gene_df.merge(selected_gene_gff_df, on='gene_name', how='left')

    # Only keep gene on chr1-22 and X
    merged_df = merged_df[merged_df['seqname'].isin([f'chr{i}' for i in range(1, 23)] + ['chrX'])]

    # Remove genes with NA start or end and print warning
    na_start = merged_df['start'].isna()
    na_end = merged_df['end'].isna()
    na_gene = merged_df[na_start | na_end]
    if len(na_gene) > 0:
        print(f'Warning: {len(na_gene)} genes with NA start or end')
        print(na_gene['gene_name'].to_string(index=False))

    merged_df = merged_df[~(na_start | na_end)]

    # Add TSS column
    merged_df['tss'] = np.where(merged_df['strand'] == '+', merged_df['start'], merged_df['end']).astype(int)

    return merged_df

def read_genes(gff_path):
    names = ['seqname', 'source', 'feature', 'start', 'end', 'score', 'strand', 'frame', 'attribute']
    genes = pd.read_csv(gff_path, sep='\t', comment='#', names=names)
    gene_names = []
    transcript_types = []
    transcript_names = []
    for attribute in genes['attribute']:
        try:
            gene_field, gene_name = attribute.split(';')[5].split('=')
            transcript_type_field, transcript_type = attribute.split(';')[6].split('=')
            transcript_field, transcript_name = attribute.split(';')[7].split('=')
            assert gene_field == 'gene_name'
            assert transcript_type_field == 'transcript_type'
            assert transcript_field == 'transcript_name'
            gene_names.append(gene_name)
            transcript_types.append(transcript_type)
            transcript_names.append(transcript_name)
        except Exception as e:
            print(f'Error: {e}')
            gene_names.append('')
            transcript_names.append('')
    genes['gene_name'] = gene_names
    genes['transcript_name'] = transcript_names
    genes['transcript_type'] = transcript_types
    return genes

def load_chr_sizes(chr_size_path):
    chr_sizes = {}
    with open(chr_size_path) as f:
        for line in f:
            chrom, size = line.strip().split()
            chr_sizes[chrom] = int(size)
    return chr_sizes

def copy_merged_loci_to_inputs(config,config_name):
    source_path = os.path.join(config["preprocessing_data_config"]["output"]["path"],config_name, 'gene_list', 'merged_loci.bed') #hua added
    dest_path = os.path.join(config["inference_config"]["input"]["root"], config["inference_config"]["input"]["locus_list_out"]) #hua locus_list_path to locus_list_out
    os.makedirs(os.path.dirname(dest_path), exist_ok=True)
    import shutil
    shutil.copy(source_path, dest_path)

if __name__ == '__main__':
    main()

