#!/usr/bin/env python3
"""
S7: lncRNA filter
Filter s5/s6 results to lncRNA loci only, using gene type annotations from s1.
Reads s5 postprocessed matrices, filters rows matching lncRNA transcript types,
saves filtered matrices and re-runs NMF signature extraction (s6 logic).

Usage:
    python -u s7_lncrna_filter.py -c config_KG1_with_lncrna.yaml
"""

import os
import re
import json
import yaml
import argparse
import numpy as np
import pandas as pd
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from src.signature import nmf
from src.utils.signature import plot_signature_heatmap_full, plot_signature_heatmap_group

LNCRNA_TYPES = [
    'lncRNA', 'lincRNA', 'antisense', 'processed_transcript',
    'retained_intron', 'sense_intronic', 'sense_overlapping',
    'macro_lncRNA', 'bidirectional_promoter_lncRNA', 'non_coding'
]

def load_lncrna_gene_names(gene_list_path):
    """Return set of transcript_names that are lncRNA type."""
    df = pd.read_csv(gene_list_path)
    lncrna_df = df[df['transcript_type'].isin(LNCRNA_TYPES)]
    print(f"  Found {len(lncrna_df)} lncRNA entries out of {len(df)} total in {os.path.basename(gene_list_path)}")
    return set(lncrna_df['transcript_name'].tolist())

def filter_and_save(s5_folder, data_type, lncrna_names, output_path, config, config_name):
    """Load s5 results, filter to lncRNA loci, save filtered outputs."""
    print(f"\n{'='*60}")
    print(f"Processing {data_type.upper()} lncRNA loci...")
    print(f"{'='*60}")

    os.makedirs(output_path, exist_ok=True)

    # Load metadata
    meta_path = os.path.join(s5_folder, data_type, 'metadata.json')
    if not os.path.exists(meta_path):
        print(f"  Skipping {data_type}: metadata not found at {meta_path}")
        return None, None

    with open(meta_path) as f:
        metadata = json.load(f)

    gene_names = metadata['gene_names']
    cap_names  = metadata['cap_names']

    # Load control and experiment matrices from s5 csvs
    ctrl_path = os.path.join(s5_folder, data_type,
                             f"{config['loci_source_config']['atac']['data_path']['control']}.csv")
    exp_path  = os.path.join(s5_folder, data_type,
                             f"{config['loci_source_config']['atac']['data_path']['experimental']}.csv")

    if not os.path.exists(ctrl_path) or not os.path.exists(exp_path):
        print(f"  Skipping {data_type}: s5 csv files not found")
        return None, None

    ctrl_df = pd.read_csv(ctrl_path, index_col=0)
    exp_df  = pd.read_csv(exp_path,  index_col=0)

    # Filter to lncRNA rows
    lncrna_mask = [g in lncrna_names for g in ctrl_df.index]
    ctrl_lncrna = ctrl_df[lncrna_mask]
    exp_lncrna  = exp_df[lncrna_mask]

    n_lncrna = lncrna_mask.count(True)
    print(f"  lncRNA loci retained: {n_lncrna} / {len(ctrl_df)}")

    if n_lncrna == 0:
        print(f"  No lncRNA loci found for {data_type}, skipping.")
        return None, None

    # Save filtered matrices
    ctrl_lncrna.to_csv(os.path.join(output_path, f"{data_type}_lncrna_control.csv"))
    exp_lncrna.to_csv(os.path.join(output_path,  f"{data_type}_lncrna_experiment.csv"))

    # Compute and save differential
    diff = np.log2((exp_lncrna.values + 1e-6) / (ctrl_lncrna.values + 1e-6))
    diff_df = pd.DataFrame(diff, index=ctrl_lncrna.index, columns=ctrl_lncrna.columns)
    diff_df.to_csv(os.path.join(output_path, f"{data_type}_lncrna_differential.csv"))

    # Save metadata
    lncrna_meta = {
        'gene_names': ctrl_lncrna.index.tolist(),
        'cap_names': cap_names,
        'n_lncrna_loci': n_lncrna,
        'lncrna_types_used': LNCRNA_TYPES
    }
    with open(os.path.join(output_path, f"{data_type}_lncrna_metadata.json"), 'w') as f:
        json.dump(lncrna_meta, f, indent=2)

    print(f"  Saved filtered matrices to {output_path}")
    return diff_df, cap_names

def run_nmf_signatures(diff_df, cap_names, data_type, output_path, config):
    """Run NMF on lncRNA differential matrix and save signature outputs."""
    if diff_df is None or len(diff_df) < 4:
        print(f"  Skipping NMF for {data_type}: too few loci ({0 if diff_df is None else len(diff_df)})")
        return

    print(f"  Running NMF on {data_type} lncRNA matrix ({diff_df.shape})...")
    nmf_config = config['signature_config']['nmf']

    sig_output = os.path.join(output_path, 'nmf', data_type)
    os.makedirs(sig_output, exist_ok=True)

    try:
        result = nmf.run_nmf(
            diff_df.values,
            gene_names=diff_df.index.tolist(),
            cap_names=cap_names,
            n_components=nmf_config.get('n_components', 4),
            min_components=nmf_config.get('min_components', 3),
            max_components=min(nmf_config.get('max_components', 40), len(diff_df) - 1),
            init=nmf_config.get('init', 'nndsvd'),
            max_iter=nmf_config.get('max_iter', 200),
            seed=config['signature_config'].get('seed', 9),
            output_path=sig_output
        )
        print(f"  NMF complete for {data_type} lncRNA. Results saved to {sig_output}")
    except Exception as e:
        print(f"  NMF failed for {data_type}: {e}")

def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("-c", "--config", help="Path to YAML config.")
    args = parser.parse_args()

    config_path = args.config
    config_name = re.sub(r'^(?:config_)?|\.ya?ml$', '', config_path)
    with open(config_path, 'r') as f:
        config = yaml.safe_load(f)

    # Paths
    s1_gene_list_dir = os.path.join(
        config['preprocessing_data_config']['output']['path'],
        config_name, 'gene_list'
    )
    s5_folder = os.path.join(
        config['postprocess_config']['output']['path'],
        config_name
    )
    output_root = os.path.join(
        os.path.dirname(config['postprocess_config']['output']['path']),
        's7_lncrna', config_name
    )

    print(f"S7 lncRNA filter started for {config_name}")
    print(f"S1 gene lists : {s1_gene_list_dir}")
    print(f"S5 input      : {s5_folder}")
    print(f"S7 output     : {output_root}")

    data_types = ['conserved', 'upregulated', 'downregulated']

    for data_type in data_types:
        gene_list_path = os.path.join(s1_gene_list_dir, f"{data_type}_gene_list.csv")
        if not os.path.exists(gene_list_path):
            print(f"\nSkipping {data_type}: gene list not found at {gene_list_path}")
            continue

        lncrna_names = load_lncrna_gene_names(gene_list_path)

        out_path = os.path.join(output_root, data_type)
        diff_df, cap_names = filter_and_save(
            s5_folder, data_type, lncrna_names, out_path, config, config_name
        )
        run_nmf_signatures(diff_df, cap_names, data_type, out_path, config)

    print("\nS7 lncRNA filter complete.")

if __name__ == "__main__":
    main()
