"""
Modular annotation pipeline with dynamic database support.

Copyright Richard Stöckl 2025.
Distributed under the Boost Software License, Version 1.0.
(See accompanying file LICENSE or copy at
https://www.boost.org/LICENSE_1_0.txt)
"""

import os
import json
import sys
from pathlib import Path

import jsonschema
import pandas as pd
import yaml
from snakemake.utils import validate
from snakemake.utils import min_version
from snakemake.logging import logger

# This is required so that Snakefile can import mapping_utils.py during workflow parsing.
SCRIPTS_DIR = Path(workflow.basedir) / "scripts"
if str(SCRIPTS_DIR) not in sys.path:
    sys.path.insert(0, str(SCRIPTS_DIR))

from mapping_utils import validate_mapping_file as validate_mapping_file_shared

# True on ARM64 hosts (Apple Silicon, aarch64 Linux); False on x86-64.
IS_ARM = os.uname().machine in ("arm64", "aarch64")
logger.info(f"Running on {'ARM64' if IS_ARM else 'x86-64'} architecture")

########## check minimum snakemake version ##########
min_version("9.0.1")


# =============================================================================
# Configuration Loading and Validation
# =============================================================================


def load_and_validate_config(config_path, schema_path):
    """Load a YAML config and validate against JSON schema."""
    with open(config_path, "r") as f:
        cfg = yaml.safe_load(f)
    with open(schema_path, "r") as f:
        schema = json.load(f)
    jsonschema.validate(cfg, schema)
    return cfg


########## load config and metadata sheets ##########


configfile: os.path.join(workflow.basedir, "../", "config", "config.yaml")


validate(
    config,
    schema=os.path.join(
        workflow.basedir, "../", "config", "schemas", "config.schema.json"
    ),
)

metadataFile = os.path.join(
    workflow.basedir, "../", "config", config["global"]["metadata"]
)
validate(
    metadataFile,
    schema=os.path.join(
        workflow.basedir, "../", "config", "schemas", "metadata.schema.json"
    ),
)

metadata = pd.read_csv(metadataFile, sep=",")
duplicate_sample_ids = metadata.loc[
    metadata["sample_id"].duplicated(keep=False), "sample_id"
].unique()
if len(duplicate_sample_ids):
    duplicate_ids = ", ".join(map(str, duplicate_sample_ids))
    raise ValueError(f"Duplicate sample_id values in metadata: {duplicate_ids}.")

metadata = metadata.set_index("sample_id", drop=False)
NAMES = metadata.sample_id.to_list()
path = metadata.path.to_dict()

if "input_type" in metadata.columns:
    input_type = metadata.input_type.fillna("dna").str.lower().to_dict()
else:
    input_type = {sample_id: "dna" for sample_id in NAMES}

invalid_input_types = [
    sid for sid, itype in input_type.items() if itype not in {"dna", "protein"}
]
if invalid_input_types:
    raise ValueError(
        f"Invalid input_type for samples: {invalid_input_types}. "
        "Allowed values are: dna, protein."
    )

LOGPATH = os.path.normpath(config["global"]["logs"])
INTERIMPATH = os.path.normpath(config["global"]["interim"])
RESULTPATH = os.path.normpath(config["global"]["results"])

PROKANOTA_VERBOSE = False
prokanota_cfg = config.get("prokanota", {})
if isinstance(prokanota_cfg, dict):
    prokanota_args = prokanota_cfg.get("args", {})
    if isinstance(prokanota_args, dict):
        PROKANOTA_VERBOSE = bool(
            prokanota_args.get("verbose", False)
            or prokanota_args.get("very_verbose", False)
        )
if not PROKANOTA_VERBOSE:
    PROKANOTA_VERBOSE = bool(config.get("prokanota_verbose", False))

LOG_REDIRECT = "2>&1 | tee {log:q}" if PROKANOTA_VERBOSE else "1> {log:q} 2>&1"

FEATURES_DEFAULTS = {
    "translation_table": 11,
    "run_rrna": True,
    "run_trna": True,
    "run_crispr": True,
    "minimum_gene_length": 90,
}


