# ==============================================================================
# SNKEMAKE WORKFLOW: Genome-summary-and-classification-workflow (MODULE 1)
# ==============================================================================
# EXECUTION COMMAND:
# snakemake --use-conda --cores 20 --configfile config.yaml
#
# NOTE: 
# This Snakefile is automated. All project-specific changes (paths, tool 
# parameters, and eukaryotic/prokaryotic toggles) should be made in 
# the 'config.yaml' file only.
# ==============================================================================

# ============================
# Standard library imports
# ============================
import os
import yaml

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

# Path interpolation: replaces "{output_dir}" with the actual path
def fix_paths(d):
    for k, v in d.items():
        if isinstance(v, dict): fix_paths(v)
        elif isinstance(v, str): d[k] = v.replace("{output_dir}", config["output_dir"])

fix_paths(config)

# Global path variables
GENOME_DIR = config["genome_dir"]
OUTPUT_DIR = config["output_dir"]
EXT = config["assembly_extension"]

# ============================
# Helper function: get genome names
# ============================
def get_genomes():
    # Scans genome directory and strips extension for wildcard use
    return [f.replace(EXT, "") for f in os.listdir(GENOME_DIR) if f.endswith(EXT)]

GENOMES = get_genomes()

# ============================
# Master rule
# ============================
rule generate_genomic_summary:
    input:
        output_csv=os.path.join(OUTPUT_DIR, "genome_summary.csv")

# ======================================================================
# Eukaryotic Genome Workflow (QUAST + EukCC + CAT)
# ======================================================================
if config["is_euk"]:

    # ----------------------------
    # QUAST: Assembly statistics
    # ----------------------------
    rule quast_euk:
        input:
            genome=expand(os.path.join(GENOME_DIR, "{genome}" + EXT), genome=GENOMES)
        output:
            output=directory(config["quast"]["output"])
        conda:
            "./envs/quast.yaml"
        shell:
            "quast.py {input.genome} -o {output} --threads {config[quast][threads]}"

    # ----------------------------
    # EukCC: Completeness estimation
    # ----------------------------
    rule eukcc:
        input:
            genome=os.path.join(GENOME_DIR, "{genome}" + EXT)
        params:
            output_dir=os.path.join(config["eukcc"]["output"], "{genome}-eukcc-out"),
            eukcc_db_dir=config["eukcc"]["database"]
        output:
            est_tab=os.path.join(config["eukcc"]["output"], "{genome}-eukcc-estimates.tsv"),
            AA_seqs=os.path.join(config["eukcc"]["output"], "{genome}-prots.faa")
        conda:
            "./envs/eukcc.yaml"
        shell:
            """
            eukcc single --db {params.eukcc_db_dir} --threads {config[eukcc][threads]} --out {params.output_dir} --keep {input.genome}
            paste <( echo "{wildcards.genome}" ) <( tail -n +2 {params.output_dir}/eukcc.csv | head -n 1 | cut -f 2,3 ) > {output.est_tab}
            sed 's/metaeuk_//' {params.output_dir}/workdir/metaeuk/*_metaeuk_cleaned.faa > {output.AA_seqs}
            """

    # ----------------------------
    # Merge EukCC metrics
    # ----------------------------
    rule combine_eukcc_estimates:
        input:
            expand(os.path.join(config["eukcc"]["output"], "{genome}-eukcc-estimates.tsv"), genome=GENOMES)
        output:
            os.path.join(config["eukcc"]["output"], "combined-eukcc-estimates.tsv")
        shell:
            """
            printf "Assembly\\tEst. Comp.\\tEst. Redund.\\n" > {output}
            cat {input} >> {output}
            """

    # ----------------------------
    # CAT: Taxonomic classification
    # ----------------------------
    rule CAT:
        conda:
            "./envs/cat.yaml"
        input:
            genome=os.path.join(GENOME_DIR, "{genome}" + EXT),
            AA_seqs=os.path.join(config["eukcc"]["output"], "{genome}-prots.faa")
        params:
            tmp_out_prefix = os.path.join(config["cat"]["output"], "{genome}-tax-dir.tmp"),
            tmp_tax = os.path.join(config["cat"]["output"], "{genome}-tax.tmp"),
            cat_db = config["cat"]["CAT_DB"],
            cat_tax = config["cat"]["CAT_TAX"],
            num_threads = config["cat"]["threads"],
            assembly_extension = EXT
        output:
            tax=os.path.join(config["cat"]["output"], "{genome}-tax.tsv")
        shell:
            """
            CAT bin -b {input.genome} -p {input.AA_seqs} -d {params.cat_db} -t {params.cat_tax} -n {params.num_threads} -o {params.tmp_out_prefix} -f 0.5 -r 3
            CAT add_names -i {params.tmp_out_prefix}.bin2classification.txt -o {params.tmp_tax} -t {params.cat_tax} --only_official
            grep -v "^#" {params.tmp_tax} | awk -F $'\\t' ' BEGIN {{ OFS=FS }} {{ if ( $2 == "taxid assigned" ) {{ print $1,$6,$7,$8,$9,$10,$11,$12 }} else {{ print $1,"NA","NA","NA","NA","NA","NA","NA" }} }} ' | head -n 1 | sed 's/: [0-9\\.]*//g' | sed 's/not classified/NA/g' | sed 's/no support/NA/g' | sed 's/{params.assembly_extension}//' > {output}
            rm -rf {wildcards.genome}*tmp* {input.AA_seqs}
            """

    # ----------------------------
    # Merge CAT outputs
    # ----------------------------
    rule combine_euk_tax_outputs:
        input:
            expand(os.path.join(config["cat"]["output"], "{genome}-tax.tsv"), genome=GENOMES)
        output:
            os.path.join(config["cat"]["output"], "CAT-taxonomies.tsv")
        shell:
            """
            printf "Assembly\\tdomain\\tphylum\\tclass\\torder\\tfamily\\tgenus\\tspecies\\n" > {output}
            cat {input} >> {output}
            """

    # ----------------------------
    # Final Eukaryotic Summary
    # ----------------------------
    rule combine_euk_outputs:
        input:
            config["quast"]["output"],
            os.path.join(config["eukcc"]["output"], "combined-eukcc-estimates.tsv"),
            tax_tsv=os.path.join(config["cat"]["output"], "CAT-taxonomies.tsv")
        params:
            report_tsv=os.path.join(config["quast"]["output"], "report.tsv"),
            eukcc_tsv=os.path.join(config["eukcc"]["output"], "combined-eukcc-estimates.tsv"),
            tax_tsv=os.path.join(config["cat"]["output"], "CAT-taxonomies.tsv")
        output:
            output_csv=os.path.join(OUTPUT_DIR, "genome_summary.csv")
        shell:
            "python3 Eukaryotic_genome_summary.py --report_tsv {params.report_tsv} --eukcc_tsv {params.eukcc_tsv} --tax_tsv {params.tax_tsv} --output_csv {output.output_csv}"

