#!/usr/bin/env python3

import os
import sys
import subprocess

# Setup
# $ pixi run -e dev postinstall
# $ pixi run -e dev aviary build # or setup db folder as in README
# Run from base dir of repo
# $ pixi run -e dev python test/run_samples_at_cmr
#
# For now, need to make sure manually that all the mqsub'd jobs finish without error at the end.
# Randomly selected via Sandpiper
# Gut: SRR29134037, SRR6028613, SRR17283995
# Soil: SRR28059992, SRR30718229, SRR26545921
# Ocean: ERR599108, SRR13153254, ERR12716859

DATA_DIR = "/work/microbiome/aviary_module_benchmarking/data/test_data"
COMMIT = subprocess.Popen(["git", "rev-parse", "HEAD"], stdout=subprocess.PIPE, text=True).communicate()[0].strip()
OUTPUT_DIR = f"/work/microbiome/aviary_module_benchmarking/aviary_test_samples/{COMMIT}"
ASSEMBLERS = {
    "megahit": "--use-megahit",
    "metaspades": "",
}

SAMPLES_LIST = "test/test_samples.tsv"
samples_list = {}
samples_env = {}
with open(SAMPLES_LIST) as f:
    for line in f:
        sample = line.strip().split("\t")[0]
        samples_list[sample] = line.strip().split("\t")[1].split(",")
        samples_env[sample] = line.strip().split("\t")[2]

os.makedirs(OUTPUT_DIR, exist_ok=True)
for assembler in ASSEMBLERS.keys():
    OUTPUT_DIR_ASM = f"{OUTPUT_DIR}/{assembler}"
    os.makedirs(OUTPUT_DIR_ASM, exist_ok=True)

with open(f"{OUTPUT_DIR}/cmds.sh", "w") as cmd_file:
    for sample, cobinning in samples_list.items():
        for assembler, asm_flag in ASSEMBLERS.items():
            sample_output = f"{OUTPUT_DIR}/{assembler}/{sample}"
            bin_info = f"{sample_output}/bins/bin_info.tsv"
            if os.path.exists(bin_info):
                print(f"Skipping sample {sample} with {assembler}, already complete")
                continue

            log_file = f"{OUTPUT_DIR}/{assembler}/{sample}.log"
            forward_reads = [f"{DATA_DIR}/{sample}_1.fastq.gz" for sample in cobinning]
            reverse_reads = [f"{DATA_DIR}/{sample}_2.fastq.gz" for sample in cobinning]
            print(f"Preparing sample {sample} with {assembler}, outputting to {sample_output}")

            aviary_cmd = (
                f"pixi run --frozen --manifest-path aviary/pixi.toml -e dev "
                f"aviary recover {asm_flag} "
                f"-o {sample_output} "
                f"-1 {' '.join(forward_reads)} "
                f"-2 {' '.join(reverse_reads)} "
                f"--binning-only --request-gpu --strict "
                f"-n 32 -t 32 --local-cores 1 "
                f"-m 280 --coassemble no "
                f"--snakemake-profile aqua --cluster-retries 3 "
                f"&> {log_file}"
            )
            cmd_file.write(aviary_cmd + "\n")

process = subprocess.Popen(
    [
        "parallel", "--delay", "600", "-j", "18", "::::", f"{OUTPUT_DIR}/cmds.sh"
    ],
    stdout=subprocess.PIPE,
    stderr=subprocess.STDOUT,
    text=True
)
for line in process.stdout:
    print(line, end="")
process.wait()

if process.returncode != 0:
    print(f"Error: Aviary sample tests failed with return code {process.returncode}.")
    print(f"Check the log files in {OUTPUT_DIR} for details.")
    sys.exit(1)
else:
    print(f"Tests appear to have run successfully, but check the log files in {OUTPUT_DIR} for details.\n")
    print("All jobs completed successfully without FAILED lines in their output. Win.")

    with open(f"{OUTPUT_DIR}/summary.tsv", "w") as summary_file:
        summary_file.write("Env\tSample\tAssembler\tScore\tGenomes\n")
        for sample in samples_list:
            for assembler in ASSEMBLERS.keys():
                print(f"Processing results for sample {sample} with {assembler} in {OUTPUT_DIR}/{assembler}/{sample}")

                sample_score = 0
                sample_genomes = 0
                with open(f"{OUTPUT_DIR}/{assembler}/{sample}/bins/bin_info.tsv") as bin_info_file:
                    header = True
                    for line in bin_info_file:
                        columns = line.strip().split("\t")

                        if header:
                            comp_col = columns.index("Completeness")
                            cont_col = columns.index("Contamination")
                            header = False
                        else:
                            completeness = float(columns[comp_col])
                            contamination = float(columns[cont_col])
                            quality = completeness - 5 * contamination

                            if quality >= 50:
                                sample_score += quality
                                sample_genomes += 1

                print(f"Sample {sample} with {assembler} has {sample_genomes} genomes with total score {sample_score}")
                summary_file.write(f"{samples_env[sample]}\t{sample}\t{assembler}\t{int(sample_score)}\t{sample_genomes}\n")

    sys.exit(0)