def resolve_feature_settings(cfg):
    """Resolve optional global feature settings with fail-fast validation."""
    user_features = cfg.get("features", {})
    resolved = {**FEATURES_DEFAULTS, **user_features}

    if not isinstance(resolved["translation_table"], int):
        raise ValueError("features.translation_table must be an integer.")
    if resolved["translation_table"] < 1 or resolved["translation_table"] > 25:
        raise ValueError("features.translation_table must be between 1 and 25.")

    if not isinstance(resolved["minimum_gene_length"], int):
        raise ValueError("features.minimum_gene_length must be an integer.")
    if resolved["minimum_gene_length"] < 1:
        raise ValueError("features.minimum_gene_length must be >= 1.")

    for key in ("run_rrna", "run_trna", "run_crispr"):
        if not isinstance(resolved[key], bool):
            raise ValueError(f"features.{key} must be a boolean.")

    return resolved


FEATURE_SETTINGS = resolve_feature_settings(config)


# =============================================================================
# Load and Validate Databases Config
# =============================================================================

DATABASES_CONFIG_PATH = os.path.join(
    workflow.basedir,
    "../",
    "config",
    config["global"].get("databases", "databases.yaml"),
)
DATABASES_SCHEMA_PATH = os.path.join(
    workflow.basedir, "../", "config", "schemas", "databases.schema.json"
)
databases_config = load_and_validate_config(
    DATABASES_CONFIG_PATH, DATABASES_SCHEMA_PATH
)


# =============================================================================
# Database Registry Helper Functions
# =============================================================================


def get_enabled_databases():
    """Return list of enabled databases sorted by order."""
    dbs = [db for db in databases_config["databases"] if db.get("enabled", True)]
    return sorted(dbs, key=lambda x: x.get("order", 100))


def get_db_by_name(name):
    """Get database config by name."""
    for db in databases_config["databases"]:
        if db["name"] == name:
            return db
    raise ValueError(f"Database '{name}' not found in config")


def get_columns_json(db_config):
    """Convert columns config to JSON string for script argument."""
    return json.dumps(db_config["columns"])


def build_mmseqs2_extra_args(db_config):
    """Build optional mmseqs2 arguments based on db config."""
    args = []
    sensitivity = db_config.get("mmseqs2_sensitivity")
    max_seqs = db_config.get("mmseqs2_max_seqs")
    min_seq_id = db_config.get("mmseqs2_min_seq_id")
    cov_mode = db_config.get("mmseqs2_cov_mode")
    coverage = db_config.get("mmseqs2_coverage")

    if sensitivity is not None:
        args.extend(["--sensitivity", str(sensitivity)])
    if max_seqs is not None:
        args.extend(["--max-seqs", str(max_seqs)])
    if min_seq_id is not None:
        args.extend(["--min-seq-id", str(min_seq_id)])
    if cov_mode is not None:
        args.extend(["--cov-mode", str(cov_mode)])
    if coverage is not None:
        args.extend(["--coverage", str(coverage)])

    return args


ENABLED_DBS = get_enabled_databases()
HAS_DATABASES = bool(ENABLED_DBS)
if not HAS_DATABASES:
    logger.warning(
        "No annotation databases are enabled (see config/databases.yaml). "
        "Continuing with feature prediction only."
    )
ENABLED_DB_NAMES = [db["name"] for db in ENABLED_DBS]


# =============================================================================
# Mapping File Validation
# =============================================================================


def validate_mapping_file(mapping_path):
    """Validate a database mapping TSV and log summary statistics."""
    mapping_path = Path(mapping_path)
    stats = validate_mapping_file_shared(mapping_path)
    logger.debug(
        f"Mapping file validated: {mapping_path} "
        f"(rows: {stats['row_count']}, sample accessions: {', '.join(stats['sample_accessions'])})"
    )


for db in ENABLED_DBS:
    validate_mapping_file(db["mapping_path"])


# =============================================================================
# Target Rule
# =============================================================================

FEATURE_TARGETS = [
    os.path.join(RESULTPATH, sample, "features", f"{sample}.faa") for sample in NAMES
]
ANNOTATION_TARGETS = (
    [
        os.path.join(RESULTPATH, sample, "annotation", f"{sample}_finalAnnotation.tsv")
        for sample in NAMES
    ]
    + [os.path.join(RESULTPATH, "common", "annotation", "finalAnnotation.tsv")]
    if HAS_DATABASES
    else []
)


