#!/usr/bin/env python

''' 
 * All rights Reserved, Designed By HIT-Bioinformatics   
 * @Title: cuteSVTrio 
 * @author: Lixin
 * @date: Apr. 19th 2025
 * @version V0.1.0
'''


from cuteSVTrio.cuteSVTrio_Description import parseArgs
from cuteSVTrio.CommandRunner import *
from cuteSVTrio.cuteSVTrio_resolveINV import run_inv
from cuteSVTrio.cuteSVTrio_resolveTRA import run_tra
from cuteSVTrio.cuteSVTrio_resolveINDEL import run_ins, run_del
from cuteSVTrio.cuteSVTrio_resolveDUP import run_dup
from cuteSVTrio.cuteSVTrio_genotype import generate_output, generate_pvcf, load_valuable_chr, load_bed, Generation_VCF_header, allele_correction
from cuteSVTrio.cuteSVTrio_phasing import run_phasing,genetic_phasing_family,genetic_no_phasing_family, generate_haplotype_read_names
from cuteSVTrio.cuteSVTrio_assembly import split_reference_chromosomes, remove_redundant_samesv, run_assembly, remove_redundant_pos
import pysam
import cigar
from multiprocessing import Pool,Manager,Queue, current_process
#from cuteSVTrio.cuteSVTrio_mendel import resolution_INDEL_mendel
import os
import argparse
import logging
import sys
import time
import gc
import pickle
import atexit
import math

dic_starnd = {1: '+', 2: '-'}
RefChangeOp=set([0,2,7,8])
signal = {1 << 2: 0, \
            1 >> 1: 1, \
            1 << 4: 2, \
            1 << 11: 3, \
            1 << 4 | 1 << 11: 4}
'''
    1 >> 1 means normal_foward read
    1 << 2 means unmapped read
    1 << 4 means reverse_complement read
    1 << 11 means supplementary alignment read
    1 << 4 | 1 << 11 means supplementary alignment with reverse_complement read
'''
def detect_flag(Flag):
    back_sig = signal[Flag] if Flag in signal else 0
    return back_sig

def analysis_inv(ele_1, ele_2, read_name, candidate, SV_size):
    if ele_1[5] == '+':
        # +-
        if ele_1[3] - ele_2[3] >= SV_size:
            if ele_2[0] + 0.5 * (ele_1[3] - ele_2[3]) >= ele_1[1]:
                candidate.append(("++", 
                                    ele_2[3], 
                                    ele_1[3], 
                                    read_name,
                                    "INV",
                                    ele_1[4]))
                # head-to-head
                # 5'->5'
        if ele_2[3] - ele_1[3] >= SV_size:
            if ele_2[0] + 0.5 * (ele_2[3] - ele_1[3]) >= ele_1[1]:
                candidate.append(("++", 
                                    ele_1[3], 
                                    ele_2[3], 
                                    read_name,
                                    "INV",
                                    ele_1[4]))
                # head-to-head
                # 5'->5'
    else:
        # -+
        if ele_2[2] - ele_1[2] >= SV_size:
            if ele_2[0] + 0.5 * (ele_2[2] - ele_1[2]) >= ele_1[1]:
                candidate.append(("--", 
                                    ele_1[2], 
                                    ele_2[2], 
                                    read_name,
                                    "INV",
                                    ele_1[4]))
                # tail-to-tail
                # 3'->3'
        if ele_1[2] - ele_2[2] >= SV_size:
            if ele_2[0] + 0.5 * (ele_1[2] - ele_2[2]) >= ele_1[1]:
                candidate.append(("--", 
                                    ele_2[2], 
                                    ele_1[2], 
                                    read_name,
                                    "INV",
                                    ele_1[4]))
                # tail-to-tail
                # 3'->3'


def analysis_bnd(ele_1, ele_2, read_name, candidate):
    '''
    *********Description*********
    *	TYPE A:		N[chr:pos[	*
    *	TYPE B:		N]chr:pos]	*
    *	TYPE C:		[chr:pos[N	*
    *	TYPE D:		]chr:pos]N	*
    *****************************
    '''
    if ele_2[0] - ele_1[1] <= 100:
        if ele_1[5] == '+':
            if ele_2[5] == '+':
                # +&+
                if ele_1[4] < ele_2[4]:
                    candidate.append(('A', 
                                        ele_1[3], 
                                        ele_2[4], 
                                        ele_2[2], 
                                        read_name,
                                        "TRA",
                                        ele_1[4]))
                    # N[chr:pos[
                else:
                    candidate.append(('D', 
                                        ele_2[2], 
                                        ele_1[4], 
                                        ele_1[3], 
                                        read_name,
                                        "TRA",
                                        ele_2[4]))
                    # ]chr:pos]N
            else:
                # +&-
                if ele_1[4] < ele_2[4]:
                    candidate.append(('B', 
                                        ele_1[3], 
                                        ele_2[4], 
                                        ele_2[3], 
                                        read_name,
                                        "TRA",
                                        ele_1[4]))
                    # N]chr:pos]
                else:
                    candidate.append(('B', 
                                        ele_2[3], 
                                        ele_1[4], 
                                        ele_1[3], 
                                        read_name,
                                        "TRA",
                                        ele_2[4]))
                    # N]chr:pos]
        else:
            if ele_2[5] == '+':
                # -&+
                if ele_1[4] < ele_2[4]:
                    candidate.append(('C', 
                                        ele_1[2], 
                                        ele_2[4], 
                                        ele_2[2], 
                                        read_name,
                                        "TRA",
                                        ele_1[4]))
                    # [chr:pos[N
                else:
                    candidate.append(('C', 
                                        ele_2[2], 
                                        ele_1[4], 
                                        ele_1[2], 
                                        read_name,
                                        "TRA",
                                        ele_2[4]))
                    # [chr:pos[N
            else:
                # -&-
                if ele_1[4] < ele_2[4]:
                    candidate.append(('D', 
                                        ele_1[2], 
                                        ele_2[4], 
                                        ele_2[3], 
                                        read_name,
                                        "TRA",
                                        ele_1[4]))
                    # ]chr:pos]N
                else:
                    candidate.append(('A', 
                                        ele_2[3], 
                                        ele_1[4], 
                                        ele_1[2], 
                                        read_name,
                                        "TRA",
                                        ele_2[4]))
                    # N[chr:pos[

