import os
import yaml
import pandas as pd
import glob
from snakemake.io import expand
from pathlib import Path
import tempfile

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

# Resolve the directory that contains this snakefile (for relative env paths)
SNAKEFILE_DIR = os.path.dirname(os.path.abspath(workflow.snakefile))

def get_genomes():
    genome_dir = config["genome_dir_WGS"]
    return [f.replace(".fasta", "") for f in os.listdir(genome_dir) if f.endswith(".fasta")]

def get_unique_genera(csv_path):
    df = pd.read_csv(csv_path)
    return df['Genus'].dropna().unique().tolist()

genome_summary_df = pd.read_csv(os.path.join(config["genome_dir_WGS"], "genome_summary.csv"))
genome_summary_df['Assembly'] = genome_summary_df['Assembly'].astype(str)
genome_summary_df['Genus']    = genome_summary_df['Genus'].fillna("UnknownGenus").astype(str)

assembly_to_genus = dict(zip(genome_summary_df['Assembly'], genome_summary_df['Genus']))

GCF_FILES  = glob.glob(os.path.join(config["GCF_dir"], "*_final_GCF.txt"))
GENUS_LIST = [os.path.basename(f).replace("_final_GCF.txt", "") for f in GCF_FILES]

assembly_list      = genome_summary_df['Assembly'].tolist()
genomes_fasta_list = [f.replace(".fasta", "") for f in os.listdir(config["genome_dir_WGS"]) if f.endswith(".fasta")]

def get_basename(path):
    return Path(path).stem

# ── FastANI pair selection ─────────────────────────────────────────────────────
def get_query_ref_pairs(fastani_dir):
    pairs = []
    if not os.path.exists(fastani_dir):
        return pairs
    for f in os.listdir(fastani_dir):
        if f.endswith(".txt"):
            query_file = os.path.join(fastani_dir, f)
            if os.path.getsize(query_file) == 0:
                continue
            with open(query_file, 'r') as fin:
                lines = fin.readlines()
                if not lines:
                    continue
                best_match = sorted(lines, key=lambda x: float(x.split()[2]), reverse=True)[0]
                parts = best_match.split()
                if len(parts) >= 2:
                    query = get_basename(parts[0])
                    ref   = get_basename(parts[1])
                    pairs.append((query, ref))
    return pairs

# ── Target rule ───────────────────────────────────────────────────────────────
rule all:
    input:
        expand(os.path.join(config["prokka_dir"], "{genome}", "{genome}.gff"),  genome=genomes_fasta_list),
        expand(os.path.join(config["prokka_dir"], "{genome}", "{genome}.fna"),  genome=genomes_fasta_list),
        expand(os.path.join(config["barrnap_dir"], "{genome}_16S.fasta"),        genome=assembly_list),
        expand(os.path.join(config["barrnap_dir"], "{genome}_filtered_16S.fasta"), genome=assembly_list),
        expand(os.path.join(config["blast_dir"],   "{genome}.blast.csv"),        genome=assembly_list),

        expand(f"{config['taxid_dir']}/{{genus}}.taxid",
               genus=get_unique_genera(config["genome_dir_WGS"] + "/genome_summary.csv")),
        expand("{output_dir}/{genus}_species_contigs.txt",
               genus=get_unique_genera(config["genome_dir_WGS"] + "/genome_summary.csv"),
               output_dir=config["GCF_dir"]),
        expand("{output_dir}/{genus}_final_GCF.txt",
               genus=get_unique_genera(config["genome_dir_WGS"] + "/genome_summary.csv"),
               output_dir=config["GCF_dir"]),

        expand(f"{config['download_dir']}/{{genus}}/{{genus}}_genome_paths.txt", genus=GENUS_LIST),
        expand(f"{config['fast_ani']}/{{assembly}}.txt",                          assembly=assembly_list),

        f"{config['output_dir']}/fastani_complete.txt",
        f"{config['output_dir']}/fastani_processing_complete.txt",

        [os.path.join(config["aai_dir"],  f"{q}_vs_{r}_aai.txt")
         for q, r in get_query_ref_pairs(config["fast_ani"])],
        [os.path.join(config["pocp_dir"], f"{q}_vs_{r}_pocp.txt")
         for q, r in get_query_ref_pairs(config["fast_ani"])],

        f"{config['output_dir']}/aai_pocp_completion_marker"