rule all:
    input:
        FEATURE_TARGETS + ANNOTATION_TARGETS,


# =============================================================================
# Feature Prediction Rules
# =============================================================================
rule predict_features:
    input:
        fasta=lambda wildcards: path[wildcards.id],
    output:
        faa=os.path.join(RESULTPATH, "{id}", "features", "{id}.faa"),
        gff=os.path.join(RESULTPATH, "{id}", "features", "{id}.gff"),
        gbk=os.path.join(RESULTPATH, "{id}", "features", "{id}.gbk"),
        fna=os.path.join(RESULTPATH, "{id}", "features", "{id}.fna"),
        tsv=os.path.join(RESULTPATH, "{id}", "features", "{id}.tsv"),
        rna=os.path.join(RESULTPATH, "{id}", "features", "{id}_rna.tsv"),
        crispr=os.path.join(RESULTPATH, "{id}", "features", "{id}_crispr.tsv"),
        genome=os.path.join(RESULTPATH, "{id}", "features", "{id}.fasta"),
    log:
        os.path.join(LOGPATH, "{id}", "features", "{id}_predict_features.log"),
    conda:
        os.path.join(workflow.basedir, "envs", "features.yaml")
    threads: 4
    params:
        sample_id="{id}",
        translation_table=FEATURE_SETTINGS["translation_table"],
        minimum_gene_length=FEATURE_SETTINGS["minimum_gene_length"],
        rrna_flag=["--run_rrna"] if FEATURE_SETTINGS["run_rrna"] else [],
        trna_flag=["--run_trna"] if FEATURE_SETTINGS["run_trna"] else [],
        crispr_flag=["--run_crispr"] if FEATURE_SETTINGS["run_crispr"] else [],
        input_type=lambda wildcards: input_type[wildcards.id],
        scriptpath=os.path.join(workflow.basedir, "scripts", "features.py"),
    message:
        "Predicting features for {wildcards.id}"
    shell:
        """
        python {params.scriptpath:q} \
        {params.sample_id:q} {input.fasta:q} \
        --input_type {params.input_type:q} \
        --faa_path {output.faa:q} \
        --gff_path {output.gff:q} \
        --gbk_path {output.gbk:q} \
        --fna_path {output.fna:q} \
        --tsv_path {output.tsv:q} \
        --genome_path {output.genome:q} \
        {params.rrna_flag:q} {params.trna_flag:q} {params.crispr_flag:q} \
        --rna_tsv_path {output.rna:q} \
        --crispr_tsv_path {output.crispr:q} \
        --translation_table {params.translation_table} \
        --minimum_gene_length {params.minimum_gene_length} \
        --threads {threads} """ + LOG_REDIRECT


# =============================================================================
# Dynamic Search Rules
# =============================================================================