def analysis_split_read(split_read, SV_size, RLength, read_name, candidate, MaxSize, query):
    '''
    read_start	read_end	ref_start	ref_end	chr	strand
    #0			#1			#2			#3		#4	#5
    '''
    SP_list = sorted(split_read, key = lambda x:x[0])

    # detect INS involoved in a translocation
    trigger_INS_TRA = 0	

    # Store Strands of INV

    if len(SP_list) == 2:
        ele_1 = SP_list[0]
        ele_2 = SP_list[1]
        if ele_1[4] == ele_2[4]:
            if ele_1[5] != ele_2[5]:
                analysis_inv(ele_1, 
                                ele_2, 
                                read_name, 
                                candidate["INV"],
                                SV_size)

            else:
                # dup & ins & del 
                a = 0
                if ele_1[5] == '-':
                    ele_1 = [RLength-SP_list[a+1][1], RLength-SP_list[a+1][0]]+SP_list[a+1][2:]
                    ele_2 = [RLength-SP_list[a][1], RLength-SP_list[a][0]]+SP_list[a][2:]
                    query = query[::-1]

                if ele_1[3] - ele_2[2] >= SV_size:
                    # if ele_2[1] - ele_1[1] >= ele_1[3] - ele_2[2]:
                    if ele_2[0] - ele_1[1] >= ele_1[3] - ele_2[2]:
                        candidate["INS"].append(((ele_1[3]+ele_2[2])/2, 
                                        ele_2[0]+ele_1[3]-ele_2[2]-ele_1[1], 
                                        read_name,
                                        str(query[ele_1[1]+int((ele_1[3]-ele_2[2])/2):ele_2[0]-int((ele_1[3]-ele_2[2])/2)]),
                                        "INS",
                                        ele_2[4]))
                    else:
                        candidate["DUP"].append((ele_2[2], 
                                            ele_1[3], 
                                            read_name,
                                            "DUP",
                                            ele_2[4]))

                delta_length = ele_2[0] + ele_1[3] - ele_2[2] - ele_1[1]
                if ele_1[3] - ele_2[2] < max(SV_size, delta_length/5) and delta_length >= SV_size:
                    if ele_2[2] - ele_1[3] <= max(100, delta_length/5) and (delta_length <= MaxSize or MaxSize == -1):
                        candidate["INS"].append(((ele_2[2]+ele_1[3])/2, 
                                            delta_length, 
                                            read_name,
                                            str(query[ele_1[1]+int((ele_2[2]-ele_1[3])/2):ele_2[0]-int((ele_2[2]-ele_1[3])/2)]),
                                            "INS",
                                            ele_2[4]))
                delta_length = ele_2[2] - ele_2[0] + ele_1[1] - ele_1[3]
                if ele_1[3] - ele_2[2] < max(SV_size, delta_length/5) and delta_length >= SV_size:
                    if ele_2[0] - ele_1[1] <= max(100, delta_length/5) and (delta_length <= MaxSize or MaxSize == -1):
                        candidate["DEL"].append((ele_1[3], 
                                            delta_length, 
                                            read_name,
                                            "DEL",
                                            ele_2[4]))
        else:
            trigger_INS_TRA = 1
            analysis_bnd(ele_1, ele_2, read_name, candidate["TRA"])

    else:
        # over three splits
        for a in range(len(SP_list[1:-1])):
            ele_1 = SP_list[a]
            ele_2 = SP_list[a+1]
            ele_3 = SP_list[a+2]

            if ele_1[4] == ele_2[4]:
                if ele_2[4] == ele_3[4]:
                    if ele_1[5] == ele_3[5] and ele_1[5] != ele_2[5]:
                        if ele_2[5] == '-':
                            # +-+
                            if ele_2[0] + 0.5 * (ele_3[2] - ele_1[3]) >= ele_1[1] and ele_3[0] + 0.5 * (ele_3[2] - ele_1[3]) >= ele_2[1]:
                                # No overlaps in split reads

                                if ele_2[2] >= ele_1[3] and ele_3[2] >= ele_2[3]:
                                    candidate["INV"].append(("++", 
                                                        ele_1[3], 
                                                        ele_2[3], 
                                                        read_name,
                                                        "INV",
                                                        ele_1[4]))
                                    # head-to-head
                                    # 5'->5'
                                    candidate["INV"].append(("--", 
                                                        ele_2[2], 
                                                        ele_3[2], 
                                                        read_name,
                                                        "INV",
                                                        ele_1[4]))
                                    # tail-to-tail
                                    # 3'->3'
                        else:
                            # -+-
                            if ele_1[1] <= ele_2[0] + 0.5 * (ele_1[2] - ele_3[3]) and ele_3[0] + 0.5 * (ele_1[2] - ele_3[3]) >= ele_2[1]:
                                # No overlaps in split reads

                                if ele_2[2] - ele_3[3] >= -50 and ele_1[2] - ele_2[3] >= -50:
                                    candidate["INV"].append(("++", 
                                                        ele_3[3], 
                                                        ele_2[3], 
                                                        read_name,
                                                        "INV",
                                                        ele_1[4]))
                                    # head-to-head
                                    # 5'->5'
                                    candidate["INV"].append(("--", 
                                                        ele_2[2], 
                                                        ele_1[2], 
                                                        read_name,
                                                        "INV",
                                                        ele_1[4]))
                                    # tail-to-tail
                                    # 3'->3'	

                    if len(SP_list) - 3 == a:
                        if ele_1[5] != ele_3[5]:
                            if ele_2[5] == ele_1[5]:
                                # ++-/--+
                                analysis_inv(ele_2, 
                                                ele_3, 
                                                read_name, 
                                                candidate["INV"], 
                                                SV_size)
                            else:
                                # +--/-++
                                analysis_inv(ele_1, 
                                                ele_2, 
                                                read_name, 
                                                candidate["INV"], 
                                                SV_size)

                    if ele_1[5] == ele_3[5] and ele_1[5] == ele_2[5]:
                        # dup & ins & del 
                        if ele_1[5] == '-':
                            ele_1 = [RLength-SP_list[a+2][1], RLength-SP_list[a+2][0]]+SP_list[a+2][2:]
                            ele_2 = [RLength-SP_list[a+1][1], RLength-SP_list[a+1][0]]+SP_list[a+1][2:]
                            ele_3 = [RLength-SP_list[a][1], RLength-SP_list[a][0]]+SP_list[a][2:]
                            query = query[::-1]

                        if ele_2[3] - ele_3[2] >= SV_size and ele_2[2] < ele_3[3]:
                            candidate["DUP"].append((ele_3[2], 
                                                ele_2[3], 
                                                read_name,
                                                "DUP",
                                                ele_2[4]))

                        if a == 0:
                            if ele_1[3] - ele_2[2] >= SV_size:
                                candidate["DUP"].append((ele_2[2], 
                                                    ele_1[3], 
                                                    read_name,
                                                    "DUP",
                                                    ele_2[4]))

                        delta_length = ele_2[0] + ele_1[3] - ele_2[2] - ele_1[1]
                        if ele_1[3] - ele_2[2] < max(SV_size, delta_length/5) and delta_length >= SV_size:
                            if ele_2[2] - ele_1[3] <= max(100, delta_length/5) and (delta_length <= MaxSize or MaxSize == -1):
                                if ele_3[2] >= ele_2[3]:
                                    candidate["INS"].append(((ele_2[2]+ele_1[3])/2, 
                                                        delta_length, 
                                                        read_name,
                                                        str(query[ele_1[1]+int((ele_2[2]-ele_1[3])/2):ele_2[0]-int((ele_2[2]-ele_1[3])/2)]),
                                                        "INS",
                                                        ele_2[4]))
                        delta_length = ele_2[2] - ele_2[0] + ele_1[1] - ele_1[3]
                        if ele_1[3] - ele_2[2] < max(SV_size, delta_length/5) and delta_length >= SV_size:
                            if ele_2[0] - ele_1[1] <= max(100, delta_length/5) and (delta_length <= MaxSize or MaxSize == -1):
                                if ele_3[2] >= ele_2[3]:
                                    candidate["DEL"].append((ele_1[3], 
                                                        delta_length, 
                                                        read_name,
                                                        "DEL",
                                                        ele_2[4]))
                        
                        if len(SP_list) - 3 == a:
                            ele_1 = ele_2
                            ele_2 = ele_3

                            delta_length = ele_2[0] + ele_1[3] - ele_2[2] - ele_1[1]
                            if ele_1[3] - ele_2[2] < max(SV_size, delta_length/5) and delta_length >= SV_size:
                                if ele_2[2] - ele_1[3] <= max(100, delta_length/5) and (delta_length <= MaxSize or MaxSize == -1):
                                    candidate["INS"].append(((ele_2[2]+ele_1[3])/2, 
                                                        delta_length, 
                                                        read_name,
                                                        str(query[ele_1[1]+int((ele_2[2]-ele_1[3])/2):ele_2[0]-int((ele_2[2]-ele_1[3])/2)]),
                                                        "INS",
                                                        ele_2[4]))

                            delta_length = ele_2[2] - ele_2[0] + ele_1[1] - ele_1[3]
                            if ele_1[3] - ele_2[2] < max(SV_size, delta_length/5) and ele_2[2] - ele_2[0] + ele_1[1] - ele_1[3] >= SV_size:
                                if ele_2[0] - ele_1[1] <= max(100, delta_length/5) and (delta_length <= MaxSize or MaxSize == -1):
                                    candidate["DEL"].append((ele_1[3], 
                                                        delta_length, 
                                                        read_name,
                                                        "DEL",
                                                        ele_2[4]))

                    if len(SP_list) - 3 == a and ele_1[5] != ele_2[5] and ele_2[5] == ele_3[5]:
                        ele_1 = ele_2
                        ele_2 = ele_3
                        ele_3 = None
                    if ele_3 == None or (ele_1[5] == ele_2[5] and ele_2[5] != ele_3[5]):
                        if ele_1[5] == '-':
                            ele_1 = [RLength-SP_list[a+1][1], RLength-SP_list[a+1][0]]+SP_list[a+1][2:]
                            ele_2 = [RLength-SP_list[a][1], RLength-SP_list[a][0]]+SP_list[a][2:]
                            query = query[::-1]
                        delta_length = ele_2[0] + ele_1[3] - ele_2[2] - ele_1[1]
                        if ele_1[3] - ele_2[2] < max(SV_size, delta_length/5) and delta_length >= SV_size:
                            if ele_2[2] - ele_1[3] <= max(100, delta_length/5) and (delta_length <= MaxSize or MaxSize == -1):
                                candidate["INS"].append(((ele_2[2]+ele_1[3])/2, 
                                                    delta_length, 
                                                    read_name,
                                                    str(query[ele_1[1]+int((ele_2[2]-ele_1[3])/2):ele_2[0]-int((ele_2[2]-ele_1[3])/2)]),
                                                    "INS",
                                                    ele_2[4]))

                        delta_length = ele_2[2] - ele_2[0] + ele_1[1] - ele_1[3]
                        if ele_1[3] - ele_2[2] < max(SV_size, delta_length/5) and delta_length >= SV_size:
                            if ele_2[0] - ele_1[1] <= max(100, delta_length/5) and (delta_length <= MaxSize or MaxSize == -1):
                                candidate["DEL"].append((ele_1[3], 
                                                    delta_length, 
                                                    read_name,
                                                    "DEL",
                                                    ele_2[4]))

            else:
                trigger_INS_TRA = 1
                analysis_bnd(ele_1, ele_2, read_name, candidate["TRA"])

                if len(SP_list) - 3 == a:
                    if ele_2[4] != ele_3[4]:
                        analysis_bnd(ele_2, ele_3, read_name, candidate["TRA"])

    if len(SP_list) >= 3 and trigger_INS_TRA == 1:
        if SP_list[0][4] == SP_list[-1][4]:
            if SP_list[0][5] != SP_list[-1][5]:
                pass
            else:
                if SP_list[0][5] == '+':
                    ele_1 = SP_list[0]
                    ele_2 = SP_list[-1]
                else:
                    ele_1 = [RLength-SP_list[-1][1], RLength-SP_list[-1][0]]+SP_list[-1][2:]
                    ele_2 = [RLength-SP_list[0][1],RLength-SP_list[0][0]]+SP_list[0][2:]
                    query = query[::-1]
                dis_ref = ele_2[2] - ele_1[3]
                dis_read = ele_2[0] - ele_1[1]
                if dis_ref < 100 and dis_read - dis_ref >= SV_size and (dis_read - dis_ref <= MaxSize or MaxSize == -1):
                    candidate["INS"].append((min(ele_2[2], ele_1[3]), 
                                        dis_read - dis_ref, 
                                        read_name,
                                        str(query[ele_1[1]+int(dis_ref/2):ele_2[0]-int(dis_ref/2)]),
                                        "INS",
                                        ele_2[4]))	

                if dis_ref <= -SV_size:
                    candidate["DUP"].append((ele_2[2], 
                                        ele_1[3], 
                                        read_name,
                                        "DUP",
                                        ele_2[4]))