# ======================================================================
# Prokaryotic Genome Workflow (QUAST + CheckM2 + GTDB-Tk)
# ======================================================================
else:

    # ----------------------------
    # QUAST: Assembly statistics
    # ----------------------------
    rule quast_prok:
        input:
            genome=expand(os.path.join(GENOME_DIR, "{genome}" + EXT), genome=GENOMES)
        output:
            output=directory(config["quast"]["output"])
        conda:
            "./envs/quast.yaml"
        shell:
            "quast.py {input.genome} -o {output} --threads {config[quast][threads]}"

    # ----------------------------
    # CheckM2: Quality assessment
    # ----------------------------
    rule checkm2:
        input:
            genome=expand(os.path.join(GENOME_DIR, "{genome}" + EXT), genome=GENOMES)
        output:
            output=directory(config["checkm2"]["output"])
        conda:
            "./envs/checkm2.yaml"
        shell:
            "checkm2 predict --input {input.genome} --threads {config[checkm2][threads]} --output-directory {output} --database_path {config[checkm2][database]}"

    # ----------------------------
    # GTDB-Tk: Taxonomic classification
    # ----------------------------
    rule gtdbtk:
        input:
            genome=expand(os.path.join(GENOME_DIR, "{genome}" + EXT), genome=GENOMES)
        output:
            output=directory(config["gtdbtk"]["output"])
        conda:
            "./envs/gtdbtk.yaml"
        shell:
            """
            export GTDBTK_DATA_PATH={config[gtdbtk][database]}
            gtdbtk classify_wf --genome_dir {GENOME_DIR} --out_dir {output} -x {EXT} --cpus {config[gtdbtk][cpus]} --skip_ani_screen
            """

    # ----------------------------
    # Final Prokaryotic Summary
    # ----------------------------
    rule genome_summary:
        input:
            config["quast"]["output"],
            config["checkm2"]["output"],
            config["gtdbtk"]["output"]
        params:
            report_tsv=os.path.join(config["quast"]["output"], "report.tsv"),
            quality_tsv=os.path.join(config["checkm2"]["output"], "quality_report.tsv"),
            gtdbtk_tsv=os.path.join(config["gtdbtk"]["output"], "classify", "gtdbtk.bac120.summary.tsv")
        output:
            output_csv=os.path.join(OUTPUT_DIR, "genome_summary.csv")
        shell:
            "python3 Prokaryotic_genome_summary.py --report_tsv {params.report_tsv} --quality_tsv {params.quality_tsv} --gtdbtk_tsv {params.gtdbtk_tsv} --output_csv {output.output_csv}"