for db in ENABLED_DBS:
    db_name = db["name"]

    if db["search_tool"] == "pyhmmer":

        rule:
            name:
                f"search_{db_name}"
            input:
                faa=os.path.join(RESULTPATH, "{id}", "features", "{id}.faa"),
            output:
                tblout=temp(os.path.join(INTERIMPATH, "{id}", db_name, "hits.tblout")),
                version=os.path.join(INTERIMPATH, "{id}", db_name, "tool_version.txt"),
            log:
                os.path.join(LOGPATH, "{id}", "search", f"search_{db_name}.log"),
            conda:
                os.path.join(workflow.basedir, "envs", "pyhmmer.yaml")
            threads: max(1, int(workflow.cores * 0.5))  # Use safe threads calculation
            params:
                db_path=db["db_path"],
                db_name=db_name,
                evalue_cutoff=db.get("evalue_cutoff", 1e-3),
                scriptpath=os.path.join(
                    workflow.basedir, "scripts", "pyhmmer_search.py"
                ),
            message:
                f"Searching {{wildcards.id}} against {db_name}"
            shell:
                """
                python {params.scriptpath:q} \
                    --db {params.db_path:q} \
                    --db-name {params.db_name:q} \
                    --faa {input.faa:q} \
                    --tblout {output.tblout:q} \
                    --toolversion {output.version:q} \
                    --threads {threads} \
                    --evalue-cutoff {params.evalue_cutoff} \
                """ + LOG_REDIRECT

    elif db["search_tool"] == "rpsblast":

        rule:
            name:
                f"search_{db_name}"
            input:
                faa=os.path.join(RESULTPATH, "{id}", "features", "{id}.faa"),
            output:
                tab=temp(os.path.join(INTERIMPATH, "{id}", db_name, "hits.tsv")),
                version=os.path.join(INTERIMPATH, "{id}", db_name, "tool_version.txt"),
            log:
                os.path.join(LOGPATH, "{id}", "search", f"search_{db_name}.log"),
            conda:
                os.path.join(workflow.basedir, "envs", "rpsblast.yaml")
            threads: max(1, int(workflow.cores * 0.5))  # Use safe threads calculation
            params:
                db_path=db["db_path"],
                db_name=db_name,
                evalue_cutoff=db.get("evalue_cutoff", 1e-3),
                scriptpath=os.path.join(
                    workflow.basedir, "scripts", "rpsblast_search.py"
                ),
            message:
                f"Searching {{wildcards.id}} against {db_name}"
            shell:
                """
                python {params.scriptpath:q} \
                    --db {params.db_path:q} \
                    --db-name {params.db_name:q} \
                    --faa {input.faa:q} \
                    --output {output.tab:q} \
                    --toolversion {output.version:q} \
                    --threads {threads} \
                    --evalue {params.evalue_cutoff} \
                """ + LOG_REDIRECT

    elif db["search_tool"] == "diamond":

        rule:
            name:
                f"search_{db_name}"
            input:
                faa=os.path.join(RESULTPATH, "{id}", "features", "{id}.faa"),
            output:
                tab=temp(os.path.join(INTERIMPATH, "{id}", db_name, "hits.tsv")),
                version=os.path.join(INTERIMPATH, "{id}", db_name, "tool_version.txt"),
            log:
                os.path.join(LOGPATH, "{id}", "search", f"search_{db_name}.log"),
            conda:
                os.path.join(workflow.basedir, "envs", "diamond.yaml")
            threads: max(1, int(workflow.cores * 0.5))  # Use safe threads calculation
            params:
                db_path=db["db_path"],
                db_name=db_name,
                evalue_cutoff=db.get("evalue_cutoff", 1e-3),
                scriptpath=os.path.join(
                    workflow.basedir, "scripts", "diamond_search.py"
                ),
            message:
                f"Searching {{wildcards.id}} against {db_name}"
            shell:
                """
                python {params.scriptpath:q} \
                    --db {params.db_path:q} \
                    --db-name {params.db_name:q} \
                    --faa {input.faa:q} \
                    --output {output.tab:q} \
                    --toolversion {output.version:q} \
                    --threads {threads} \
                    --evalue {params.evalue_cutoff} \
                """ + LOG_REDIRECT

    elif db["search_tool"] == "mmseqs2":

        rule:
            name:
                f"search_{db_name}"
            input:
                faa=os.path.join(RESULTPATH, "{id}", "features", "{id}.faa"),
            output:
                tab=temp(os.path.join(INTERIMPATH, "{id}", db_name, "hits.tsv")),
                version=os.path.join(INTERIMPATH, "{id}", db_name, "tool_version.txt"),
            log:
                os.path.join(LOGPATH, "{id}", "search", f"search_{db_name}.log"),
            conda:
                os.path.join(workflow.basedir, "envs", "mmseqs2.yaml")
            threads: max(1, int(workflow.cores * 0.5))  # Use safe threads calculation
            params:
                db_path=db["db_path"],
                db_name=db_name,
                evalue_cutoff=db.get("evalue_cutoff", 1e-3),
                extra_args=build_mmseqs2_extra_args(db),
                scriptpath=os.path.join(
                    workflow.basedir, "scripts", "mmseqs2_search.py"
                ),
            message:
                f"Searching {{wildcards.id}} against {db_name}"
            shell:
                """
                python {params.scriptpath:q} \
                    --db {params.db_path:q} \
                    --db-name {params.db_name:q} \
                    --faa {input.faa:q} \
                    --output {output.tab:q} \
                    --toolversion {output.version:q} \
                    --threads {threads} \
                    --evalue {params.evalue_cutoff} \
                    {params.extra_args:q} \
                """ + LOG_REDIRECT