# ── Annotation ────────────────────────────────────────────────────────────────
rule annotate_prokka_fasta:
    input:
        fasta=lambda wildcards: os.path.join(config["genome_dir_WGS"], f"{wildcards.genome}.fasta")
    output:
        gff=os.path.join(config["prokka_dir"], "{genome}", "{genome}.gff"),
        fna=os.path.join(config["prokka_dir"], "{genome}", "{genome}.fna")
    conda:
        "envs/prokka.yaml"
    shell:
        """
        mkdir -p $(dirname {output.gff})
        prokka --outdir $(dirname {output.gff}) --prefix {wildcards.genome} --force --norrna --notrna {input.fasta}
        """

rule run_barrnap:
    input:
        fasta=os.path.join(config["prokka_dir"], "{genome}", "{genome}.fna")
    output:
        fasta=os.path.join(config["barrnap_dir"], "{genome}_16S.fasta")
    conda:
        "envs/barrnap.yaml"
    shell:
        """
        mkdir -p $(dirname {output.fasta})
        barrnap {input.fasta} --outseq {output.fasta}
        """

rule filter_16S:
    input:
        fasta=os.path.join(config["barrnap_dir"], "{genome}_16S.fasta")
    output:
        filtered_fasta=os.path.join(config["barrnap_dir"], "{genome}_filtered_16S.fasta")
    shell:
        """
        awk '/^>/ {{if($0 ~ /16S_rRNA/) {{print; getline; print}} }}' {input.fasta} > {output.filtered_fasta}
        """

# ── BLAST ─────────────────────────────────────────────────────────────────────
rule setup_blastdb:
    output:
        os.path.join(config["blast_db_dir"], config["blast_db_name"] + ".nhr")
    conda:
        "envs/blast.yaml"
    shell:
        """
        mkdir -p {config[blast_db_dir]}
        export BLASTDB={config[blast_db_dir]}
        cd {config[blast_db_dir]}
        if [ ! -f {config[blast_db_name]}.nhr ]; then
            DB={config[blast_db_name]}
            HTTPS_URL="https://ftp.ncbi.nlm.nih.gov/blast/db/${{DB}}.tar.gz"
            echo "[NOSE] Downloading BLAST DB: ${{DB}}"
            # -c resumes partial downloads; --tries=5 retries on transient failures
            wget -c --tries=5 --timeout=300 --progress=bar:force \
                "${{HTTPS_URL}}" -O "${{DB}}.tar.gz" \
            || update_blastdb.pl --decompress "${{DB}}" \
            || {{ echo "[ERROR] BLAST DB download failed after 5 retries"; exit 1; }}
            echo "[NOSE] Extracting ${{DB}}.tar.gz"
            tar -xzf "${{DB}}.tar.gz"
            rm -f "${{DB}}.tar.gz"
        else
            echo "[NOSE] BLAST DB already present: {config[blast_db_name]}.nhr"
        fi
        """

rule blast:
    input:
        fasta=os.path.join(config["barrnap_dir"], "{genome}_filtered_16S.fasta"),
        db_installed=os.path.join(config["blast_db_dir"], config["blast_db_name"] + ".nhr")
    output:
        os.path.join(config["blast_dir"], "{genome}.blast.csv")
    params:
        db=os.path.join(config["blast_db_dir"], config["blast_db_name"])
    conda:
        "envs/blast.yaml"
    shell:
        """
        mkdir -p "{config[blast_dir]}"
        echo "query_id,subject_id,percentage_identity,alignment_length,mismatches,gap_opens,q_start,q_end,s_start,s_end,evalue,bit_score,taxid,sci_name" > "{output}"
        blastn -query "{input.fasta}" -task megablast -db "{params.db}" -word_size 28 -evalue 0.05 \
            -max_target_seqs 100 -outfmt "10 qseqid sseqid pident length mismatch gapopen qstart qend sstart send evalue bitscore" >> "{output}"
        """

