import pandas as pd
import pysam

# read config files
configfile: workflow.source_path("../config/config.yaml")


########################################
# define regions (chunks) of fasta file
########################################

def gen_ref_regions(fasta_file, chunk_size, read_length):
    fastafh = pysam.FastaFile(fasta_file)
    regions = []
    chunk_id = 0
    for ref_index in range(fastafh.nreferences):
        ref_name = fastafh.references[ref_index]
        ref_length = fastafh.lengths[ref_index]
        # chunkify reference
        for chunk_start in range(0, ref_length, chunk_size):
            chunk_end = min(chunk_start + chunk_size + read_length, ref_length)
            regions.append((chunk_id, ref_name, chunk_start, chunk_end))
            chunk_id += 1
    df = pd.DataFrame(regions, columns=['id', 'ref', 'start', 'end'])
    df.set_index('id', inplace=True, drop=False)
    return df

# regions (chunks) of input fasta files to process in parallel
GENOME_REGIONS = gen_ref_regions(config['genome_fasta'], config['chunk_size'], config['read_length'])

###########################################################
# generate genome reads
###########################################################

OUTPUT_GEN_GENOME_READS_PIPE = 'genome{i}/gen_reads_pipe.fa'

rule gen_genome_reads:
    input:
        fasta = config['genome_fasta']
    params:
        ref = lambda wildcards: GENOME_REGIONS.loc[int(wildcards.i)].ref,
        start = lambda wildcards: GENOME_REGIONS.loc[int(wildcards.i)].start,
        end = lambda wildcards: GENOME_REGIONS.loc[int(wildcards.i)].end,
        read_length = config['read_length'],
        coverage = config['coverage'],
        error_rate = config['error_rate']
    output:
        pipe(OUTPUT_GEN_GENOME_READS_PIPE)
    script:
        'scripts/gen_reads.py'


###########################################################
# star load index
###########################################################

OUTPUT_STAR_LOAD_INDEX = 'star_load_index_to_shmem'
PARAM_STAR_LOAD_INDEX_PREFIX = 'star/'

rule star_load_index:
    input:
        genome_dir = config['star_index_dir']
    output:
        touch(OUTPUT_STAR_LOAD_INDEX)
    params:
        outFileNamePrefix = PARAM_STAR_LOAD_INDEX_PREFIX
    shell:
        'STAR --genomeLoad LoadAndExit '
        '--genomeDir {input.genome_dir} '
        '--outFileNamePrefix {params.outFileNamePrefix} '

###########################################################
# star align reads to genome
###########################################################

OUTPUT_STAR_GENOME_ALIGNED = 'genome{i}/star/Aligned.out.bam'
PARAM_STAR_GENOME_PREFIX = 'genome{i}/star/'

rule star_genome:
    input:
        load_index = rules.star_load_index.output,
        fasta = rules.gen_genome_reads.output[0],
        genome_dir = config['star_index_dir']
    output:
        aligned = pipe(OUTPUT_STAR_GENOME_ALIGNED)
    params:
        outFileNamePrefix = PARAM_STAR_GENOME_PREFIX,
        MultimapNMax = 100        
    threads: 2
    shell:
        'STAR --runThreadN {threads} '
        '--genomeDir {input.genome_dir} '
        '--genomeLoad LoadAndKeep '
#        '--genomeLoad NoSharedMemory '
        '--readFilesIn {input.fasta} '
        '--outFileNamePrefix {params.outFileNamePrefix} '
        '--outSAMtype BAM Unsorted '
        '--outSAMmode NoQS '
        '--outSAMstrandField intronMotif '
        '--outSAMattributes NH HI AS NM MD MC '
        '--outSAMunmapped Within KeepPairs '
        '--outSAMorder Paired '
        '--outSAMmultNmax 1 '
        '--outBAMcompression 0 '
        '--outFilterMultimapNmax {params.MultimapNMax} '
        '--winAnchorMultimapNmax {params.MultimapNMax} '
        '--peOverlapNbasesMin 10 '
        '--alignSplicedMateMapLminOverLmate 0.5 '
        '--alignSJstitchMismatchNmax 5 -1 5 5 '
        '--chimSegmentMin 10 '
        '--chimOutType WithinBAM HardClip '
        '--chimJunctionOverhangMin 10 '
        '--chimScoreDropMax 30 '
        '--chimScoreJunctionNonGTAG 0 '
        '--chimScoreSeparation 1 '
        '--chimSegmentReadGapMax 3 '
        '--chimMultimapNmax {params.MultimapNMax}'


###########################################################
# process star genome output
###########################################################

OUTPUT_GENOME_REGION_H5 = 'genome{i}/mapping.h5'

rule process_genome_alignments:
    input:
        fasta = config['genome_fasta'],
        aligned = rules.star_genome.output.aligned
    output:
        h5 = OUTPUT_GENOME_REGION_H5
    script:
        'scripts/proc_aln.py'

###########################################################
# genome reduce/aggregate regions
###########################################################

OUTPUT_GENOME_H5 = 'genome.h5'

rule aggregate_genome:
    input:
        fasta = config['genome_fasta'],
        h5_list = expand(OUTPUT_GENOME_REGION_H5, i = GENOME_REGIONS.id.unique())
    output:
        h5 = OUTPUT_GENOME_H5
    script:
        'scripts/aggregate_h5.py'