# =============================================================================
# Dynamic Parse Rules
# =============================================================================

for db in ENABLED_DBS:
    db_name = db["name"]

    if db["search_tool"] == "pyhmmer":

        rule:
            name:
                f"parse_{db_name}"
            input:
                tblout=os.path.join(INTERIMPATH, "{id}", db_name, "hits.tblout"),
                mapping=db["mapping_path"],
            output:
                tsv=os.path.join(INTERIMPATH, "{id}", db_name, "parsed_annotation.tsv"),
            log:
                os.path.join(LOGPATH, "{id}", "parse", f"parse_{db_name}.log"),
            conda:
                os.path.join(
                    workflow.basedir,
                    "envs",
                    "polars_arm.yaml" if IS_ARM else "polars.yaml",
                )
            threads: 1  # Parsing is single-threaded
            params:
                db_name=db_name,
                mapping_key=db.get("mapping_key", "accession"),
                evalue_cutoff=db.get("evalue_cutoff", 1e-3),
                columns_json=get_columns_json(db),
            message:
                f"Parsing {{wildcards.id}} {db_name} results"
            shell:
                """
                python {workflow.basedir:q}/scripts/pyhmmer_parse.py \
                    --tblout {input.tblout:q} \
                    --mapping {input.mapping:q} \
                    --output {output.tsv:q} \
                    --db-name {params.db_name:q} \
                    --mapping-key {params.mapping_key:q} \
                    --evalue-cutoff {params.evalue_cutoff} \
                    --columns {params.columns_json:q} \
                """ + LOG_REDIRECT

    elif db["search_tool"] == "rpsblast":

        rule:
            name:
                f"parse_{db_name}"
            input:
                rpsblast_output=os.path.join(INTERIMPATH, "{id}", db_name, "hits.tsv"),
                mapping=db["mapping_path"],
            output:
                tsv=os.path.join(INTERIMPATH, "{id}", db_name, "parsed_annotation.tsv"),
            log:
                os.path.join(LOGPATH, "{id}", "parse", f"parse_{db_name}.log"),
            conda:
                os.path.join(
                    workflow.basedir,
                    "envs",
                    "polars_arm.yaml" if IS_ARM else "polars.yaml",
                )
            threads: 1  # Parsing is single-threaded
            params:
                db_name=db_name,
                mapping_key=db.get("mapping_key", "accession"),
                evalue_cutoff=db.get("evalue_cutoff", 1e-3),
                columns_json=get_columns_json(db),
            message:
                f"Parsing {{wildcards.id}} {db_name} results"
            shell:
                """
                python {workflow.basedir:q}/scripts/rpsblast_parse.py \
                    --rpsblast-output {input.rpsblast_output:q} \
                    --mapping {input.mapping:q} \
                    --output {output.tsv:q} \
                    --db-name {params.db_name:q} \
                    --mapping-key {params.mapping_key:q} \
                    --evalue-cutoff {params.evalue_cutoff} \
                    --columns {params.columns_json:q} \
                """ + LOG_REDIRECT

    elif db["search_tool"] == "diamond":

        rule:
            name:
                f"parse_{db_name}"
            input:
                diamond_output=os.path.join(INTERIMPATH, "{id}", db_name, "hits.tsv"),
                mapping=db["mapping_path"],
            output:
                tsv=os.path.join(INTERIMPATH, "{id}", db_name, "parsed_annotation.tsv"),
            log:
                os.path.join(LOGPATH, "{id}", "parse", f"parse_{db_name}.log"),
            conda:
                os.path.join(
                    workflow.basedir,
                    "envs",
                    "polars_arm.yaml" if IS_ARM else "polars.yaml",
                )
            threads: 1  # Parsing is single-threaded
            params:
                db_name=db_name,
                mapping_key=db.get("mapping_key", "accession"),
                evalue_cutoff=db.get("evalue_cutoff", 1e-3),
                columns_json=get_columns_json(db),
            message:
                f"Parsing {{wildcards.id}} {db_name} results"
            shell:
                """
                python {workflow.basedir:q}/scripts/diamond_parse.py \
                    --diamond-output {input.diamond_output:q} \
                    --mapping {input.mapping:q} \
                    --output {output.tsv:q} \
                    --db-name {params.db_name:q} \
                    --mapping-key {params.mapping_key:q} \
                    --evalue-cutoff {params.evalue_cutoff} \
                    --columns {params.columns_json:q} \
                """ + LOG_REDIRECT

    elif db["search_tool"] == "mmseqs2":

        rule:
            name:
                f"parse_{db_name}"
            input:
                mmseqs_output=os.path.join(INTERIMPATH, "{id}", db_name, "hits.tsv"),
                mapping=db["mapping_path"],
            output:
                tsv=os.path.join(INTERIMPATH, "{id}", db_name, "parsed_annotation.tsv"),
            log:
                os.path.join(LOGPATH, "{id}", "parse", f"parse_{db_name}.log"),
            conda:
                os.path.join(
                    workflow.basedir,
                    "envs",
                    "polars_arm.yaml" if IS_ARM else "polars.yaml",
                )
            threads: 1  # Parsing is single-threaded
            params:
                db_name=db_name,
                mapping_key=db.get("mapping_key", "accession"),
                evalue_cutoff=db.get("evalue_cutoff", 1e-3),
                columns_json=get_columns_json(db),
            message:
                f"Parsing {{wildcards.id}} {db_name} results"
            shell:
                """
                python {workflow.basedir:q}/scripts/mmseqs2_parse.py \
                    --mmseqs-output {input.mmseqs_output:q} \
                    --mapping {input.mapping:q} \
                    --output {output.tsv:q} \
                    --db-name {params.db_name:q} \
                    --mapping-key {params.mapping_key:q} \
                    --evalue-cutoff {params.evalue_cutoff} \
                    --columns {params.columns_json:q} \
                """ + LOG_REDIRECT


