import yaml
import os 
import numpy as np
import argparse
import re
from src.ChromnitronDataset import ChromnitronDataset, create_from_original_data
from src.utils.vis import plot_histogram, plot_chipseq_tracks, plot_variability 

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)
    print('started for '+ config_name) 
    with open(config_path, 'r') as f:
        config = yaml.safe_load(f)
    
    input_path = config['extract_data_config']['output']['path']
    output_path_root = config['postprocess_config']['output']['path']
    
    data_types = ['upregulated', 'downregulated', 'conserved']
    for data_type in data_types:
        print(f"\n{'='*60}")
        print(f"Processing {data_type.upper()} data...")
        print(f"{'='*60}")
        
        output_path = os.path.join(output_path_root,config_name, data_type)
        os.makedirs(output_path, exist_ok=True)
        npy_path = os.path.join(input_path,config_name, f"{data_type}_data.npy")
        gene_list_path = os.path.join(input_path,config_name, f"{data_type}_gene_list_no_chrX.csv")
        cap_list_path = original_cap_path = os.path.join(config['merge_inference_config']['output']['path'],config['loci_source_config']['atac']['data_path']['experimental'].split(',')[0], 'cap_list.txt')
        celltypes_list_path = os.path.join(input_path, config_name,f"{data_type}_celltypes.txt")

        ### STEP 1: Create ChromnitronDataset
        print(f"Loading dataset from {npy_path}...")
        dataset = create_from_original_data(npy_path, gene_list_path, cap_list_path, celltypes_list_path, mode='max')
        
        ### STEP 2: Compute differential matrix
        print("Computing log2 fold change differential...")
        dataset.compute_differential(log_transform=True, base='2')
        
        # Histogram (if enabled)
        if config['postprocess_config']['plotting'].get('plot_histograms', True):
            print("Generating histogram after differential computation...")
            diff_matrix = dataset.get_differential()
            plot_histogram(
                data=diff_matrix.flatten(),
                bins=100,
                plot_title=f"{data_type.capitalize()} - Log2 Fold Change Distribution",
                x_name="Log2 Fold Change",
                y_name="Frequency",
                filename=f"{data_type}_histogram_1_differential",
                data_save_path=output_path
            )
            print(f"Saved differential histogram to {output_path}/{data_type}_histogram_1_differential.png")
        else:
            print("Skipping histogram generation (disabled in config)")

        ### STEP 3: Variance filtering (if enabled)
        if config['postprocess_config']['variance_filtering'].get('enable', True):
            n_genes = config['postprocess_config']['variance_filtering'].get('n_genes', 500)
            n_caps = config['postprocess_config']['variance_filtering'].get('n_caps', 100)
            print(f"Applying variance filtering: top {n_genes} genes, top {n_caps} CAPs...")
            var_results = dataset.variance_filter(n_genes=n_genes, n_caps=n_caps, condition='differential')
        else:
            print("Skipping variance filtering (disabled in config)")
            # Create a mock var_results for plotting if variability plot is enabled
            var_results = None
        
        # Histogram (if both variance filtering and histograms are enabled)
        if (config['postprocess_config']['variance_filtering'].get('enable', True) and 
            config['postprocess_config']['plotting'].get('plot_histograms', True)):
            print("Generating histogram after variance filtering...")
            filtered_diff_matrix = dataset.get_differential()
            n_genes = config['postprocess_config']['variance_filtering'].get('n_genes', 500)
            n_caps = config['postprocess_config']['variance_filtering'].get('n_caps', 100)
            plot_histogram(
                data=filtered_diff_matrix.flatten(),
                bins=100,
                plot_title=f"{data_type.capitalize()} - After Variance Filtering (Top {n_genes} genes, {n_caps} CAPs)",
                x_name="Log2 Fold Change",
                y_name="Frequency",
                filename=f"{data_type}_histogram_2_variance_filtered",
                data_save_path=output_path
            )
            print(f"Saved variance filtered histogram to {output_path}/{data_type}_histogram_2_variance_filtered.png")
        else:
            if not config['postprocess_config']['variance_filtering'].get('enable', True):
                print("Skipping variance filtering histogram (variance filtering disabled)")
            else:
                print("Skipping variance filtering histogram (plotting disabled in config)")
        
        # Plot variability (if both variance filtering and variability plotting are enabled)
        if (config['postprocess_config']['variance_filtering'].get('enable', True) and 
            config['postprocess_config']['plotting'].get('plot_variability', True) and 
            var_results is not None):
            variability_plot_path = os.path.join(output_path, f"{data_type}_variability.png")
            plot_variability(var_results, save_path=variability_plot_path, show_labels=True, max_labels=50)
            print(f"Saved variability plot to {variability_plot_path}")
        else:
            if not config['postprocess_config']['variance_filtering'].get('enable', True):
                print("Skipping variability plot (variance filtering disabled)")
            else:
                print("Skipping variability plot (plotting disabled in config)")
        
        ### Plot ChIP-seq tracks before noise removal (if enabled)
        if config['postprocess_config']['plotting'].get('plot_chipseq', True):
            print("Generating ChIP-seq tracks before noise removal...")
            gene_names = dataset.filtered_gene_names or dataset.gene_names
            cap_names = dataset.filtered_cap_names or dataset.cap_names
            
            # Plot all three conditions
            conditions_to_plot = [
                (dataset.get_control(), dataset.control_name, "raw"),
                (dataset.get_experiment(), dataset.experiment_name, "raw"), 
                (dataset.get_differential(), "Differential", "log2fc")
            ]
            
            for matrix, condition_name, data_type_suffix in conditions_to_plot:
                plot_chipseq_tracks(
                    matrix, 
                    cap_names,
                    loci_list=gene_names,
                    n_caps_per_page=100,  # Smaller pages for multiple conditions
                    track_height=0.25,   # Smaller tracks
                    figsize_width=250,   # Smaller width
                    smooth_window=3,     # Less smoothing for raw data
                    log_scale=False,
                    show_percentiles=True,
                    data_save_path=output_path,
                    filename_prefix=f"{data_type}_{condition_name.lower().replace(' ', '_')}_before_noise_removal",
                    output_format="pdf"
                )
                print(f"Saved {condition_name} tracks (before noise removal) to {output_path}")
        else:
            print("Skipping ChIP-seq tracks before noise removal (disabled in config)")
        
        ### STEP 4: Remove noise by keeping only top X% of signal (if enabled)
        if config['postprocess_config']['remove_noise'].get('enable', True):
            top_percent = config['postprocess_config']['remove_noise'].get('top_percent', 10.0)
            conditions = config['postprocess_config']['remove_noise'].get('conditions', ['all'])
            print(f"Removing noise (keeping top {top_percent}% extreme values)...")
            dataset.remove_noise(top_percent=top_percent, conditions=conditions)
        else:
            print("Skipping noise removal (disabled in config)")
        
        # Histogram (if both noise removal and histograms are enabled)
        if (config['postprocess_config']['remove_noise'].get('enable', True) and 
            config['postprocess_config']['plotting'].get('plot_histograms', True)):
            print("Generating histogram after noise removal...")
            final_diff_matrix = dataset.get_differential()
            top_percent = config['postprocess_config']['remove_noise'].get('top_percent', 10.0)
            plot_histogram(
                data=final_diff_matrix.flatten(),
                bins=100,
                plot_title=f"{data_type.capitalize()} - After Noise Removal (Top {top_percent}% Extreme Values)",
                x_name="Log2 Fold Change",
                y_name="Frequency",
                filename=f"{data_type}_histogram_3_noise_removed",
                data_save_path=output_path
            )
            print(f"Saved noise removed histogram to {output_path}/{data_type}_histogram_3_noise_removed.png")
        else:
            if not config['postprocess_config']['remove_noise'].get('enable', True):
                print("Skipping noise removal histogram (noise removal disabled)")
            else:
                print("Skipping noise removal histogram (plotting disabled in config)")
        
        ### Plot ChIP-seq tracks AFTER noise removal (if both noise removal and chipseq plotting are enabled)
        if (config['postprocess_config']['remove_noise'].get('enable', True) and 
            config['postprocess_config']['plotting'].get('plot_chipseq', True)):
            print("Generating ChIP-seq tracks AFTER noise removal...")
            
            # Plot all three conditions after noise removal
            conditions_to_plot_after = [
                (dataset.get_control(), dataset.control_name, "raw_clean"),
                (dataset.get_experiment(), dataset.experiment_name, "raw_clean"), 
                (dataset.get_differential(), "Differential", "log2fc_clean")
            ]
            
            for matrix, condition_name, data_type_suffix in conditions_to_plot_after:
                plot_chipseq_tracks(
                    matrix, 
                    cap_names,
                    loci_list=gene_names,
                    n_caps_per_page=100,  # Smaller pages for multiple conditions
                    track_height=0.25,   # Smaller tracks
                    figsize_width=250,   # Smaller width
                    smooth_window=3,     # Less smoothing
                    log_scale=False,
                    show_percentiles=True,
                    data_save_path=output_path,
                    filename_prefix=f"{data_type}_{condition_name.lower().replace(' ', '_')}_after_noise_removal",
                    output_format="pdf"
                )
                print(f"Saved {condition_name} tracks (after noise removal) to {output_path}")
        else:
            if not config['postprocess_config']['remove_noise'].get('enable', True):
                print("Skipping ChIP-seq tracks after noise removal (noise removal disabled)")
            else:
                print("Skipping ChIP-seq tracks after noise removal (plotting disabled in config)")

        ### STEP 5: Save processed dataset
        print("Saving processed dataset...")
        dataset.save(output_path)
        print(f"Saved dataset to {output_path}")
        
        # Summary
        print("Dataset summary:")
        dataset.summary()
    print('Done for '+ config_name) 

if __name__ == "__main__":
    main()
