"""Bundled Salmon bulk RNA-seq counts and QC workflow."""

import json
import shlex
from pathlib import Path


DOWNLOADS = {}


def resolve_input(value):
    if not value.startswith(("http://", "https://")):
        return value
    filename = Path(value).name
    DOWNLOADS[filename] = value
    return f".workflow-inputs/{filename}"


RAW_SAMPLES = config.get("samples") or config.get("rnaseq_salmon_samples", {})
REFERENCES = {
    key: resolve_input(value)
    if key in {"transcriptome_fasta", "annotation_gtf"} and value
    else value
    for key, value in config["references"].items()
}
SALMON_CONFIG = dict(config.get("salmon", {}))
if SALMON_CONFIG.get("decoys"):
    SALMON_CONFIG["decoys"] = resolve_input(SALMON_CONFIG["decoys"])
THREADS = int(config.get("threads", 4))
COMMANDS = config.get("commands", {})
FASTQC = COMMANDS.get("fastqc", "fastqc")
MULTIQC = COMMANDS.get("multiqc", "multiqc")
SALMON = COMMANDS.get("salmon", "salmon")
PYTHON = COMMANDS.get("python", "python")
AGGREGATION_SCRIPT = str(Path(workflow.basedir) / "scripts" / "aggregate_salmon_quant.py")
ANNOTATION = REFERENCES.get("annotation_gtf")


def _as_list(value):
    if not value:
        return []
    return [value] if isinstance(value, str) else list(value)


def _normalized_samples():
    normalized = {}
    for sample, metadata in RAW_SAMPLES.items():
        r1 = [resolve_input(value) for value in _as_list(metadata["r1"])]
        r2 = [resolve_input(value) for value in _as_list(metadata.get("r2"))]
        normalized[sample] = {
            **metadata,
            "layout": metadata.get("layout", "PE" if r2 else "SE"),
            "strandedness": metadata.get("strandedness", "unknown"),
            "salmon_libtype": metadata.get("salmon_libtype", "A"),
            "salmon_libtype_source": metadata.get(
                "salmon_libtype_source",
                "from_config" if "salmon_libtype" in metadata else "infer_from_salmon",
            ),
            "row_indices": metadata.get("row_indices", []),
            "r1": r1,
            "r2": r2,
        }
    return normalized


def _fastq_files():
    configured = config.get("fastq_files")
    if configured:
        return {
            unit: {
                **metadata,
                "path": resolve_input(metadata["path"]),
            }
            for unit, metadata in configured.items()
        }
    files = {}
    for sample, metadata in RNASEQ_SALMON.items():
        for read in ("r1", "r2"):
            for index, path in enumerate(metadata[read], start=1):
                files[f"{sample}__{read}_{index}"] = {
                    "sample": sample,
                    "read": read,
                    "path": path,
                }
    return files


RNASEQ_SALMON = _normalized_samples()
FASTQ_FILES = _fastq_files()


rule all:
    input:
        "fastqc/multiqc/multiqc_report.html",
        "rnaseq_salmon/multiqc/multiqc_report.html",
        "rnaseq_salmon/matrices/tpm.tsv",
        "rnaseq_salmon/matrices/num_reads.tsv",
        "rnaseq_salmon/matrices/effective_length.tsv",
        "rnaseq_salmon/matrices/samples.tsv",
        "rnaseq_salmon/matrices/tx2gene_coverage.json",


rule download_input:
    output:
        ".workflow-inputs/{filename}"
    params:
        url=lambda wildcards: DOWNLOADS[wildcards.filename]
    shell:
        "curl -fsSL {params.url:q} -o {output:q}"


rule config_snapshot:
    output:
        "provenance/config.resolved.json"
    params:
        config_state=lambda wildcards: json.dumps(config, sort_keys=True)
    run:
        resolved = dict(config)
        resolved["rnaseq_salmon_samples"] = RNASEQ_SALMON
        resolved["references"] = REFERENCES
        resolved["salmon"] = SALMON_CONFIG
        if config.get("fastq_files"):
            resolved["fastq_files"] = FASTQ_FILES
        destination = Path(output[0])
        destination.parent.mkdir(parents=True, exist_ok=True)
        destination.write_text(
            json.dumps(resolved, indent=2, sort_keys=True) + "\n",
            encoding="utf-8",
        )