########################################
# define groups / chunks of transcripts
########################################

# def gen_transcript_chunks(fasta_file, chunk_size):
#     # open fasta file
#     fastafh = pysam.FastaFile(fasta_file)
#     chunks = []
#     for chunk_id, i in enumerate(range(0, fastafh.nreferences, chunk_size)):
#         chunks.append((chunk_id, fastafh.references[i:i + chunk_size]))
#     return chunks

# rule map_transcript_fasta:
#     input:
#         fasta = config['transcript_fasta']
#     output:
#         OUTPUT_TRANSCRIPT_CHUNK
#     params:
#         chunk_size = 100
#     run:
#         pass

###########################################################
# generate transcript reads
###########################################################

OUTPUT_GEN_TRANSCRIPT_READS_PIPE = 'transcript/gen_reads_pipe.fa'

rule gen_transcript_reads:
    input:
        fasta = config['transcript_fasta']
    params:
        ref = None,
        start = None,
        end = None,
        read_length = config['read_length'],
        coverage = config['coverage'],
        error_rate = config['error_rate']
    output:
        # OUTPUT_GEN_TRANSCRIPT_READS_PIPE
        pipe(OUTPUT_GEN_TRANSCRIPT_READS_PIPE)
    script:
        'scripts/gen_reads.py'


###########################################################
# star align reads to transcripts
###########################################################

OUTPUT_STAR_TRANSCRIPT_ALIGNED = 'transcript/star/Aligned.out.bam'
PARAM_STAR_TRANSCRIPT_PREFIX = 'transcript/star/'

rule star_transcript:
    input:
        load_index = rules.star_load_index.output,
        fasta = rules.gen_transcript_reads.output[0],
        genome_dir = config['star_index_dir']
    output:
        aligned = pipe(OUTPUT_STAR_TRANSCRIPT_ALIGNED)
    params:
        outFileNamePrefix = PARAM_STAR_TRANSCRIPT_PREFIX,
        MultimapNMax = 100        
    threads: 2
    shell:
        'STAR --runThreadN {threads} '
        '--genomeDir {input.genome_dir} '
        '--genomeLoad LoadAndKeep '
#        '--genomeLoad NoSharedMemory '
        '--readFilesIn {input.fasta} '
        '--outFileNamePrefix {params.outFileNamePrefix} '
        '--outSAMtype BAM Unsorted '
        '--outSAMmode NoQS '
        '--outSAMstrandField intronMotif '
        '--outSAMattributes NH HI AS NM MD MC '
        '--outSAMunmapped Within KeepPairs '
        '--outSAMorder Paired '
        '--outSAMmultNmax 1 '
        '--outBAMcompression 0 '
        '--outFilterMultimapNmax {params.MultimapNMax} '
        '--winAnchorMultimapNmax {params.MultimapNMax} '
        '--peOverlapNbasesMin 10 '
        '--alignSplicedMateMapLminOverLmate 0.5 '
        '--alignSJstitchMismatchNmax 5 -1 5 5 '
        '--chimSegmentMin 10 '
        '--chimOutType WithinBAM HardClip '
        '--chimJunctionOverhangMin 10 '
        '--chimScoreDropMax 30 '
        '--chimScoreJunctionNonGTAG 0 '
        '--chimScoreSeparation 1 '
        '--chimSegmentReadGapMax 3 '
        '--chimMultimapNmax {params.MultimapNMax} '

###########################################################
# process star transcript output
###########################################################

OUTPUT_TRANSCRIPT_H5 = 'transcript.h5'

rule process_transcript_alignments:
    input:
        gtf = config['transcript_gtf'],
        fasta = config['transcript_fasta'],
        aligned = rules.star_transcript.output.aligned
    output:
        h5 = OUTPUT_TRANSCRIPT_H5
    script:
        'scripts/proc_aln_transcript.py'


###########################################################
# cleanup star genome
###########################################################

OUTPUT_STAR_REMOVE_INDEX = 'star_remove_index_from_shmem'

def star_remove_index_input(wildcards):
    input = {'genome_dir': config['star_index_dir']}
    if config['genome_run']:
        input['genome_h5'] = rules.aggregate_genome.output.h5
    if config['transcript_run']:
        input['transcript_h5'] = rules.process_transcript_alignments.output.h5
    return input

rule star_remove_index:
    input:
        unpack(star_remove_index_input)
    output:
        touch(OUTPUT_STAR_REMOVE_INDEX)
    params:
        outFileNamePrefix = PARAM_STAR_LOAD_INDEX_PREFIX
    shell:
        'STAR --genomeLoad Remove '
        '--genomeDir {input.genome_dir} '
        '--outFileNamePrefix {params.outFileNamePrefix} '

###########################################################
# default rule
###########################################################

def get_main_input(wildcards):
    input = [
        OUTPUT_STAR_REMOVE_INDEX
    ]
    if config['genome_run']:
        input.append(OUTPUT_GENOME_H5)
    if config['transcript_run']:
        input.append(OUTPUT_TRANSCRIPT_H5)
    return input

        
rule all:
    default_target: True
    input:
        get_main_input