def acquire_clip_pos(deal_cigar):
    seq = list(cigar.Cigar(deal_cigar).items())
    if seq[0][1] == 'S':
        first_pos = seq[0][0]
    else:
        first_pos = 0
    if seq[-1][1] == 'S':
        last_pos = seq[-1][0]
    else:
        last_pos = 0

    bias = 0
    for i in seq:
        if i[1] == 'M' or i[1] == 'D' or i[1] == '=' or i[1] == 'X':
            bias += i[0]
    return [first_pos, last_pos, bias]

def organize_split_signal(primary_info, Supplementary_info, total_L, SV_size, 
    min_mapq, max_split_parts, read_name, candidate, MaxSize, query):
    split_read = list()
    if len(primary_info) > 0:
        split_read.append(primary_info)
        min_mapq = 0
    for i in Supplementary_info:
        seq = i.split(',')
        local_chr = seq[0]
        local_start = int(seq[1])
        local_cigar = seq[3]
        local_strand = seq[2]
        local_mapq = int(seq[4])
        if local_mapq >= min_mapq:
            local_set = acquire_clip_pos(local_cigar)
            if local_strand == '+':
                 split_read.append([local_set[0], total_L-local_set[1], local_start, 
                     local_start+local_set[2], local_chr, local_strand])
            else:
                try:
                    split_read.append([local_set[1], total_L-local_set[0], local_start, 
                        local_start+local_set[2], local_chr, local_strand])
                except:
                    pass
    if len(split_read) <= max_split_parts or max_split_parts == -1:
        analysis_split_read(split_read, SV_size, total_L, read_name, candidate, MaxSize, query)

def generate_combine_sigs(sigs, Chr_name, read_name, svtype, candidate, merge_dis):
    if len(sigs) == 0:
        pass
    elif len(sigs) == 1:
        if svtype == 'INS':
            candidate.append((sigs[0][0], 
                                            sigs[0][1], 
                                            read_name,
                                            sigs[0][2],
                                            svtype,
                                            Chr_name))
        else:
            candidate.append((sigs[0][0], 
                                            sigs[0][1], 
                                            read_name,
                                            svtype,
                                            Chr_name))
    else:
        temp_sig = sigs[0]
        if svtype == "INS":
            temp_sig += [sigs[0][0]]
            for i in sigs[1:]:
                if i[0] - temp_sig[3] <= merge_dis:
                    temp_sig[1] += i[1]
                    temp_sig[2] += i[2]
                    temp_sig[3] = i[0]
                else:
                    candidate.append((temp_sig[0], 
                                                        temp_sig[1], 
                                                        read_name,
                                                        temp_sig[2],
                                                        svtype,
                                                        Chr_name))
                    temp_sig = i
                    temp_sig.append(i[0])
            candidate.append((temp_sig[0], 
                                                temp_sig[1], 
                                                read_name,
                                                temp_sig[2],
                                                svtype,
                                                Chr_name))
        else:
            temp_sig += [sum(sigs[0])]
            # merge_dis_bias = max([i[1]] for i in sigs)
            for i in sigs[1:]:
                if i[0] - temp_sig[2] <= merge_dis:
                    temp_sig[1] += i[1]
                    temp_sig[2] = sum(i)
                else: 
                    candidate.append((temp_sig[0], 
                                                        temp_sig[1], 
                                                        read_name,
                                                        svtype,
                                                        Chr_name))
                    temp_sig = i
                    temp_sig.append(i[0])
            candidate.append((temp_sig[0], 
                                                temp_sig[1], 
                                                read_name,
                                                svtype,
                                                Chr_name))

OPLIST=[
    pysam.CBACK,
    pysam.CDEL,
    pysam.CDIFF,
    pysam.CEQUAL,
    pysam.CHARD_CLIP,
    pysam.CINS,
    pysam.CMATCH,
    pysam.CPAD,
    pysam.CREF_SKIP,
    pysam.CSOFT_CLIP
]
RefChangeOp=set([0,2,7,8])

#QUERY CHANGE, REF CHANGE
CHANGETABLE={
    pysam.CMATCH:     (True,True),
    pysam.CINS:       (True,False),
    pysam.CDEL:       (False,True),
    pysam.CREF_SKIP:  (False,True),
    pysam.CPAD:       (False,False),
    pysam.CEQUAL:     (True,True),
    pysam.CDIFF:      (True,True)
}

CHANGEOP=[CHANGETABLE[i] if i in CHANGETABLE.keys() else (False,False) for i in range(max(OPLIST)+1)]
REFCHANGEOP=[CHANGETABLE[i][1] if i in CHANGETABLE.keys() else False for i in range(max(OPLIST)+1)]
INDELOP=[(i==pysam.CDEL or i==pysam.CINS) for i in range(max(OPLIST)+1)]

def parse_read(read, candidate, Chr_name, SV_size, min_mapq, max_split_parts, min_read_len, min_siglength, merge_del_threshold, merge_ins_threshold, MaxSize, family_mode ,family_member):
    if read.query_length < min_read_len:
        return []
    Combine_sig_in_same_read_ins = list()
    Combine_sig_in_same_read_del = list()
    #new start
    process_signal = detect_flag(read.flag)
    if read.mapq >= min_mapq:
        pos_start = read.reference_start # 0-based
        pos_end = read.reference_end
        sig_start=pos_start
        softclip_left = 0
        softclip_right = 0
        hardclip_left = 0
        hardclip_right = 0
        shift_ins_read = 0
        if read.cigar[0][0] == 4:
            softclip_left = read.cigar[0][1]
        elif read.cigar[0][0] == 5:
            hardclip_left = read.cigar[0][1]
        
        shift_ins_read=-hardclip_left
        for op, oplen in read.cigartuples:
            # calculate offset of an ins sig in read
            if op != 2:#might be fixed later
                shift_ins_read += oplen
            if oplen >= min_siglength and INDELOP[op]:
                if op==2:
                    Combine_sig_in_same_read_del.append([sig_start, oplen])
                    sig_start += oplen
                else:
                    Combine_sig_in_same_read_ins.append([sig_start, oplen,
                        str(read.query_sequence[shift_ins_read-oplen:shift_ins_read])])
            else:
                # if op in RefChangeOp:
                if REFCHANGEOP[op]:
                    sig_start += oplen

        
        if read.cigar[-1][0] == 4:
            softclip_right = read.cigar[-1][1]
        elif read.cigar[-1][0] == 5:
            hardclip_right = read.cigar[-1][1]

        if hardclip_left != 0:
            softclip_left = hardclip_left
        if hardclip_right != 0:
            softclip_right = hardclip_right

    # ************Combine signals in same read********************
    generate_combine_sigs(Combine_sig_in_same_read_ins, Chr_name, "%s/%s/%s"%(family_mode,family_member,read.query_name), "INS", candidate["INS"], merge_ins_threshold)
    generate_combine_sigs(Combine_sig_in_same_read_del, Chr_name, "%s/%s/%s"%(family_mode,family_member,read.query_name), "DEL", candidate["DEL"], merge_del_threshold)
    if process_signal == 1 or process_signal == 2: # 0 / 16
        Tags = read.get_tags()
        if read.mapq >= min_mapq:
            if process_signal == 1:
                primary_info = [softclip_left, read.query_length-softclip_right, pos_start, 
                pos_end, Chr_name, dic_starnd[process_signal]]
            else:
                primary_info = [softclip_right, read.query_length-softclip_left, pos_start, 
                pos_end, Chr_name, dic_starnd[process_signal]]
        else:
            primary_info = []

        for i in Tags:
            if i[0] == 'SA':
                Supplementary_info = i[1].split(';')[:-1]
                organize_split_signal(primary_info, Supplementary_info, read.query_length, 
                    SV_size, min_mapq, max_split_parts, "%s/%s/%s"%(family_mode,family_member,read.query_name), candidate, MaxSize, read.query_sequence)
    return candidate

