import csv
import pandas as pd
import os
import glob
import numpy as np
from sklearn.preprocessing import MinMaxScaler

# =====================================================
# Step 0: Detect project prefix and create output folder
# =====================================================
def detect_project_prefix():
    candidates = glob.glob("*_analysis.txt")
    if not candidates:
        raise FileNotFoundError("No docking file found (expected pattern '*_analysis.txt').")
    docking_file = candidates[0]
    prefix = docking_file.replace("_analysis.txt", "")
    output_dir = f"results_{prefix}"
    os.makedirs(output_dir, exist_ok=True)
    print(f"Detected project prefix: {prefix}")
    print(f"Output directory: {output_dir}")
    return prefix, docking_file, output_dir


# =====================================================
# Step 1: Filter Homo sapiens data (including AI)
# =====================================================
def filter_human_sources(prefix, docking_file, outdir):
    print("\n=== Step 1: Filter Homo sapiens data ===")

    # Docking
    try:
        df_dock = pd.read_csv(docking_file, sep=';')
        human_dock = df_dock[df_dock['Type'].str.endswith('_HUMAN', na=False)]
        dock_out = os.path.join(outdir, f"{prefix}_analysis_human.txt")
        human_dock.to_csv(dock_out, sep=';', index=False)
        print(f"Docking: {len(df_dock)} -> {len(human_dock)} (Homo sapiens)")
    except Exception as e:
        print(f"Error reading docking file: {e}")

    # 2D and 3D similarity
    for src, out_name in [
        ("Similarity_2D_MACCS_detail.txt", "Similarity_2D_MACCS_detail_human.txt"),
        ("Similarity_3D_Rigid-LS-align_detail.txt", "Similarity_3D_Rigid-LS-align_detail_human.txt")
    ]:
        try:
            df = pd.read_csv(src, sep='\t')
            human_df = df[df['Organism'].str.strip().eq('Homo sapiens')]
            out_path = os.path.join(outdir, out_name)
            human_df.to_csv(out_path, sep='\t', index=False)
            print(f"{src}: {len(df)} -> {len(human_df)} (Homo sapiens)")
        except Exception as e:
            print(f"Error reading {src}: {e}")

    # AI file
    try:
        ai_file = "MPNN-CNN-TargetPred_with_species.txt"
        if os.path.exists(ai_file):
            df_ai = pd.read_csv(ai_file, sep='\t')
            if 'Organism' in df_ai.columns:
                human_ai = df_ai[df_ai['Organism'].str.contains('Homo sapiens', case=False, na=False)]
                ai_out = os.path.join(outdir, "MPNN-CNN-TargetPred_human.txt")
                human_ai.to_csv(ai_out, sep='\t', index=False)
                print(f"AI file: {len(df_ai)} -> {len(human_ai)} (Homo sapiens)")
            else:
                print("Warning: 'Organism' column not found in AI file.")
        else:
            print("Warning: MPNN-CNN-TargetPred_with_species.txt not found.")
    except Exception as e:
        print(f"Warning: AI filtering skipped: {e}")

    print(f"\nHuman-filtered files saved in: {outdir}\n")


# =====================================================
# Step 1.5: Filter high-confidence targets
# =====================================================
def filter_high_confidence(prefix, outdir):
    print("\n=== Step 1.5: Filter high-confidence targets (Score / Similarity / Prediction) ===")

    # Docking Score ≤ -9
    dock_file = os.path.join(outdir, f"{prefix}_analysis_human.txt")
    try:
        df_dock = pd.read_csv(dock_file, sep=';')
        if 'Docking_Score' in df_dock.columns:
            df_dock = df_dock[df_dock['Docking_Score'] <= -9]
            print(f"Docking: filtered to {len(df_dock)} (Score ≤ -9)")
        df_dock.to_csv(dock_file, sep=';', index=False)
    except Exception as e:
        print(f"Warning: Docking filtering skipped: {e}")

    # 2D Similarity ≥ 0.8
    sim2d_file = os.path.join(outdir, "Similarity_2D_MACCS_detail_human.txt")
    try:
        df_2d = pd.read_csv(sim2d_file, sep='\t')
        if 'Similarity' in df_2d.columns:
            df_2d = df_2d[df_2d['Similarity'].astype(float) >= 0.8]
            print(f"2D Similarity: filtered to {len(df_2d)} (≥ 0.8)")
        df_2d.to_csv(sim2d_file, sep='\t', index=False)
    except Exception as e:
        print(f"Warning: 2D filtering skipped: {e}")

    # 3D Similarity ≥ 0.7
    sim3d_file = os.path.join(outdir, "Similarity_3D_Rigid-LS-align_detail_human.txt")
    try:
        df_3d = pd.read_csv(sim3d_file, sep='\t')
        if 'Similarity' in df_3d.columns:
            df_3d = df_3d[df_3d['Similarity'].astype(float) >= 0.7]
            print(f"3D Similarity: filtered to {len(df_3d)} (≥ 0.7)")
        df_3d.to_csv(sim3d_file, sep='\t', index=False)
    except Exception as e:
        print(f"Warning: 3D filtering skipped: {e}")

    # AI Score ≥ 0.9 and Homo sapiens
    ai_file = os.path.join(outdir, "MPNN-CNN-TargetPred_human.txt")
    try:
        if os.path.exists(ai_file):
            df_ai = pd.read_csv(ai_file, sep='\t')
            if {'Score', 'Organism'}.issubset(df_ai.columns):
                df_ai = df_ai[df_ai['Organism'].str.contains('Homo sapiens', case=False, na=False)]
                df_ai = df_ai[df_ai['Score'].astype(float) >= 0.9]
                ai_out = os.path.join(outdir, "MPNN-CNN-TargetPred_filtered.txt")
                df_ai.to_csv(ai_out, sep='\t', index=False)
                print(f"AI Prediction: filtered to {len(df_ai)} (Homo sapiens & Score ≥ 0.9)")
            else:
                print("Warning: 'Score' or 'Organism' column not found in AI file.")
        else:
            print("Warning: MPNN-CNN-TargetPred_human.txt not found.")
    except Exception as e:
        print(f"Warning: AI filtering skipped: {e}")

    print(f"\nAll filtered human data updated in {outdir}\n")