# ── Taxonomy ──────────────────────────────────────────────────────────────────
rule download_taxdump:
    output:
        names_dmp_path=config["names_dmp_path"]
    run:
        if not os.path.exists(config["names_dmp_path"]):
            url      = "https://ftp.ncbi.nih.gov/pub/taxonomy/taxdump.tar.gz"
            zip_path = os.path.join(os.path.dirname(config["names_dmp_path"]), "taxdump.tar.gz")
            shell("wget -O {zip_path} {url}")
            shell("tar -xzf {zip_path} -C {os.path.dirname(config['names_dmp_path'])} names.dmp")
            shell("rm {zip_path}")

rule get_taxid:
    input:
        names_dmp=config["names_dmp_path"]
    output:
        taxid_file=os.path.join(config['taxid_dir'], "{genus}.taxid")
    shell:
        r"""
        grep -P '\t{wildcards.genus}\t' {input} | \
        grep -P '\tscientific name\t' | \
        awk -F'\t' '{{print $1}}' | sort -n > {output}
        if [ ! -s {output} ]; then
            echo "null" > {output}
        fi
        """

# ── NCBI dataset download ─────────────────────────────────────────────────────
rule get_genome_summary:
    input:
        taxid_file=f"{config['taxid_dir']}/{{genus}}.taxid"
    output:
        "{GCF_dir}/{genus}_species_contigs.txt"
    resources:
        ncbi_api=1
    conda:
        "envs/ncbi.yaml"
    shell:
        r"""
        mkdir -p $(dirname {output})
        if [ "$(cat {input.taxid_file})" == "null" ]; then
            echo "# No taxid found for {wildcards.genus}" > {output}
            exit 0
        fi
        ALL_TAXIDS=$(cat {input.taxid_file})
        VALID_TAXIDS_FILE="{output}.valid_taxids.tmp"
        > "$VALID_TAXIDS_FILE"
        for TAXID in $ALL_TAXIDS; do
            if datasets summary genome taxon "$TAXID" --from-type 2>/dev/null | jq -e '.reports and (.reports | length > 0)' >/dev/null; then
                echo "$TAXID" >> "$VALID_TAXIDS_FILE"
            fi
        done
        if [ ! -s "$VALID_TAXIDS_FILE" ]; then
            echo "# No genome data available for {wildcards.genus}" > {output}
            rm -f "$VALID_TAXIDS_FILE"
            exit 0
        fi
        BEST_TAXID=$(sort -n < "$VALID_TAXIDS_FILE" | tail -n 1)
        ATTEMPT=0; SUCCESS=false
        while [ $ATTEMPT -lt 5 ]; do
            datasets summary genome taxon "$BEST_TAXID" --from-type > {output}.tmp 2>/dev/null
            if [ -s {output}.tmp ] && grep -q "reports" {output}.tmp; then
                jq -r '.reports[] | "\(.accession)\t\(.organism.organism_name)\t\(.assembly_stats.number_of_contigs)"' {output}.tmp > {output}
                rm {output}.tmp; SUCCESS=true; break
            else
                ATTEMPT=$((ATTEMPT+1)); sleep $((ATTEMPT * 2))
            fi
        done
        if [ "$SUCCESS" = false ]; then
            echo "# NCBI retries failed for {wildcards.genus}" > {output}
            rm -f {output}.tmp
        fi
        rm -f "$VALID_TAXIDS_FILE"
        """

rule filter_gcf:
    input:
        "{GCF_dir}/{genus}_species_contigs.txt"
    output:
        "{GCF_dir}/{genus}_species_contigs_GCF.txt"
    conda:
        "envs/ncbi.yaml"
    shell:
        "grep '^GCF_' {input} > {output} || true"