def init_reading_process(sam_path):
    global samfile
    samfile=pysam.AlignmentFile(sam_path)

def cleanup():
    global samfile
    if samfile!=None:
        samfile.close()
        samfile=None

SVTYPES=["DEL", "INS", "DUP", "INV", "TRA"]
samfile=None
def single_pipe(sam_path, min_length, min_mapq, max_split_parts, min_read_len, temp_dir, 
                task, min_siglength, merge_del_threshold, merge_ins_threshold, MaxSize, bed_regions,family_mode,family_member):
    candidate = {}
    candidate["DEL"]=list()
    candidate["INS"]=list()
    candidate["DUP"]=list()
    candidate["INV"]=list()
    candidate["TRA"]=list()
    reads_info_list= list()
    Chr_name = task[0]
    global samfile
    for read in samfile.fetch(Chr_name, task[1], task[2]):
        # handle_read(read, task[1], bed_regions, candidate, reads_info_list, Chr_name, min_length, min_mapq, max_split_parts, min_read_len, min_siglength, merge_del_threshold, merge_ins_threshold, MaxSize)
        if read.flag == 256 or read.flag == 272:
            continue
        pos_start = read.reference_start # 0-based
        pos_end = read.reference_end
        in_bed = False
        if bed_regions != None:
            for bed_region in bed_regions:
                if pos_end <= bed_region[0] or pos_start >= bed_region[1]:
                    continue
                else:
                    in_bed = True
                    break
        else:
            in_bed = True
        if read.reference_start >= task[1] and in_bed:
            parse_read(read, candidate, Chr_name, min_length, min_mapq, max_split_parts, 
                                min_read_len, min_siglength, merge_del_threshold, 
                                merge_ins_threshold, MaxSize, family_mode, family_member)
            if read.mapq >= min_mapq:
                is_primary = 0
                if read.flag in [0, 16]:
                    is_primary = 1
                reads_info_list.append((pos_start, pos_end, is_primary, "%s/%s/%s"%(family_mode,family_member,read.query_name), Chr_name))
    pid=current_process().pid
    for sv_type in SVTYPES:
        with open("%s%s.%s.signatures/%s%s.pickle"%(temp_dir,family_mode,family_member,pid,sv_type),"ab") as f:
            pickle.dump(candidate[sv_type],f)
    with open("%s%s.%s.signatures/%sreads.pickle"%(temp_dir,family_mode,family_member,pid),"ab") as f:
        pickle.dump(reads_info_list,f)
    logging.info("Finished %s:%d-%d."%(Chr_name, task[1], task[2]))	
    gc.collect()
    return None

def multi_run_wrapper(args):
    return single_pipe(*args)

#old_file_sig[]=mem_sig[DEL: -2,-1,0,1,2, INS: -2,-1,0,1,2,3, DUP: -2,-1,0,1,2, INV: -2,-1,0,1,2,3, TRA: -2,-1,0,1,2,3,4, reads: -1,0,1,2,3]
def process_process_sigs_type(args):
    sv_type, temporary_dir, pids, write_old_sigs, family_mode, family_member=args
    #read
    type_candidates=[]
    for pid in pids:
        with open("%s%s.%s.signatures/%s%s.pickle"%(temporary_dir,family_mode,family_member,pid,sv_type), "rb") as f:
            while True:
                try:
                    candidate=pickle.load(f)
                    type_candidates.extend(candidate)
                except EOFError:
                    break
    #write
    if sv_type=="DEL":
        type_candidates.sort(key=lambda x: (x[-1], int(x[0]), x[1], x[2]))
        type_candidates=remove_duplicates_sorted(type_candidates)
        if write_old_sigs:
            with open("%s/%s.%s.DEL.sigs"%(temporary_dir,family_mode,family_member),"w") as f:
                if len(type_candidates)!=0:
                    for ele in type_candidates:
                        line="%s\t%s\t%d\t%d\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2])
                        print(line, end="",file=f)

    elif sv_type=="INS":
        type_candidates.sort(key=lambda x: (x[-1], int(x[0]), x[1], x[2], x[3]))
        type_candidates=remove_duplicates_sorted(type_candidates)
        if write_old_sigs:
            with open("%s/%s.%s.INS.sigs"%(temporary_dir,family_mode,family_member),"w") as f:
                if len(type_candidates)!=0:
                    for ele in type_candidates:
                        line="%s\t%s\t%d\t%d\t%s\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2], ele[3])
                        print(line, end="",file=f)
    elif sv_type=="DUP":
        type_candidates.sort(key=lambda x: (x[-1], int(x[0]), int(x[1]), x[2]))
        type_candidates=remove_duplicates_sorted(type_candidates)
        if write_old_sigs:
            with open("%s/%s.%s.DUP.sigs"%(temporary_dir,family_mode,family_member),"w") as f:
                if len(type_candidates)!=0:
                    for ele in type_candidates:
                        line="%s\t%s\t%d\t%d\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2])
                        print(line, end="",file=f)
    elif sv_type=="INV":
        type_candidates.sort(key=lambda x: (x[-1], x[0], int(x[1]), x[2], x[3]))
        type_candidates=remove_duplicates_sorted(type_candidates)
        if write_old_sigs:
            with open("%s/%s.%s.INV.sigs"%(temporary_dir,family_mode,family_member),"w") as f:
                if len(type_candidates)!=0:
                    for ele in type_candidates:
                        line="%s\t%s\t%s\t%d\t%d\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2], ele[3])
                        print(line, end="",file=f)
    elif sv_type=="TRA":
        type_candidates.sort(key=lambda x: (x[-1], x[2], x[0], int(x[1]), x[3], x[4], x[5]))
        type_candidates=remove_duplicates_sorted(type_candidates)
        if write_old_sigs:
            with open("%s/%s.%s.TRA.sigs"%(temporary_dir,family_mode,family_member),"w") as f:
                if len(type_candidates)!=0:
                    for ele in type_candidates:
                        line="%s\t%s\t%s\t%d\t%s\t%d\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2], ele[3], ele[4])
                        print(line, end="",file=f)
    elif sv_type=="reads":
        type_candidates.sort(key=lambda x: (x[-1]))#reads file not deduped, might be resolved later
        if write_old_sigs:
            with open("%s/%s.%s.reads.sigs"%(temporary_dir,family_mode,family_member), "w") as f:
                if len(type_candidates)!=0:
                    for ele in type_candidates:
                        line="%s\t%d\t%d\t%d\t%s\n"%(ele[-1], ele[0], ele[1], ele[2], ele[3])
                        print(line,end="",file=f)
    index={}
    reads_count={}
    with open("%s/%s.%s.%s.pickle"%(temporary_dir,family_mode,family_member,sv_type),"wb") as f:
        chr=None
        startl=0
        start=0
        if sv_type=="reads":
            if len(type_candidates)!=0:
                for i in range(len(type_candidates)):
                    ele=type_candidates[i]
                    if ele[-1]!=chr:
                        if chr==None:
                            chr=ele[-1]
                        else:
                            dump=pickle.dumps(type_candidates[startl:i])
                            f.write(dump)
                            index[chr]=start
                            reads_count[chr]=i-startl
                            start+=len(dump)
                            chr=ele[-1]
                            startl=i
                pickle.dump(type_candidates[startl:],f)
                index[chr]=start
                reads_count[chr]=len(type_candidates)-startl
        else:
            if len(type_candidates)!=0:
                for i in range(len(type_candidates)):
                    ele=type_candidates[i]
                    if ele[-1]!=chr:
                        if chr==None:
                            chr=ele[-1]
                        else:
                            dump=pickle.dumps(type_candidates[startl:i])
                            f.write(dump)
                            index[chr]=start
                            start+=len(dump)
                            chr=ele[-1]
                            startl=i
                pickle.dump(type_candidates[startl:],f)
                index[chr]=start
    return (sv_type,index,reads_count)