# =============================================================================
# Merge Annotations
# =============================================================================


def get_all_parsed_annotations(wildcards):
    """Get paths to all parsed annotation files for a sample."""
    return [
        os.path.join(INTERIMPATH, wildcards.id, db["name"], "parsed_annotation.tsv")
        for db in ENABLED_DBS
    ]


def get_db_results_json(wildcards):
    """Build JSON array of database results for merge script."""
    results = []
    for db in ENABLED_DBS:
        results.append(
            {
                "name": db["name"],
                "path": os.path.join(
                    INTERIMPATH, wildcards.id, db["name"], "parsed_annotation.tsv"
                ),
                "order": db.get("order", 100),
            }
        )
    return json.dumps(results)


rule merge_sample_annotations:
    input:
        base_table=os.path.join(RESULTPATH, "{id}", "features", "{id}.tsv"),
        parsed=get_all_parsed_annotations,
    output:
        final=os.path.join(RESULTPATH, "{id}", "annotation", "{id}_finalAnnotation.tsv"),
    log:
        os.path.join(LOGPATH, "{id}", "finalize", "merge_annotations.log"),
    conda:
        os.path.join(
            workflow.basedir, "envs", "polars_arm.yaml" if IS_ARM else "polars.yaml"
        )
    threads: 1  # Merging is single-threaded
    params:
        db_results_json=get_db_results_json,
        scriptpath=os.path.join(workflow.basedir, "scripts", "merge_annotations.py"),
    message:
        "Merging annotations for {wildcards.id}"
    shell:
        """
        python {params.scriptpath:q} \
            --base-table {input.base_table:q} \
            --output {output.final:q} \
            --db-results {params.db_results_json:q} \
        """ + LOG_REDIRECT


rule collect_master_table:
    input:
        tables=expand(
            os.path.join(RESULTPATH, "{id}", "annotation", "{id}_finalAnnotation.tsv"),
            id=NAMES,
        ),
    output:
        final=os.path.join(RESULTPATH, "common", "annotation", "finalAnnotation.tsv"),
    log:
        os.path.join(LOGPATH, "collect_master_table.log"),
    threads: 1  # Merging is single-threaded
    message:
        "Collecting master annotation table"
    run:
        import pandas as pd

        dfs = [pd.read_csv(path, sep="\t") for path in input.tables]
        combined = pd.concat(dfs, ignore_index=True)
        combined.to_csv(output.final, sep="\t", index=False)