rule select_final_gcf:
    input:
        gcf_list = "{GCF_dir}/{genus}_species_contigs_GCF.txt",
        lpsn_ref = config["lpsn_ref"]
    output:
        final_gcf = "{GCF_dir}/{genus}_final_GCF.txt"
    conda:
        "envs/ncbi.yaml"
    shell:
        r"""
        if [ ! -s {input.gcf_list} ]; then
            touch {output.final_gcf}
            exit 0
        fi
        awk '
        NR==FNR {{
            split($0, lpsn_parts, ",");
            genus = lpsn_parts[1]; species = lpsn_parts[2];
            if (genus != "" && species != "") {{
                valid_name = tolower(genus" "species);
                lpsn[valid_name] = 1;
            }}
            next;
        }}
        {{
            ncbi_species = tolower($2" "$3);
            gcf_id = $1; contigs = $NF;
            if (ncbi_species in lpsn) {{
                if (!(ncbi_species in seen) || contigs < seen[ncbi_species]) {{
                    seen[ncbi_species] = contigs;
                    best_gcf[ncbi_species] = gcf_id;
                }}
            }}
        }}
        END {{ for (s in best_gcf) print best_gcf[s]; }}
        ' {input.lpsn_ref} {input.gcf_list} > {output.final_gcf}
        """

# ── Genome downloads ──────────────────────────────────────────────────────────
rule download_genomes:
    input:
        gcf_file=f"{config['GCF_dir']}/{{genus}}_final_GCF.txt"
    output:
        genome_paths=f"{config['download_dir']}/{{genus}}/{{genus}}_genome_paths.txt"
    params:
        species_dir=f"{config['download_dir']}/{{genus}}"
    conda:
        "envs/ncbi.yaml"
    shell:
        r"""
        mkdir -p $(dirname {output.genome_paths})
        touch {output.genome_paths}
        if [ ! -s {input.gcf_file} ]; then exit 0; fi
        while read gcf_id; do
            RETRY=0; MAX=5; SUCCESS=false
            until $SUCCESS || [ $RETRY -ge $MAX ]; do
                datasets download genome accession $gcf_id \
                    --filename {params.species_dir}/$gcf_id.zip --include genome \
                    && SUCCESS=true || (RETRY=$((RETRY+1)); sleep 10)
            done
            if ! $SUCCESS; then echo "Failed: $gcf_id" >&2; exit 1; fi
            unzip -o {params.species_dir}/$gcf_id.zip -d {params.species_dir}/temp_extracted
            for fna in $(find {params.species_dir}/temp_extracted -name "*_genomic.fna"); do
                mv "$fna" {params.species_dir}/
                echo "{params.species_dir}/$(basename $fna)" >> {output.genome_paths}
            done
            rm -rf {params.species_dir}/temp_extracted {params.species_dir}/$gcf_id.zip
        done < {input.gcf_file}
        """

# ── FastANI ───────────────────────────────────────────────────────────────────
rule run_fastani:
    input:
        query_fasta=config['genome_dir_WGS'] + "/{assembly}.fasta",
        genome_paths=lambda wildcards: os.path.join(
            config['download_dir'],
            str(assembly_to_genus.get(wildcards.assembly, "UnknownGenus")),
            f"{assembly_to_genus.get(wildcards.assembly, 'UnknownGenus')}_genome_paths.txt"
        )
    output:
        fastani_output=config['fast_ani'] + "/{assembly}.txt"
    resources:
        fastani_jobs=1
    conda:
        "envs/fastani.yaml"
    shell:
        """
        if [ ! -s {input.genome_paths} ]; then
            touch {output.fastani_output}
        else
            fastANI -q {input.query_fasta} --rl {input.genome_paths} -o {output.fastani_output}
        fi
        """

rule get_query_ref_pairs_trigger:
    input:
        expand(f"{config['fast_ani']}/{{assembly}}.txt", assembly=assembly_list)
    output:
        touch(f"{config['output_dir']}/fastani_complete.txt")
    run:
        pass