# =====================================================
# Step 2: Extract and label IDs
# =====================================================
def read_ids_from_txt(filepath, column_index, skip_lines=1, limit=None, delimiter='\t'):
    ids = set()
    if not os.path.exists(filepath):
        return ids
    with open(filepath, 'r', encoding='utf-8') as f:
        reader = csv.reader(f, delimiter=delimiter)
        for _ in range(skip_lines):
            next(reader, None)
        for i, row in enumerate(reader):
            if limit and i >= limit:
                break
            if len(row) > column_index and row[column_index].strip():
                ids.add(row[column_index].strip())
    return ids


def write_output(filepath, data):
    with open(filepath, 'w', encoding='utf-8') as f:
        if isinstance(data, list):
            f.write('\n'.join(data))
        else:
            f.write(data)


def extract_and_label(prefix, outdir):
    print("\n=== Step 2: Extract and label IDs ===")

    docking_file = os.path.join(outdir, f"{prefix}_analysis_human.txt")
    docking = read_ids_from_txt(docking_file, 1, skip_lines=1, delimiter=';')
    MACCS = read_ids_from_txt(os.path.join(outdir, "Similarity_2D_MACCS_detail_human.txt"), 17, skip_lines=1, delimiter='\t')
    Rigid = read_ids_from_txt(os.path.join(outdir, "Similarity_3D_Rigid-LS-align_detail_human.txt"), 17, skip_lines=1, delimiter='\t')
    AI = read_ids_from_txt(os.path.join(outdir, "MPNN-CNN-TargetPred_filtered.txt"), 2, skip_lines=1, delimiter='\t')

    print(f"docking: {len(docking)}")
    print(f"MACCS: {len(MACCS)}")
    print(f"Rigid: {len(Rigid)}")
    print(f"AI: {len(AI)}")

    union_set = docking.union(MACCS, Rigid, AI)
    union_similarity = MACCS.union(Rigid)
    print(f"Total unique IDs: {len(union_set)}")

    label_data = []
    for item in sorted(union_set):
        tag = ''
        if item in docking:
            tag += 'a'
        if item in union_similarity:
            tag += 'b'
        if item in AI:
            tag += 'c'
        label_data.append(f'{item},{tag}')

    label_file = os.path.join(outdir, f"{prefix}_label.csv")
    write_output(label_file, '\n'.join(label_data))
    print(f"Label file created: {label_file} ({len(label_data)} records)")


# =====================================================
# Helper: Convert ligand potency to nM
# =====================================================
def potency_to_nM(value):
    if pd.isna(value) or value == '':
        return np.nan
    s = str(value).strip().upper()
    # Extract numeric part
    import re
    match = re.search(r'([<>~]?\s*[\d.]+)', s)
    if not match:
        return np.nan
    num_str = match.group(1).replace('<', '').replace('>', '').replace('~', '').strip()
    try:
        num = float(num_str)
    except:
        return np.nan
    # Convert to nM
    if 'MM' in s or 'MILLIMOLAR' in s:
        return num * 1e6
    elif 'UM' in s or 'MICROMOLAR' in s:
        return num * 1e3
    elif 'NM' in s or 'NANOMOLAR' in s:
        return num
    elif 'PM' in s or 'PICOMOLAR' in s:
        return num / 1e3
    else:
        return num


# =====================================================
# Step 3: Merge all datasets and weighted scoring
# =====================================================
def safe_merge(target, source, prefix):
    source = source.rename(columns={
        f'{prefix}UniProt': 'UniProt_ID',
        f'{prefix}UniProtID': 'UniProt_ID'
    })
    source = source.drop_duplicates(subset=['UniProt_ID'])
    return pd.merge(target, source, on='UniProt_ID', how='left')


