from glob import glob
from os import environ
from os.path import basename, dirname, join, splitext
import re


work_dir = config["workdir"]
subreads_dir = config["inputs"]["readsdir"]
scaffolds_dir = join(work_dir, "scaffolds")
subread_alignments_dir = join(work_dir, "alignments-subreads")
alignments_dir = join(work_dir, "alignments")
gap_windows_dir = join(work_dir, "gap-windows")
results_dir = join(work_dir, "results")
assembly_dir = join(work_dir, "assembly")


reference = config["inputs"]["reference"]
subreads_type = config["inputs"]["reads_type"]
reference_scaffold = join(scaffolds_dir, splitext(basename(reference))[0] + ".id_{scaffold_id}.fasta")
reference_windows = join(gap_windows_dir, "gap-windows.{scaffold_id}.lst")
alignments_by_scaffold = join(alignments_dir, "{scaffold_id}.bam")
result_by_scaffold = join(results_dir, "{scaffold_id}.fasta")
assembly_by_scaffold = join(assembly_dir, "{scaffold_id}.fasta")
result = config["outputs"]["assembly"]

seed_file = join(work_dir, "seed")
alignments_marker = join(work_dir, ".alignments.done")
arrow_marker = join(work_dir, ".arrow.done")
max_threads = 8


# Fix issues with HDF5 and lustre
environ["HDF5_USE_FILE_LOCKING"] = "FALSE"


def db_files(fasta):
    root = dirname(fasta)
    dbname = splitext(basename(fasta))[0]
    hidden_dam_file_suffixes = ["", ".bps", ".hdr", ".idx"]

    def __db_file(suffix):
        if len(suffix) > 0:
            hidden = "."
        else:
            hidden = ""
            suffix = ".dam"

        if root:
            return "{}/{}{}{}".format(root, hidden, dbname, suffix)
        else:
            return "{}{}{}".format(hidden, dbname, suffix)

    return [__db_file(suffix)  for suffix in hidden_dam_file_suffixes]


def scaffold_ids():
    from os import linesep

    with open(indexed_fasta(reference)[1]) as headers:
        return [line.partition("\t")[0]  for line in headers]


def reference_by_scaffold():
    return [reference_scaffold.replace("{scaffold_id}", scaffold_id, 1)  for scaffold_id in scaffold_ids()]


def results_by_scaffold():
    return [result_by_scaffold.replace("{scaffold_id}", scaffold_id, 1)  for scaffold_id in scaffold_ids()]


def name(filename, remove_ext=False):
    from os.path import basename, splitext

    if remove_ext:
        filename = splitext(filename)[0]

    return basename(filename)


def all_subread_alignments():
    subreads = glob(join(subreads_dir, "*" + subreads_type))

    return (subread.replace(subreads_dir, subread_alignments_dir, 1).replace(subreads_type, ".bam")  for subread in subreads)


def get_all_alignments_by_scaffold():
    return [alignments_by_scaffold.replace("{scaffold_id}", scaffold_id, 1)  for scaffold_id in scaffold_ids()]


def get_all_assemblies_by_scaffold():
    return [assembly_by_scaffold.replace("{scaffold_id}", scaffold_id, 1)  for scaffold_id in scaffold_ids()]


def indexed_fasta(fasta):
    return [fasta, fasta + ".fai"]


def log(id):
    from os.path import join

    return join(config["logdir"], id + ".log")


import csv
import re


class ResultAssembler:

    def __init__(self, inputs, output):
        from os.path import basename, splitext

        self.ref_fasta = inputs["reference"][0]
        self.ref_index = inputs["reference"][1]
        self.result_fasta = inputs["result"]
        self.assembly_fasta = output
        self.scaffold = splitext(basename(self.result_fasta))[0]


    def run(self):
        self.scaffold_length = self.get_scaffold_length()
        self.ref_sequence = self.read_reference_scaffold()
        self.insertion_infos, self.insertions = self.read_insertions()

        with open(self.assembly_fasta, "w") as assembly_file:
            self.write_header(assembly_file)
            self.assemble_arrow_results(assembly_file)
            self.copy_unmodified_suffix(assembly_file)


    def get_scaffold_length(self):
        with open(self.ref_index) as index_file:
            index_reader = csv.reader(index_file, delimiter='\t')
            for scaffold_info in index_reader:
                if scaffold_info[0] == self.scaffold:
                    self.scaffold_length = int(scaffold_info[1])


    def read_reference_scaffold(self):
        from subprocess import check_output
        from shlex import quote

        scaffold = quote(self.scaffold)
        fasta = quote(self.ref_fasta)
        get_sequence_cmd = "seqkit grep --line-width=0 -np {} {} | tail -n+2".format(scaffold, fasta)

        return check_output(get_sequence_cmd, shell=True, encoding="ascii")


    def read_insertions(self):
        from subprocess import PIPE, Popen
        from shlex import quote

        get_sequence_cmd = ["seqkit", "seq", "--line-width=0", self.result_fasta]
        sequence_process = Popen(get_sequence_cmd, stdout=PIPE, bufsize=1, encoding="ascii")

        insertion_infos = list()
        insertions = list()

        expect_header = True
        for line in sequence_process.stdout:
            if expect_header:
                if line[0] != ">":
                    raise Exception("expecting header!")

                insertion_info = self.parse_arrow_header(line)
                insertion_infos.append(insertion_info)
                expect_header = False
            else:
                if line[0] == ">":
                    raise Exception("expecting sequence!")

                insertions.append(line.strip())
                expect_header = True

        return insertion_infos, insertions


    def write_header(self, assembly_file):
        assembly_file.write(">{}\n".format(self.scaffold))


    def copy_unmodified_suffix(self, assembly_file):
        begin = 0

        if len(self.insertion_infos) > 0:
            begin = self.insertion_infos[-1].end

        self.copy_reference_region(assembly_file, begin, self.scaffold_length)


    def assemble_arrow_results(self, assembly_file):
        current_pos = 0
        for insertion_info, insertion in zip(self.insertion_infos, self.insertions):
            self.copy_reference_region(assembly_file, current_pos, insertion_info.begin)
            assembly_file.write(insertion)
            current_pos = insertion_info.end


    def copy_reference_region(self, assembly_file, begin, end):
        assembly_file.write(self.ref_sequence[begin:end])


    # Example: translocated_gaps_504_31700_36617|quiver
    arrow_header_format = re.compile(r'^>?(?P<scaffold>translocated_gaps_[0-9]+)_(?P<begin>[0-9]+)_(?P<end>[0-9]+)\|(?P<algorithm>quiver|arrow)')


    class InsertionInfo:
        def __init__(self, scaffold, begin, end, algorithm):
            self.scaffold = str(scaffold)
            self.begin = int(begin) - 1
            self.end = int(end)
            self.algorithm = str(algorithm)


        def __str__(self):
            return "{}_{}_{}|{}".format(self.scaffold, self.begin, self.end, self.algorithm)


    def parse_arrow_header(self, header):
        match = self.arrow_header_format.search(header)

        if not match:
            raise Exception("ill-formatted arrow header: " + header)

        return self.InsertionInfo(match["scaffold"], match["begin"], match["end"], match["algorithm"])