rule process_fastani:
    input:
        ani_file=config['fast_ani'] + "/{query}.txt"
    output:
        genome_file=config['temp_files'] + "/{query}_matched.fna"
    shell:
        r"""
        mkdir -p "{config[temp_files]}"
        if [ ! -s "{input.ani_file}" ]; then
            touch "{output.genome_file}"; exit 0
        fi
        best_ref_info=$(awk 'BEGIN {{max=-1}} $3+0>max {{max=$3; split($2,a,"/"); path=$2; name=a[length(a)]}} END {{print path "|" name}}' "{input.ani_file}")
        best_ref_path=$(echo "$best_ref_info" | cut -d'|' -f1)
        best_ref_name=$(echo "$best_ref_info" | cut -d'|' -f2)
        if [ -z "$best_ref_path" ] || [ ! -f "$best_ref_path" ]; then
            touch "{output.genome_file}"
        else
            cp "$best_ref_path" "{config[temp_files]}/$best_ref_name"
            touch "{output.genome_file}"
        fi
        """

rule cleanup_matched_fna:
    input:
        expand(f"{config['temp_files']}/{{assembly}}_matched.fna", assembly=genome_summary_df['Assembly'])
    output:
        touch(f"{config['output_dir']}/cleanup_marker.txt")
    shell:
        "rm -f {input}"

rule process_remaining_fna:
    input:
        marker=f"{config['output_dir']}/cleanup_marker.txt"
    output:
        touch(f"{config['output_dir']}/fastani_processing_complete.txt")
    conda:
        "envs/prokka.yaml"
    shell:
        r"""
        FNA_FILES=$(find {config[temp_files]} -maxdepth 1 -name "*.fna" 2>/dev/null)
        if [ -z "$FNA_FILES" ]; then touch {output}; exit 0; fi
        for FNA_FILE in $FNA_FILES; do
            GENOME_NAME=$(basename "$FNA_FILE" .fna)
            mkdir -p {config[prokka_dir]}/$GENOME_NAME
            prokka --outdir {config[prokka_dir]}/$GENOME_NAME --prefix $GENOME_NAME --force "$FNA_FILE"
        done
        touch {output}
        """

rule trigger_aai_pocp:
    input:
        f"{config['output_dir']}/fastani_processing_complete.txt"
    output:
        touch(f"{config['output_dir']}/aai_pocp_completion_marker")
    run:
        pass

# ── AAI / POCP ────────────────────────────────────────────────────────────────
rule run_aai:
    input:
        marker=f"{config['output_dir']}/aai_pocp_completion_marker"
    output:
        aai_result=os.path.join(config["aai_dir"], "{query}_vs_{reference}_aai.txt")
    params:
        query_faa=lambda wildcards: os.path.join(config["prokka_dir"], wildcards.query, f"{wildcards.query}.faa"),
        ref_faa=lambda wildcards:   os.path.join(config["prokka_dir"], wildcards.reference, f"{wildcards.reference}.faa")
    conda:
        "envs/aai.yaml"
    shell:
        "ruby {config[aai_rb]} --seq1 {params.query_faa} --seq2 {params.ref_faa} > {output.aai_result}"

rule run_pocp:
    input:
        marker=f"{config['output_dir']}/aai_pocp_completion_marker"
    output:
        pocp_result=os.path.join(config["pocp_dir"], "{query}_vs_{reference}_pocp.txt")
    params:
        query_faa=lambda wildcards: os.path.join(config["prokka_dir"], wildcards.query, f"{wildcards.query}.faa"),
        ref_faa=lambda wildcards:   os.path.join(config["prokka_dir"], wildcards.reference, f"{wildcards.reference}.faa"),
        evalue="10",
        tmpdir=tempfile.mkdtemp(prefix="pocp_tmp_")
    conda:
        "envs/aai.yaml"
    shell:
        r"""
        WORKDIR=$(mktemp -d -p {params.tmpdir})
        cp {params.query_faa} $WORKDIR/query.faa
        cp {params.ref_faa}   $WORKDIR/ref.faa
        cd $WORKDIR
        bash {config[pocp_script]} query.faa ref.faa {params.evalue} > result.txt || exit 1
        mv result.txt {output.pocp_result}
        cd && rm -rf $WORKDIR
        """