def integrate_all(prefix, outdir):
    print("\n=== Step 3: Data integration (Pandas) ===")

    target_file = os.path.join(outdir, f"{prefix}_label.csv")
    jst_file = os.path.join(outdir, f"{prefix}_analysis_human.txt")
    sim2d_file = os.path.join(outdir, "Similarity_2D_MACCS_detail_human.txt")
    sim3d_file = os.path.join(outdir, "Similarity_3D_Rigid-LS-align_detail_human.txt")
    mpnn_file = os.path.join(outdir, "MPNN-CNN-TargetPred_filtered.txt")
    target_df = pd.read_csv(target_file, header=None, names=['UniProt_ID', 'Original_Info'])
    jst = pd.read_csv(jst_file, sep=';').add_prefix('Docking_')
    sim2d = pd.read_csv(sim2d_file, sep='\t').add_prefix('2D_')
    sim3d = pd.read_csv(sim3d_file, sep='\t').add_prefix('3D_')
    mpnn = pd.read_csv(mpnn_file, sep='\t').add_prefix('MPNN_') if os.path.exists(mpnn_file) else pd.DataFrame()

    merged = target_df.copy()
    merged = safe_merge(merged, jst, 'Docking_')
    merged = safe_merge(merged, sim2d, '2D_')
    merged = safe_merge(merged, sim3d, '3D_')
    if not mpnn.empty:
        merged = safe_merge(merged, mpnn, 'MPNN_')

    # === Step 4: Weighted scoring (full version) ===
    print("\n=== Step 4: Weighted scoring (mean-fill + normalization) ===")
    
    # Prepare Ligand_Potency_nM
    if 'Docking_Ligand_Potency' in merged.columns:
        merged['Ligand_Potency_nM'] = merged['Docking_Ligand_Potency'].apply(potency_to_nM)
    elif 'Ligand_Potency' in merged.columns:
        merged['Ligand_Potency_nM'] = merged['Ligand_Potency'].apply(potency_to_nM)
    else:
        merged['Ligand_Potency_nM'] = np.nan
    
    # Prepare Original_Info_Score
    score_map = {'a': 0.7, 'b': 0.7, 'c': 0.7, 'ab': 1.5, 'ac': 1.5, 'bc': 1.5, 'abc': 3}
    merged['Original_Info_Score'] = merged['Original_Info'].map(score_map)
    merged.loc[merged['Original_Info'] == '', 'Original_Info_Score'] = np.nan
    
    # --- Define all weighted features ---
    features = {
        'Original_Info_Score': 2.0,
        'Docking_Score_Ratio': 1.6,
        'MPNN_Score': 1.5,
        '2D_Similarity': 1.0,
        '3D_Similarity': 1.0,
        'Docking_Atom_Efficiency': 1.2,
        'Ligand_Potency_nM': -1.5,
        'Docking_Score': -2.5
    }
    
    # Ensure all feature columns exist and fill missing values with column mean
    for col in features:
        if col not in merged.columns:
            merged[col] = np.nan
        merged[col] = pd.to_numeric(merged[col], errors='coerce')
        merged[col] = merged[col].fillna(merged[col].mean())
    
    # Normalize ALL features together using MinMaxScaler
    scaler = MinMaxScaler()
    normalized = pd.DataFrame(
        scaler.fit_transform(merged[list(features.keys())]), 
        columns=features.keys()
    )
    
    # Apply weights to normalized features
    for col in features:
        normalized[col] *= features[col]
    
    # Calculate Total_Score
    merged['Total_Score'] = normalized.sum(axis=1)
    
    # Priority + sort
    merged['priority'] = ((merged['2D_Similarity'] == 1) | (merged['3D_Similarity'] == 1)).astype(int)
    merged = merged.sort_values(by=['priority', 'Total_Score'], ascending=[False, False]).reset_index(drop=True)
    merged = merged.drop(columns=['priority'])
    
    cols = ['Total_Score'] + [c for c in merged.columns if c != 'Total_Score']
    merged = merged[cols]
    
    out_file = os.path.join(outdir, f"{prefix}_target_weighted_sorted.csv")
    merged.to_csv(out_file, index=False)
    print(f"✅ Weighted and sorted file written: {out_file}")
    print(f"Rows: {len(merged)}, Columns: {len(merged.columns)}")
    print("\nPreview (first 3 rows, first 5 cols):")
    print(merged.iloc[:3, :5].to_string(index=False))


# =====================================================
# Entry point
# =====================================================
def main():
    print("\n=== Unified Integration Script (Docking + 2D + 3D + AI + Filtering + Weighted Sorting) ===\n")
    prefix, docking_file, outdir = detect_project_prefix()
    filter_human_sources(prefix, docking_file, outdir)
    filter_high_confidence(prefix, outdir)
    extract_and_label(prefix, outdir)
    integrate_all(prefix, outdir)
    print(f"\nAll steps completed for project [{prefix}].")
    print(f"All output files are in: {outdir}\n")


if __name__ == "__main__":
    main()