#-----------------------------------------------------------------------------
# BEGIN Rules
#-----------------------------------------------------------------------------


localrules:
    ALL,
    generate_seed,
    split_reference,
    faindex,
    pbindex,
    reference_windows,
    all_alignments,
    arrow,
    assemble_result_scaffold,
    assemble_result


rule ALL:
    input:
        result


rule generate_seed:
    output: seed_file
    run:
        from random import getrandbits

        with open(output[0], 'w') as seed:
            seed.write(str(getrandbits(64)))


rule split_reference:
    input: reference
    output: reference_by_scaffold()
    shell:
        "rmdir {scaffolds_dir} ; seqkit split --by-id --out-dir={scaffolds_dir} {input}"


rule faindex:
    input: "{name}.fasta"
    output: "{name}.fasta.fai"
    shell:
        "samtools faidx {input}"


rule pbindex:
    input: "{name}.bam"
    output: "{name}.bam.pbi"
    shell:
        "pbindex {input}"


rule pbalign:
    input:
        ref=reference,
        reads=join(subreads_dir, "{name}" + subreads_type),
        seed=seed_file
    output: temp(join(subread_alignments_dir, "{name}.bam"))
    log: log("pbalign.{name}")
    threads: max_threads
    shell:
        "pbalign --log-file={log} --nproc={threads} --seed=$(< {input[seed]}) {input[reads]} {input[ref]} {output} &>> {log}"


rule reference_windows:
    input:
        *db_files(reference),
        indexed_fasta(reference)
    output: reference_windows
    shell:
        "DBshow -un {input[0]} | grep -E '^>{wildcards.scaffold_id}\s' | awk -f mk-gap-windows.awk > {output}"


rule alignments_by_scaffold:
    input: all_subread_alignments()
    output:
        alignments_by_scaffold
    params:
        additional_threads = lambda _, threads: threads - 1
    threads: 8
    shell: """
        samtools merge -@ {params.additional_threads} -R {wildcards.scaffold_id} {output} {input}
        samtools view -h {output} | \\
            sed -E 's/(@RG.*ID:|RG:Z:)([a-f0-9]{{1,8}})-[A-F0-9]{{1,8}}/\\1\\2/' | \\
            awk -F'\\t' '($1 == "@RG") {{ if (!visited[$2]) {{ visited[$2] = 1; print }} }} ($1 != "@RG") {{ print }}' | \\
            samtools view -b -o {output}.fixed.bam
        mv {output}.fixed.bam {output}
    """


rule all_alignments:
    input:
        get_all_alignments_by_scaffold()
    output:
        touch(alignments_marker)


rule arrow_by_scaffold:
    input:
        *indexed_fasta(reference_scaffold),
        alignments_by_scaffold + ".pbi",
        ref_windows=reference_windows,
        alignment=alignments_by_scaffold
    output:
        result_by_scaffold
    log: log("arrow.{scaffold_id}")
    threads: max_threads
    shell:
        "if (( $(wc -l < {input[ref_windows]}) > 0 )); then variantCaller --algorithm=best --debug -j {threads} {input[alignment]} --referenceWindowsFile={input[ref_windows]} --reference={input[0]} -o {output}; else touch {output}; fi &> {log}"


rule arrow:
    input:
        *indexed_fasta(reference),
        results_by_scaffold()
    output:
        touch(arrow_marker)


rule assemble_result_scaffold:
    input:
        reference=indexed_fasta(reference_scaffold),
        result=result_by_scaffold
    output:
        assembly=assembly_by_scaffold
    log: log("assembly.{scaffold_id}")
    run:
        ResultAssembler(input, output["assembly"]).run()


rule assemble_result:
    input:
        get_all_assemblies_by_scaffold()
    output:
        result
    log: log("assembly")
    shell:
        "cat {input} > {output}"