def write_sigs(temporary_dir, candidates, reads_info_list, prefix=""):
    index={}
    with open("%s/%sDEL.sigs"%(temporary_dir,prefix),"w") as f:
        index["DEL"]={}
        chr=None
        start=0
        bytecount=0
        if len(candidates["DEL"])!=0:
            for ele in candidates["DEL"]:
                line="%s\t%s\t%d\t%d\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2])
                if ele[-1]!=chr:
                    if chr==None:
                        chr=ele[-1]
                    else:
                        index["DEL"][chr]=start
                        chr=ele[-1]
                        start=bytecount
                bytecount+=len(line)
                print(line, end="",file=f)
            index["DEL"][chr]=start
    with open("%s/%sINS.sigs"%(temporary_dir,prefix),"w") as f:
        index["INS"]={}
        chr=None
        start=0
        bytecount=0
        if len(candidates["INS"])!=0:
            for ele in candidates["INS"]:
                line="%s\t%s\t%d\t%d\t%s\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2], ele[3])
                if ele[-1]!=chr:
                    if chr==None:
                        chr=ele[-1]
                    else:
                        index["INS"][chr]=start
                        chr=ele[-1]
                        start=bytecount
                bytecount+=len(line)
                print(line, end="",file=f)
            index["INS"][chr]=start
    with open("%s/%sDUP.sigs"%(temporary_dir,prefix),"w") as f:
        index["DUP"]={}
        chr=None
        start=0
        bytecount=0
        if len(candidates["DUP"])!=0:
            for ele in candidates["DUP"]:
                line="%s\t%s\t%d\t%d\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2])
                if ele[-1]!=chr:
                    if chr==None:
                        chr=ele[-1]
                    else:
                        index["DUP"][chr]=start
                        chr=ele[-1]
                        start=bytecount
                bytecount+=len(line)
                print(line, end="",file=f)
            index["DUP"][chr]=start
    with open("%s/%sINV.sigs"%(temporary_dir,prefix),"w") as f:
        index["INV"]={}
        chr=None
        start=0
        bytecount=0
        if len(candidates["INV"])!=0:
            for ele in candidates["INV"]:
                line="%s\t%s\t%s\t%d\t%d\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2], ele[3])
                if ele[-1]!=chr:
                    if chr==None:
                        chr=ele[-1]
                    else:
                        index["INV"][chr]=start
                        chr=ele[-1]
                        start=bytecount
                bytecount+=len(line)
                print(line, end="",file=f)
            index["INV"][chr]=start
    with open("%s/%sTRA.sigs"%(temporary_dir,prefix),"w") as f:
        index["TRA"]={}
        chr=None
        start=0
        bytecount=0
        if len(candidates["TRA"])!=0:
            for ele in candidates["TRA"]:
                line="%s\t%s\t%s\t%d\t%s\t%d\t%s\n"%(ele[-2], ele[-1], ele[0], ele[1], ele[2], ele[3], ele[4])
                if ele[-1]!=chr:
                    if chr==None:
                        chr=ele[-1]
                    else:
                        index["TRA"][chr]=start
                        chr=ele[-1]
                        start=bytecount
                bytecount+=len(line)
                print(line, end="",file=f)
            index["TRA"][chr]=start
    with open("%s/%sreads.sigs"%(temporary_dir,prefix), "w") as f:
        for ele in reads_info_list:
            print("%s\t%d\t%d\t%d\t%s\n"%(ele[-1], ele[0], ele[1], ele[2], ele[3]), end="", file=f)
    with open("%s/sigindex.pickle"%temporary_dir,"wb") as f:
        pickle.dump(index,f)
    return index

def remove_duplicates_sorted(sorted):
    if len(sorted) == 0:
        return []
    i=0
    j=0
    while i < len(sorted):
        if sorted[i] != sorted[j]:
            j += 1
            sorted[j] = sorted[i]
        i += 1
    
    return sorted[:j+1]

def remove_duplicates(data):
    seen = set()
    result = []
    for item in data:
        if item not in seen:
            result.append(item)
            seen.add(item)
    
    return result

def add_candidates(result):
    global candidates
    candidate=result[0][0]
    candidates["DEL"].extend(candidate["DEL"])
    candidates["INS"].extend(candidate["INS"])
    candidates["DUP"].extend(candidate["DUP"])
    candidates["INV"].extend(candidate["INV"])
    candidates["TRA"].extend(candidate["TRA"])
    candidates["reads_info"].extend(result[0][1])