rule fastqc_raw:
    input:
        lambda wildcards: FASTQ_FILES[wildcards.unit]["path"]
    output:
        directory("fastqc/raw/{unit}")
    threads: THREADS
    conda:
        "envs/salmon.yaml"
    shell:
        "mkdir -p {output:q} && {FASTQC:q} -t {threads} -o {output:q} {input:q}"


rule multiqc_fastq:
    input:
        expand("fastqc/raw/{unit}", unit=FASTQ_FILES)
    output:
        "fastqc/multiqc/multiqc_report.html"
    conda:
        "envs/salmon.yaml"
    shell:
        "mkdir -p fastqc/multiqc && "
        "{MULTIQC:q} --force --cl-config 'no_version_check: true' --no-megaqc-upload "
        "fastqc/raw -o fastqc/multiqc"


rule salmon_index:
    input:
        transcriptome=lambda wildcards: REFERENCES["transcriptome_fasta"],
        decoys=lambda wildcards: (
            [SALMON_CONFIG["decoys"]] if SALMON_CONFIG.get("decoys") else []
        ),
    output:
        directory("rnaseq_salmon/index")
    threads: THREADS
    params:
        kmer=lambda wildcards: int(SALMON_CONFIG.get("kmer", 31)),
        decoy_args=lambda wildcards, input: (
            f"--decoys {shlex.quote(str(input.decoys[0]))}" if input.decoys else ""
        ),
    conda:
        "envs/salmon.yaml"
    shell:
        "{SALMON:q} --no-version-check index -t {input.transcriptome:q} "
        "-i {output:q} -k {params.kmer} -p {threads} {params.decoy_args}"


rule salmon_quant:
    input:
        index="rnaseq_salmon/index",
        r1=lambda wildcards: RNASEQ_SALMON[wildcards.sample]["r1"],
        r2=lambda wildcards: RNASEQ_SALMON[wildcards.sample]["r2"],
    output:
        "rnaseq_salmon/quant/{sample}/quant.sf"
    threads: THREADS
    params:
        layout=lambda wildcards: RNASEQ_SALMON[wildcards.sample]["layout"],
        libtype=lambda wildcards: RNASEQ_SALMON[wildcards.sample]["salmon_libtype"],
        outdir=lambda wildcards: f"rnaseq_salmon/quant/{wildcards.sample}",
    conda:
        "envs/salmon.yaml"
    shell:
        r"""
        if [ {params.layout:q} = PE ]; then
          {SALMON:q} --no-version-check quant -i {input.index:q} -l {params.libtype:q} \
            -1 {input.r1:q} -2 {input.r2:q} \
            -p {threads} --validateMappings -o {params.outdir:q}
        else
          {SALMON:q} --no-version-check quant -i {input.index:q} -l {params.libtype:q} \
            -r {input.r1:q} \
            -p {threads} --validateMappings -o {params.outdir:q}
        fi
        """


rule multiqc_salmon:
    input:
        expand("rnaseq_salmon/quant/{sample}/quant.sf", sample=RNASEQ_SALMON)
    output:
        "rnaseq_salmon/multiqc/multiqc_report.html"
    conda:
        "envs/salmon.yaml"
    shell:
        "mkdir -p rnaseq_salmon/multiqc && "
        "{MULTIQC:q} --force --cl-config 'no_version_check: true' --no-megaqc-upload "
        "rnaseq_salmon/quant -o rnaseq_salmon/multiqc"


rule salmon_aggregate:
    input:
        quants=expand("rnaseq_salmon/quant/{sample}/quant.sf", sample=RNASEQ_SALMON),
        config_path="provenance/config.resolved.json",
        annotation=lambda wildcards: [ANNOTATION] if ANNOTATION else [],
    output:
        tpm="rnaseq_salmon/matrices/tpm.tsv",
        num_reads="rnaseq_salmon/matrices/num_reads.tsv",
        effective_length="rnaseq_salmon/matrices/effective_length.tsv",
        samples="rnaseq_salmon/matrices/samples.tsv",
        coverage="rnaseq_salmon/matrices/tx2gene_coverage.json",
    params:
        quant_args=lambda wildcards: " ".join(
            "--quant "
            + shlex.quote(f"{sample}=rnaseq_salmon/quant/{sample}/quant.sf")
            for sample in sorted(RNASEQ_SALMON)
        )
    conda:
        "envs/salmon.yaml"
    shell:
        "mkdir -p rnaseq_salmon/matrices && "
        "{PYTHON:q} {AGGREGATION_SCRIPT:q} "
        "--config {input.config_path:q} "
        "--outdir rnaseq_salmon/matrices "
        "{params.quant_args}"
