import pandas as pd
import os
import re
import yaml

# Load config
with open("config.yaml", "r") as file:
    config = yaml.safe_load(file)

SUMMARY_FILE = os.path.join(config["genome_dir_WGS"], "genome_summary.csv")
NAMES_DMP    = config["names_dmp_path"]

rule all:
    input:
        "filtering_complete.marker"

rule filter_summary_inplace:
    input:
        csv   = SUMMARY_FILE,
        names = NAMES_DMP
    output:
        marker = touch("filtering_complete.marker")
    run:
        print(f"Loading {input.csv}...")

        # Auto-detect delimiter
        with open(input.csv, 'r') as f:
            first_line = f.readline()
            sep = '\t' if '\t' in first_line else ','

        df = pd.read_csv(input.csv, sep=sep)
        df.columns = df.columns.str.strip()

        if 'Genus' not in df.columns:
            raise KeyError(
                f"'Genus' column not found. Available: {list(df.columns)}")

        # ── Clean GTDB-Tk suffixes ─────────────────────────────────────────
        # GTDB appends _A, _B, _C or _001, _002 to genus and species names.
        # Remove trailing _[letters/digits] from Genus and Species columns.
        def strip_gtdb_suffix(val):
            if pd.isna(val):
                return val
            return re.sub(r'_[A-Za-z0-9]+$', '', str(val).strip())

        df['Genus'] = df['Genus'].apply(strip_gtdb_suffix)

        if 'Species' in df.columns:
            df['Species'] = df['Species'].apply(strip_gtdb_suffix)

        print("Cleaned GTDB-Tk genus/species suffixes.")

        # ── Validate genera against NCBI names.dmp ─────────────────────────
        unique_genera = set(df['Genus'].dropna().unique())
        valid_genera  = set()

        print(f"Validating {len(unique_genera)} genera against NCBI names.dmp...")
        with open(input.names, 'r') as f:
            for line in f:
                if "\tscientific name\t" in line:
                    parts = line.split('|')
                    name  = parts[1].strip()
                    if name in unique_genera:
                        valid_genera.add(name)

        initial_count = len(df)
        filtered_df   = df[df['Genus'].isin(valid_genera)]
        filtered_df.to_csv(input.csv, sep=sep, index=False)

        print(f"Initial rows : {initial_count}")
        print(f"Rows kept    : {len(filtered_df)}")
        print(f"Saved back to {input.csv}")
