import os
import zarr
import yaml
import numpy as np
import pandas as pd
from tqdm import tqdm
import argparse
import re
import sys
def main():
    # Load config
    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
    print('working on ' + config_path)
    with open(config_path, 'r') as f:
        config = yaml.safe_load(f)
    config_name=re.sub(r'^(?:config_)?|\.ya?ml$', '',config_path)

    output_path = config['extract_data_config']['output']['path']+'/'+config_name
    os.makedirs(output_path, exist_ok=True)
    print('output: ' + output_path)
    preprocessing_data_path = config['preprocessing_data_config']['output']['path']
    postprocess_data_path = config['merge_inference_config']['output']['path']

    experimental_celltype = config['loci_source_config']['atac']['data_path']['experimental']
    control_celltype = config['loci_source_config']['atac']['data_path']['control']

    # Load loci groups
    loci_groups_list = ['upregulated', 'downregulated', 'conserved']
    loci_groups_df_dict = {}
    for loci_group in loci_groups_list:
        loci_groups_df = pd.read_csv(os.path.join(preprocessing_data_path,config_name, 'gene_list', f'{loci_group}_gene_list.csv'))
        loci_groups_df = remove_chr_X(loci_groups_df)
        loci_groups_df_dict[loci_group] = loci_groups_df
        # Save loci groups
        loci_groups_df.to_csv(os.path.join(output_path, f'{loci_group}_gene_list_no_chrX.csv'), index=False)

    # Load inference data
    celltypes = experimental_celltype.split(',')+ control_celltype.split(',') #hua changed [experimental_celltype, control_celltype]
    data_dict = {}
    for celltype in celltypes:
        data_path = os.path.join(postprocess_data_path, celltype, 'data.zarr')
        data = zarr.open(data_path, mode='r')
        data_dict[celltype] = data

    # Extract data for each loci group (500bp is better) for each celltype
    for loci_group in loci_groups_list:
        data_celltype_list = []
        for celltype in celltypes:
            loci_groups_df = loci_groups_df_dict[loci_group]
            cell_type_loci_data = extract_data(loci_groups_df, data_dict[celltype])
            data_celltype_list.append(cell_type_loci_data) # Shape (n_loci, n_cap, 2[max, mean])
        data_loci = np.array(data_celltype_list) # Shape (n_celltypes, n_loci, n_cap, 2[max, mean])

        # Save data and annotations
        np.save(os.path.join(output_path, f'{loci_group}_data.npy'), data_loci)
        #print('writing to: '+ os.path.join(output_path, f'{loci_group}_data.npy'))
        # Save cell types, loci groups, and loci
        with open(os.path.join(output_path, f'{loci_group}_celltypes.txt'), 'w') as f:
            for celltype in celltypes:
                f.write(celltype + '\n')

def extract_data(loci_groups_df, data):
    # Extract data for each loci group (500bp is better)
    extracted_data = []
    for index, row in tqdm(loci_groups_df.iterrows(), total=len(loci_groups_df), disable=not sys.stderr.isatty()):
        chr_name = row['seqname']
        start = row['tss'] - 500
        end = row['tss'] + 500
        loci_data = data['chrs'][chr_name][start:end, :]
        extracted_data.append((loci_data.max(axis=0), loci_data.mean(axis=0)))
    if len(extracted_data) != 0:
        extracted_data = np.array(extracted_data).transpose(0, 2, 1)
    return extracted_data # Shape (n_loci, n_cap, 2[max, mean])

def remove_chr_X(loci_groups_df):
    loci_groups_df = loci_groups_df[loci_groups_df['seqname'] != 'chrX']
    return loci_groups_df

if __name__ == '__main__':
    main()
