import pandas as pd
import os

# =====================================================
# Step 1: File path configuration
# =====================================================
MPNN_FILE = "MPNN-CNN-TargetPred.txt"
MAP_FILE = "/home/databank/ydn/d3carp-similarty/Ligand_similarity/BindingDB/D3CARP_Ligand_Database.txt"
DISEASE_DIR = "/home/databank/ydn/d3carp-similarty/Target-Disease/"
DISEASE_FILES = ["UniP_Disease_recorded.txt", "UniP_Disease_recorded_0.txt"]
UNIPROT_FILE = "uniprot_gene_species.tsv"
TREMBL_FILE = "uniprot_trembl_subset.tsv"

# =====================================================
# Step 2: Load D3CARP Ligand-Target mapping
# =====================================================
print("\n=== Step 1: Load D3CARP mapping file ===")
map_df = pd.read_csv(MAP_FILE, sep="\t", dtype=str)
if not {"uniprot_id", "Organism"}.issubset(map_df.columns):
    raise ValueError("Mapping file is missing 'uniprot_id' or 'Organism' column.")

uniprot_to_organism = (
    map_df.dropna(subset=["uniprot_id", "Organism"])
    .drop_duplicates(subset=["uniprot_id"])
    .set_index("uniprot_id")["Organism"]
    .to_dict()
)
print(f"Loaded {len(uniprot_to_organism)} UniProt to Organism mappings from D3CARP_Ligand_Database.txt.")

# =====================================================
# Step 3: Load supplemental Target-Disease mappings
# =====================================================
print("\n=== Step 2: Load supplemental Target-Disease data ===")
supplement_dict = {}

for fname in DISEASE_FILES:
    fpath = os.path.join(DISEASE_DIR, fname)
    if os.path.isfile(fpath):
        df_disease = pd.read_csv(
            fpath, sep="\t", header=None, names=["UniProtID", "Entry", "Diseases"], dtype=str
        )
        df_disease["Species"] = df_disease["Entry"].str.extract(r"_([A-Z0-9]+)$")[0]
        df_disease = df_disease.dropna(subset=["Species"])
        supplement_dict.update(df_disease.set_index("UniProtID")["Species"].to_dict())
        print(f"Loaded {len(df_disease)} records from {fname}.")
    else:
        print(f"Warning: file not found -> {fpath}")

print(f"Total supplemental mappings: {len(supplement_dict)}")
uniprot_to_organism.update(supplement_dict)

# =====================================================
# Step 4: Load UniProt SwissProt mapping
# =====================================================
print("\n=== Step 3: Load official UniProt mapping ===")
if os.path.isfile(UNIPROT_FILE):
    df_uniprot = pd.read_csv(UNIPROT_FILE, sep="\t", dtype=str)
    if {"Entry", "Organism"}.issubset(df_uniprot.columns):
        df_uniprot = df_uniprot.dropna(subset=["Entry", "Organism"]).drop_duplicates(subset=["Entry"])
        extra_map = df_uniprot.set_index("Entry")["Organism"].to_dict()
        uniprot_to_organism.update(extra_map)
        print(f"Loaded {len(extra_map)} extra mappings from UniProt file.")
    else:
        print("Warning: UniProt file is missing required columns.")
else:
    print("Warning: UniProt mapping file not found. Skipping this step.")

# =====================================================
# Step 5: Load TrEMBL subset mappings
# =====================================================
print("\n=== Step 4: Load UniProt TrEMBL subset ===")
if os.path.isfile(TREMBL_FILE):
    df_trembl = pd.read_csv(TREMBL_FILE, sep="\t", dtype=str)
    if {"Entry", "Organism"}.issubset(df_trembl.columns):
        df_trembl = df_trembl.dropna(subset=["Entry", "Organism"]).drop_duplicates(subset=["Entry"])
        trembl_dict = df_trembl.set_index("Entry")["Organism"].to_dict()
        uniprot_to_organism.update(trembl_dict)
        print(f"Loaded {len(trembl_dict)} TrEMBL mappings from uniprot_trembl_subset.tsv.")
    else:
        print("Warning: TrEMBL file missing required columns.")
else:
    print("Warning: TrEMBL subset file not found. Skipping this step.")

# =====================================================
# Step 6: Load MPNN predictions and merge mappings
# =====================================================
print("\n=== Step 5: Load MPNN predictions and merge ===")
df_mpnn = pd.read_csv(MPNN_FILE, sep="\t", dtype=str)
if "UniProtID" not in df_mpnn.columns:
    raise ValueError("MPNN-CNN-TargetPred.txt is missing 'UniProtID' column.")

df_mpnn["Organism"] = df_mpnn["UniProtID"].map(uniprot_to_organism)

# =====================================================
# Step 7: Statistics and outputs
# =====================================================
total = len(df_mpnn)
mapped = df_mpnn["Organism"].notna().sum()
human = df_mpnn["Organism"].str.contains("Homo sapiens", na=False).sum()

print(f"MPNN file loaded: {total} records.")
print(f"Successfully mapped {mapped} ({mapped/total:.2%}) UniProtID to organism info.")
print(f"Homo sapiens entries: {human} ({human/total:.2%})")

# Output main file
output_file = "MPNN-CNN-TargetPred_with_species.txt"
df_mpnn.to_csv(output_file, sep="\t", index=False)
print(f"\nOutput written to: {output_file}")

# Output unmatched IDs
unmatched = df_mpnn[df_mpnn["Organism"].isna()][["UniProtID"]]
if len(unmatched) > 0:
    unmatched_file = "MPNN_species_unmatched.txt"
    unmatched.to_csv(unmatched_file, sep="\t", index=False)
    print(f"Unmatched records: {len(unmatched)} (saved to {unmatched_file})")
else:
    print("All UniProtIDs successfully mapped.")

# Species distribution summary
print("\nTop 10 species:")
print(df_mpnn["Organism"].value_counts(dropna=False).head(10))