# candidates={}
def main_ctrl(args, argv):
    logging.info(args)
    family_mode_index_ls = ["M1","M2"]
    family_member_set = [["1","2","3"],["1","2"]]
    if args.family_mode not in family_mode_index_ls:
        raise FileNotFoundError("[Errno 2] Wrong family mode: '%s'"%args.family_mode)
    if args.family_mode == "M1" :
        if args.input_offspring is None :
            raise FileNotFoundError("[Errno 2] Missing BAM files for family member: offspring.")
        if args.input_parent_1 is None :
            raise FileNotFoundError("[Errno 2] Missing BAM files for family member: father.")
        if args.input_parent_2 is None :
            raise FileNotFoundError("[Errno 2] Missing BAM files for family member: mother.")
        bam_path_list = [args.input_offspring,args.input_parent_1,args.input_parent_2]
    if args.family_mode == "M2" :
        if args.input_offspring is None :
            raise FileNotFoundError("[Errno 2] Missing BAM files for family member: offspring.")
        if args.input_parent_1 is None :
            raise FileNotFoundError("[Errno 2] Missing BAM files for family member: parent.")
        if args.input_parent_2 is not None :
            raise FileNotFoundError("[Errno 2] Wrong extra BAM files for family member: parent.")
        bam_path_list = [args.input_offspring,args.input_parent_1]
    
    family_member_ls = family_member_set[family_mode_index_ls.index(args.family_mode)]
    # check the temporary files
    if args.reference is None :
        raise FileNotFoundError("[Errno 2] Missing reference files.")
    if not os.path.isfile(args.reference):
        raise FileNotFoundError("[Errno 2] No such file: '%s'"%args.reference)
    if args.work_dir is None :
        raise FileNotFoundError("[Errno 2] Missing work dir.")
    if args.work_dir[-1] == '/':
        temporary_dir = args.work_dir
    else:
        temporary_dir = args.work_dir+'/'
    if not os.path.exists(args.work_dir):
        raise FileNotFoundError("[Errno 2] No such directory: '%s'"%args.work_dir)
    if args.output is None :
        raise FileNotFoundError("[Errno 2] Missing output vcf files.")
    
    if args.sequencing_platform is None and args.gold_standard_version is None :
        run_max_cluster_bias_INS = args.max_cluster_bias_INS
        run_diff_ratio_merging_INS = args.diff_ratio_merging_INS
        run_max_cluster_bias_DEL = args.max_cluster_bias_DEL
        run_diff_ratio_merging_DEL = args.diff_ratio_merging_DEL
    else :
        if args.sequencing_platform is not None and args.sequencing_platform.lower() != "HiFI".lower() and args.sequencing_platform.lower() != "ONT".lower()  :
            raise FileNotFoundError("[Errno 2] Wrong sequencing platform: '%s'. Must be HiFI or ONT"%args.sequencing_platform)
        if args.gold_standard_version is not None and args.gold_standard_version.lower() != "NIST".lower() and args.gold_standard_version.lower() != "T2T".lower() :
            raise FileNotFoundError("[Errno 2] Wrong gold standard version: '%s'. Must be NIST or T2T"%args.gold_standard_version)
        if args.sequencing_platform is None :
            run_sequencing_platform = "HiFI".lower()
        else :
            run_sequencing_platform = args.sequencing_platform.lower()
        if args.gold_standard_version is None :
            run_gold_standard_version = "T2T".lower()
        else :
            run_gold_standard_version = args.gold_standard_version.lower()
        if run_sequencing_platform == "HiFI".lower() and run_gold_standard_version == "T2T".lower() :
            run_max_cluster_bias_INS = 50
            run_diff_ratio_merging_INS = 0.05
            run_max_cluster_bias_DEL = 100
            run_diff_ratio_merging_DEL = 0.05
        elif run_sequencing_platform == "HiFI".lower() and run_gold_standard_version == "NIST".lower() :
            run_max_cluster_bias_INS = 150
            run_diff_ratio_merging_INS = 0.05
            run_max_cluster_bias_DEL = 150
            run_diff_ratio_merging_DEL = 0.05
        elif run_sequencing_platform == "ONT".lower() and run_gold_standard_version == "T2T".lower() :
            run_max_cluster_bias_INS = 100
            run_diff_ratio_merging_INS = 0.1
            run_max_cluster_bias_DEL = 100
            run_diff_ratio_merging_DEL = 0.3
        elif run_sequencing_platform == "ONT".lower() and run_gold_standard_version == "NIST".lower() :
            run_max_cluster_bias_INS = 300
            run_diff_ratio_merging_INS = 0.05
            run_max_cluster_bias_DEL = 400
            run_diff_ratio_merging_DEL = 0.15
        else :
            raise FileNotFoundError("[Errno 2] Wrong sequencing platform: '%s'. Must be HiFI or ONT"%args.sequencing_platform)
            raise FileNotFoundError("[Errno 2] Wrong gold standard version: '%s'. Must be NIST or T2T"%args.gold_standard_version)

    logging.info("max_cluster_bias_INS:%s"%(str(run_max_cluster_bias_INS)))
    logging.info("diff_ratio_merging_INS:%s"%(str(run_diff_ratio_merging_INS)))
    logging.info("max_cluster_bias_DEL:%s"%(str(run_max_cluster_bias_DEL)))
    logging.info("diff_ratio_merging_DEL:%s"%(str(run_diff_ratio_merging_DEL)))
    logging.info("parents_phasing:%s"%(str(args.parents_phasing)))
    if args.execute_stage in [0,1] :
        split_reference_chromosomes(temporary_dir,args.reference)

    input_file_ls = [args.input_offspring,args.input_parent_1,args.input_parent_2]
    if args.execute_stage in [0,1] :
        for family_member_i in range(len(family_member_ls)) :            
            family_member = family_member_ls[family_member_i]
            for item in SVTYPES:
                if os.path.exists("%s/%s.%s.%s.sigs"%(temporary_dir,args.family_mode,family_member,item)):
                    raise FileExistsError("[Errno 2] File exists: '%s/%s.%s.%s.sigs'"%(temporary_dir,args.family_mode,family_member,item))
                if os.path.exists("%s/%s.%s.%s.pickle"%(temporary_dir,args.family_mode,family_member,item)):
                    raise FileExistsError("[Errno 2] File exists: '%s/%s.%s.%s.pickle'"%(temporary_dir,args.family_mode,family_member,item))
            if os.path.exists("%s/%s.%s.signatures"%(temporary_dir,args.family_mode,family_member)):
                    raise FileExistsError("[Errno 2] File exists: '%s/%s.%s.signatures'"%(temporary_dir,args.family_mode,family_member))
            
        contigINFO_fam_ls = []
        for family_member_i in range(len(family_member_ls)) :
            input_file = input_file_ls[family_member_i]
            family_member = family_member_ls[family_member_i]
            samfile = pysam.AlignmentFile(input_file)
            contig_num = len(samfile.get_index_statistics())
            logging.info("The total number of chromsomes: %d"%(contig_num))

            Task_list = list()
            chr_name_list = list()
            contigINFO = list()

            ref_ = samfile.get_index_statistics()
            
            total_mapped=0
            for i in ref_:
                total_mapped+=i[1]
            mapped_unit=total_mapped/args.threads/10
            logging.info(total_mapped)
            logging.info(mapped_unit)
            for i in ref_:
                chr_name_list.append(i[0])
                local_ref_len = samfile.get_reference_length(i[0])
                contigINFO.append([i[0], local_ref_len])
                if total_mapped==0 or i[1]<=mapped_unit:
                    batch_size=args.batches
                else:
                    batch_size=local_ref_len/(int(i[1]/mapped_unit)+1)
                logging.info("%s/%d"%(str(i[0]),batch_size))
                if local_ref_len < batch_size:
                    Task_list.append([i[0], 0, local_ref_len])
                else:
                    pos = 0
                    task_round = int(local_ref_len/batch_size)
                    for j in range(task_round):
                        Task_list.append([i[0], pos, pos+batch_size])
                        pos += batch_size
                    if pos < local_ref_len:
                        Task_list.append([i[0], pos, local_ref_len])
            bed_regions = load_bed(args.include_bed, Task_list)

            contigINFO_fam_ls.append(contigINFO)

            candidates={}
            candidates["DEL"]=list()
            candidates["INS"]=list()
            candidates["DUP"]=list()
            candidates["INV"]=list()
            candidates["TRA"]=list()
            reads_info_list=list()
            candidates["reads_info"]=reads_info_list

            atexit.register(cleanup)
            analysis_pools = Pool(processes=int(args.threads), initializer=init_reading_process, initargs=(input_file,))
            os.mkdir("%s%s.%s.signatures"%(temporary_dir,args.family_mode,family_member))
            results=[]#use this is faster than make a long paras list
            for i in range(len(Task_list)):
                paras = [(input_file, 
                            args.min_size, 
                            args.min_mapq, 
                            args.max_split_parts, 
                            args.min_read_len, 
                            temporary_dir, 
                            Task_list[i], 
                            args.min_siglength, 
                            args.merge_del_threshold/2, 
                            args.merge_ins_threshold/2, 
                            args.max_size,
                            None if bed_regions == None else bed_regions[i],
                            args.family_mode,
                            family_member)]
                analysis_pools.map_async(multi_run_wrapper, paras)
            pids = [process.pid for process in analysis_pools._pool]
            analysis_pools.close()
            analysis_pools.join()
            logging.info("Rebuilding signatures of structural variants.")

            analysis_pools = Pool(processes=int(args.threads))
            paras=[]
            for sv_type in SVTYPES:
                paras.append((sv_type,temporary_dir,pids,args.write_old_sigs,args.family_mode,family_member))
            paras.append(("reads",temporary_dir,pids,args.write_old_sigs,args.family_mode,family_member))
            results=analysis_pools.map_async(process_process_sigs_type, paras)
            analysis_pools.close()
            analysis_pools.join()
            sigs_index={}
            for r in results.get():
                if r!=None:
                    sigs_index[r[0]]=r[1]
                    if r[0]=="reads":
                        sigs_index["reads_count"]=r[2]
            with open("%s/%s.%s.sigindex.pickle"%(temporary_dir,args.family_mode,family_member),"wb") as f:
                pickle.dump(sigs_index,f)
            del reads_info_list
            del results
            del candidates
            gc.collect()
            logging.info("Rebuilding signatures completed.")
            samfile.close()
        contigINFO_dict = {}
        for contigINFO in contigINFO_fam_ls :
            for i in contigINFO :
                if i[0] not in contigINFO_dict :
                    contigINFO_dict[i[0]] = i[1]
                else :
                    if i[1] > contigINFO_dict[i[0]] :
                        contigINFO_dict[i[0]] = i[1]
        contigINFO = [[key, value] for key, value in contigINFO_dict.items()]
        with open("%s/%s.contigINFO.pickle"%(temporary_dir,args.family_mode),"wb") as f:
            pickle.dump(contigINFO,f)
        logging.info("Contig INFO complete.")
    
    if args.execute_stage in [0,2] :
        minimum_support_reads_list = [min(round(float(x)),5) for x in args.min_support_list.split(",")]
        if len(minimum_support_reads_list) != len(family_member_ls) :
            raise FileExistsError("[Errno 2] The length of number of family members and minimum support reads is not same:%d/%d"%(len(minimum_support_reads_list),len(family_member_ls)))

        for family_member_i in range(len(family_member_ls)) :            
            family_member = family_member_ls[family_member_i]
            if os.path.exists("%s/%s.%s.results"%(temporary_dir,args.family_mode,family_member)):
                    raise FileExistsError("[Errno 2] File exists: '%s/%s.%s.results'"%(temporary_dir,args.family_mode,family_member))
        
        result = list()
        chr_ls = {"DEL":[],"INS":[],"INV":[],"DUP":[],"TRA":[]}
        for family_member in family_member_ls :
            with open("%s/%s.%s.sigindex.pickle"%(temporary_dir,args.family_mode,family_member), 'rb') as f:
                sigs_index=pickle.load(f)
                f.close()
            for sv_type in chr_ls :
                for chr in sigs_index[sv_type] :
                    if chr not in chr_ls[sv_type] :
                        chr_ls[sv_type].append(chr)

        logging.info("Clustering structural variants.")
        analysis_pools = Pool(processes=int(args.threads))
        try :
            
            # +++++DEL+++++
            for chr in chr_ls["DEL"]:
                para = [(temporary_dir, 
                        chr, 
                        min(minimum_support_reads_list),
                        run_diff_ratio_merging_DEL, 
                        run_max_cluster_bias_DEL, 
                        # args.diff_ratio_filtering_DEL, 
                        minimum_support_reads_list, 
                        args.gt_round,
                        args.remain_reads_ratio,
                        args.merge_del_threshold,
                        args.read_pos_interval,
                        args.family_mode,
                        args.performing_phasing)]
                result.append(analysis_pools.map_async(run_del, para))

            # +++++INS+++++
            for chr in chr_ls["INS"]:
                para = [(temporary_dir, 
                        chr, 
                        min(minimum_support_reads_list), 
                        run_diff_ratio_merging_INS, 
                        run_max_cluster_bias_INS, 
                        minimum_support_reads_list, 
                        args.gt_round,
                        args.remain_reads_ratio,
                        args.merge_ins_threshold,
                        args.read_pos_interval,
                        args.family_mode,
                        args.performing_phasing,
                        args.all_ins_singnature_reads)]
                result.append(analysis_pools.map_async(run_ins, para))
                #logging.info(para)
            
            if not args.performing_assembly :
                # +++++INV+++++
                for chr in chr_ls["INV"]:
                    para = [(temporary_dir, 
                            chr, 
                            min(minimum_support_reads_list), 
                            args.max_cluster_bias_INV, 
                            minimum_support_reads_list, 
                            args.min_size, 
                            args.max_size,
                            args.gt_round,
                            args.read_pos_interval,
                            args.family_mode,
                            args.performing_phasing)]
                    result.append(analysis_pools.map_async(run_inv, para))
            
            
                # +++++DUP+++++
                for chr in chr_ls["DUP"]:
                    para = [(temporary_dir, 
                            chr, 
                            min(minimum_support_reads_list), 
                            args.max_cluster_bias_DUP,
                            minimum_support_reads_list, 
                            args.min_size, 
                            args.max_size,
                            args.gt_round,
                            args.read_pos_interval,
                            args.family_mode,
                            args.performing_phasing)]
                    result.append(analysis_pools.map_async(run_dup, para))
            
            
            if args.run_TRA :
                # +++++TRA+++++
                for chr in chr_ls["TRA"]:
                    para = [(temporary_dir, 
                            chr, 
                            min(minimum_support_reads_list), 
                            args.diff_ratio_filtering_TRA, 
                            args.max_cluster_bias_TRA, 
                            minimum_support_reads_list, 
                            bam_path_list,
                            args.gt_round,
                            args.read_pos_interval,
                            args.family_mode,
                            args.performing_phasing)]
                    result.append(analysis_pools.map_async(run_tra, para))
            
        finally :
            analysis_pools.close()
            analysis_pools.join()
        
        phasing_svtype_ls = ["INS","DEL","DUP","INV"]
        no_phasing_svtype_ls = ["TRA"]
        phasing_fam_results=[]
        no_phasing_fam_results=[]
        for family_member_i in range(len(family_member_ls)) :
            phasing_results = {}
            no_phasing_results = {}
            for res in result:
                try:
                    chr, svs = res.get()[0]
                    if svs[0][0][1] in phasing_svtype_ls :
                        if chr not in phasing_results.keys() :
                            phasing_results[chr]=[]
                        phasing_results[chr].extend(svs[family_member_i])
                    else :
                        if chr not in no_phasing_results.keys() :
                            no_phasing_results[chr]=[]
                        no_phasing_results[chr].extend(svs[family_member_i])
                except:
                    pass
            phasing_fam_results.append(phasing_results)
            no_phasing_fam_results.append(no_phasing_results)

        for family_member_i in range(len(phasing_fam_results)) :
            phasing_results = phasing_fam_results[family_member_i]
            for chr in phasing_results.keys() :
                phasing_results[chr] = sorted(phasing_results[chr], key=lambda x: int(x[2]))
                phasing_results[chr] = allele_correction(chr,family_member_i,phasing_results[chr],minimum_support_reads_list[family_member_i])
        for no_phasing_results in no_phasing_fam_results :
            for chr in no_phasing_results.keys() :
                no_phasing_results[chr] = sorted(no_phasing_results[chr], key=lambda x: x[2])
        
        if args.performing_phasing :
            logging.info("Phasing structural variants.")
            result = list()
            
            # phasing
            for chr in phasing_fam_results[0].keys():
                sv_fam_ls = []
                for family_member_i in range(len(family_member_ls)) :
                    sv_fam_ls.append(phasing_fam_results[family_member_i][chr])
                result.append(genetic_phasing_family(temporary_dir,
                                                     chr, 
                                                     sv_fam_ls, 
                                                     args.family_mode, 
                                                     args.read_pos_interval,
                                                     minimum_support_reads_list,
                                                     args.phase_all_ctgs,
                                                     args.parents_phasing,
                                                     input_file_ls,
                                                     args.remap_merge_k,
                                                     args.remap_minimizer_window,
                                                     args.assembly_correction_threshold,
                                                     args.performing_assembly,
                                                     args.multiple_phasing))
            
            phasing_fam_results=[]
            for family_member_i in range(len(family_member_ls)) :
                phasing_results = {}
                if family_member_i == 0 :
                    child_hap1_results = {}
                    child_hap2_results = {}
                    father_hap1_results = {}
                    father_hap2_results = {}
                    mother_hap1_results = {}
                    mother_hap2_results = {}
                for res in result:
                    try:
                        chr, svs, ch_hap_1, ch_hap_2, fa_hap_1, fa_hap_2, mo_hap_1, mo_hap_2 = res[0], res[1], res[2], res[3], res[4], res[5], res[6], res[7]
                        if chr not in phasing_results.keys() :
                            phasing_results[chr]=[]
                            if family_member_i == 0 :
                                child_hap1_results[chr]=[]
                                child_hap2_results[chr]=[]
                                father_hap1_results[chr]=[]
                                father_hap2_results[chr]=[]
                                mother_hap1_results[chr]=[]
                                mother_hap2_results[chr]=[]
                        phasing_results[chr].extend(svs[family_member_i])
                        if family_member_i == 0 :
                            child_hap1_results[chr].extend(ch_hap_1)
                            child_hap2_results[chr].extend(ch_hap_2)
                            father_hap1_results[chr].extend(fa_hap_1)
                            father_hap2_results[chr].extend(fa_hap_2)
                            mother_hap1_results[chr].extend(mo_hap_1)
                            mother_hap2_results[chr].extend(mo_hap_2)
                    except:
                        pass
                phasing_fam_results.append(phasing_results)
            
            if args.output_haplotype_reads :
                generate_haplotype_read_names(temporary_dir,
                                              args.family_mode,
                                              phasing_fam_results,
                                              child_hap1_results,
                                              child_hap2_results,
                                              father_hap1_results,
                                              father_hap2_results,
                                              mother_hap1_results,
                                              mother_hap2_results)
            
            if args.performing_assembly :
                if args.family_assembly_reads :
                    if os.path.exists(args.work_dir+"/assembly.fq.results"):
                        raise FileExistsError("[Errno 2] No such directory: '%s'"%(args.work_dir+"/assembly.fq.results"))
                    else :
                        os.makedirs(args.work_dir+"/assembly.fq.results/father", exist_ok=True)
                        os.makedirs(args.work_dir+"/assembly.fq.results/mother", exist_ok=True)
                        os.makedirs(args.work_dir+"/assembly.fq.results/num", exist_ok=True)
                # assembly
                logging.info("Correcting structural variants by local assembly.")
                sv_extension_scope = 1000
                nearsv_maxnum = 50
                nearsv_minnum = 30
                pahsing_pools = Pool(processes=args.threads)
                for chr in phasing_fam_results[0].keys():
                    sv_fam_ls = []
                    for family_member_i in range(len(family_member_ls)) :
                        sv_fam_ls.append(phasing_fam_results[family_member_i][chr])
                    nearby_list = []
                    for sv_i in range(len(sv_fam_ls[0])-1) :
                        if int(sv_fam_ls[0][sv_i+1][2]) - int(sv_fam_ls[0][sv_i][2]) < sv_extension_scope :
                            if len(nearby_list) == 0 or (len(nearby_list) > 0 and nearby_list[-1][1] != -1) :
                                if len(nearby_list) > 0 and max([abs(x) for x in nearby_list[-1][2]]) < 50 :
                                    nearby_list.pop()
                                nearby_list.append([sv_i,-1,[int(sv_fam_ls[0][sv_i][3])]])
                            else :
                                nearby_list[-1][2].append(int(sv_fam_ls[0][sv_i][3]))
                                if len(nearby_list[-1][2]) >= nearsv_maxnum + nearsv_minnum :
                                    nearby_list.append([sv_i-nearsv_minnum+1,-1,nearby_list[-1][2][nearsv_maxnum:]])
                                    nearby_list[-2][1] = sv_i-nearsv_minnum
                                    nearby_list[-2][2] = nearby_list[-2][2][0:nearsv_maxnum]
                        elif len(nearby_list) > 0 and nearby_list[-1][1] == -1 :
                            nearby_list[-1][1] = sv_i
                            nearby_list[-1][2].append(int(sv_fam_ls[0][sv_i][3]))
                    if len(nearby_list) > 0 and nearby_list[-1][1] == -1 :
                        nearby_list[-1][1] = len(sv_fam_ls[0])-1
                    for near_sv_index in range(len(nearby_list)) :
                        if near_sv_index == 0 :
                            near_sta = 0
                            if len(nearby_list) > 1 :
                                near_end = nearby_list[near_sv_index+1][0] - 1
                            else :
                                near_end = len(sv_fam_ls[0]) - 1
                            homo_range = [0-near_sta,near_end-near_sta]
                        elif near_sv_index == len(nearby_list) - 1:
                            near_sta = nearby_list[near_sv_index-1][1]+1
                            near_end = len(sv_fam_ls[0])-1
                            homo_range = [nearby_list[near_sv_index][0]-near_sta,near_end-near_sta]
                        else :
                            near_sta = nearby_list[near_sv_index-1][1]+1
                            if len(nearby_list) > 1 :
                                near_end = nearby_list[near_sv_index+1][0]-1
                            else :
                                near_end = len(sv_fam_ls[0])-1
                            homo_range = [nearby_list[near_sv_index][0]-near_sta,near_end-near_sta]
                        
                        para = [(temporary_dir,
                                 chr, 
                                 [svs[near_sta:near_end+1] for svs in sv_fam_ls], 
                                 child_hap1_results[chr][near_sta:near_end+1], 
                                 child_hap2_results[chr][near_sta:near_end+1], 
                                 father_hap1_results[chr][near_sta:near_end+1], 
                                 father_hap2_results[chr][near_sta:near_end+1], 
                                 mother_hap1_results[chr][near_sta:near_end+1], 
                                 mother_hap2_results[chr][near_sta:near_end+1], 
                                 [near_sta,near_end],
                                 homo_range, 
                                 args.family_mode, 
                                 input_file_ls, 
                                 minimum_support_reads_list, 
                                 args.remap_merge_k, 
                                 args.remap_minimizer_window, 
                                 args.assembly_correction_threshold,
                                 args.assembly_correction_setting,
                                 args.family_assembly_reads, 
                                 args.similarity_supplement_threshold,
                                 args.assembly_accelerate)]
                        result.append(pahsing_pools.map_async(run_assembly, para))

                pahsing_pools.close()
                pahsing_pools.join()
                
                phasing_fam_results=[]
                for family_member_i in range(len(family_member_ls)) :
                    phasing_results = {}
                    for res in result:
                        try:
                            chr, svs = res.get()[0]
                            if chr not in phasing_results.keys() :
                                phasing_results[chr]=[]
                            phasing_results[chr].extend(svs[family_member_i])
                        except:
                            pass
                    phasing_fam_results.append(phasing_results)
            
            # No phasing
            result = list()
            for chr in no_phasing_fam_results[0].keys():
                sv_fam_ls = []
                for family_member_i in range(len(family_member_ls)) :
                    sv_fam_ls.append(no_phasing_fam_results[family_member_i][chr])
                result.append(genetic_no_phasing_family(chr, 
                                                     sv_fam_ls, 
                                                     args.family_mode, 
                                                     args.read_pos_interval,
                                                     minimum_support_reads_list,
                                                     args.phase_all_ctgs,
                                                     args.parents_phasing))
            no_phasing_fam_results=[]
            for family_member_i in range(len(family_member_ls)) :
                no_phasing_results = {}
                for res in result:
                    try:
                        chr, svs = res[0], res[1]
                        if chr not in no_phasing_results.keys() :
                            no_phasing_results[chr]=[]
                        no_phasing_results[chr].extend(svs[family_member_i])
                    except:
                        pass
                no_phasing_fam_results.append(no_phasing_results)
        
        fam_results=[]
        for family_member_i in range(len(family_member_ls)) :
            results = {}
            for chr in phasing_fam_results[family_member_i].keys():
                if chr not in results.keys() :
                    results[chr]=[]
                results[chr].extend(phasing_fam_results[family_member_i][chr])
            for chr in no_phasing_fam_results[family_member_i].keys():
                if chr not in results.keys() :
                    results[chr]=[]
                results[chr].extend(no_phasing_fam_results[family_member_i][chr])
            fam_results.append(results)
        for results in fam_results :
            for chr in results.keys() :
                results[chr] = sorted(results[chr], key=lambda x: (int(x[2]),int(x[3])))
                
        # remove duplicate variants
        for chr in fam_results[0].keys() : 
            fam_sv_ls = []
            for fam_i in range(len(fam_results)) :
                fam_sv_ls.append(fam_results[fam_i][chr])
            fam_sv_ls = remove_redundant_samesv(chr, fam_sv_ls, args.family_mode)
            fam_sv_ls = remove_redundant_pos(chr, fam_sv_ls, args.family_mode)
            
            for fam_i in range(len(fam_results)) :
                fam_results[fam_i][chr] = fam_sv_ls[fam_i]

        empty_chr_ls = []
        for chr in fam_results[0].keys():
            is_empty = True
            for results in fam_results :
                if results[chr] != [] :
                    is_empty = False
            if is_empty :
                empty_chr_ls.append(chr)
        for chr in empty_chr_ls :
            logging.info("No SV in %s"%(chr))
            for results in fam_results :
                del results[chr]
        
        no_reference_chr_ls = []
        fa_file = pysam.FastaFile(args.reference)
        reference_chrs = fa_file.references
        for chr in fam_results[0].keys():
            if chr not in reference_chrs:
                no_reference_chr_ls.append(chr)
        for chr in no_reference_chr_ls :
            logging.info("No reference sequence in %s of %s"%(args.reference,chr))
            for results in fam_results :
                del results[chr]
        fa_file.close()
        
        logging.info("Writing to your output file.")
        with open("%s/%s.contigINFO.pickle"%(temporary_dir,args.family_mode), "rb") as f:
            contigINFO = pickle.load(f)

        svid = dict()
        svid["INS"] = 0
        svid["DEL"] = 0
        svid["BND"] = 0
        svid["DUP"] = 0
        svid["INV"] = 0
        chroms=sorted(results.keys())
        os.mkdir("%s%s.results"%(temporary_dir,args.family_mode))
        
        for chrom in chroms:
            res_ls = []
            for i in range(len(fam_results)) :
                res_ls.append(fam_results[i][chrom])
            generate_output(args, res_ls, chrom, temporary_dir)
        
        file = open(args.output, 'w')
        Generation_VCF_header(file, contigINFO, args.sample, argv)
        header_line = "#CHROM\tPOS\tID\tREF\tALT\tQUAL\tFILTER\tINFO\tFORMAT"
        for family_member in family_member_ls :
            header_line = header_line + "\t" + family_member
        header_line += "\n"
        file.write(header_line)
        for chrom in chroms:
            pickle_data = []
            with open("%s%s.results/%s.pickle"%(temporary_dir,args.family_mode,chrom), "rb") as f:
                while True:
                    try:
                        lines = pickle.load(f)
                        pickle_data.append(lines)
                    except EOFError:
                        break
            for i in range(len(pickle_data)) :
                for j in range(len(pickle_data[i])) :
                    qual_ls = []
                    filter_ls = []
                    s,t = pickle_data[i][j]
                    file.write(t.replace("<SVID>",str(svid[s])))
                    svid[s]+=1
        file.close()

    if args.execute_stage == 1 or args.retain_work_dir:
        pass
    else:
        logging.info("Cleaning temporary files.")
        if args.Ivcf != None:
            cmd_remove_tempfile = ("rm -r %ssignatures %s*.sigs %s*.pickle"%(temporary_dir, temporary_dir, temporary_dir))
        else:
            cmd_remove_tempfile = ("rm -r %ssignatures %sresults %s*.sigs %s*.pickle"%(temporary_dir, temporary_dir, temporary_dir, temporary_dir))
        exe(cmd_remove_tempfile)

def setupLogging(debug=False):
    logLevel = logging.DEBUG if debug else logging.INFO
    logFormat = "%(asctime)s [%(levelname)s] %(message)s"
    logging.basicConfig( stream=sys.stderr, level=logLevel, format=logFormat )
    logging.info("Running %s" % " ".join(sys.argv))


def run(argv):
    args = parseArgs(argv)
    setupLogging(False)
    starttime = time.time()
    main_ctrl(args, argv)
    logging.info("Finished in %0.2f seconds."%(time.time() - starttime))

if __name__ == '__main__':
    run(sys.argv[1:])
