import pandas as pd
import os
import sys

def check_dir_exists(dir_path, message):

    if not os.path.exists(dir_path):
        print('Directory {0} does not exist'.format(message))
        exit(0)

def merge_protein_info_from_files(protein_id):

    predictors = ['DISEMBL_HL', 'DISEMBL_LC', 'DISEMBL_R465', 'ESPRITZ_DISPROT', 'ESPRITZ_NMR', 'ESPRITZ_XRAY',
                  'GLOBPLOT', 'IUPRED2A_L', 'IUPRED2A_S', 'ISUNSTRUCT', 'RONN']
    columns = ['AA', 'SS', 'PHI', 'PSI']

    global spider3_dir
    global disorder_dir
    global pondr_dir
    global in_dir
    global dn_dir

    # process spider3 data
    protein = {}

    try:

        f_spd33 = open(spider3_dir + '{}.spd33'.format(protein_id), 'r')

    except FileNotFoundError:

        print('{0}: no Spider3 output file'.format(protein_id))
        exit(0)

    for line in f_spd33:

        if not line.startswith('#'):

            aa_info = line.split()
            protein[int(aa_info[0])] = {'AA': aa_info[1], 'SS': aa_info[2],
                                        'PHI': float(aa_info[4]), 'PSI': float(aa_info[5]),
                                        'DN_LEFT_REPEAT': 0, 'DN_RIGHT_REPEAT': 0,
                                        'IN_LEFT_REPEAT': 0, 'IN_RIGHT_REPEAT': 0}

    f_spd33.close()

    # disorder
    try:

        f_disorder = open(disorder_dir + '{0}.disorder.txt'.format(protein_id), 'r')

    except FileNotFoundError:

        print('{0}: no MassPred output file '.format(protein_id))
        exit(0)

    for line in f_disorder:

        dis_info = line.strip().split()
        start = int(dis_info[3])
        end = int(dis_info[4])
        pred = dis_info[6].replace('DisEMBL_Hot-loops', 'DISEMBL_HL').replace('DisEMBL_Loops/coils', 'DISEMBL_LC').replace('DisEMBL_Remark-465', 'DISEMBL_R465').replace('-', '_').upper()

        if pred != 'VSL2B':
            for i in range(start, end + 1):
                protein[i][pred] = dis_info[5]

    f_disorder.close()

    # pondr
    try:

        f_pondr = open(pondr_dir + '{0}.txt'.format(protein_id), 'r')

    except FileNotFoundError:

        print('{0}: no VLXT output file '.format(protein_id))
        exit(0)

    read_aa = False
    for line in f_pondr:

        if read_aa:

            pondr_info = line.strip().split()
            protein[int(pondr_info[0])]['VLXT'] = 'D' if float(pondr_info[2]) > 0.5 else 'O'

        if line.startswith('Num Res'):
            read_aa = True

    f_pondr.close()
    predictors.append('VLXT')

    # dn repeats
    dn_repeats = []

    try:

        f_dn = open(dn_dir + '{0}.dn.txt'.format(protein_id), 'r')

        for line in f_dn:
            dn_repeat = {}
            dn_info = line.strip().split(',')

            # left repeat
            for i in range(int(dn_info[1]), int(dn_info[2]) + 1):
                protein[i]['DN_LEFT_REPEAT'] = 1

            # right repeat
            for i in range(int(dn_info[3]), int(dn_info[4]) + 1):
                protein[i]['DN_RIGHT_REPEAT'] = 1

            dn_repeat['l_start'] = int(dn_info[1])
            dn_repeat['l_end'] = int(dn_info[2])
            dn_repeat['r_start'] = int(dn_info[3])
            dn_repeat['r_end'] = int(dn_info[4])
            dn_repeat['length'] = int(dn_info[5])
            dn_repeat['l_str'] = dn_info[6]
            dn_repeat['r_str'] = dn_info[7]



            dn_repeats.append(dn_repeat)

        f_dn.close()

    except FileNotFoundError:
        print('{0}: no DN repeats file '.format(protein_id))

    # in repeats
    in_repeats = []
    try:

        f_in = open(in_dir + '{0}.in.txt'.format(protein_id), 'r')

        for line in f_in:
            in_repeat = {}
            in_info = line.strip().split(',')

            # left repeat
            for i in range(int(in_info[1]), int(in_info[2]) + 1):
                protein[i]['IN_LEFT_REPEAT'] = 1

            # right repeat
            for i in range(int(in_info[3]), int(in_info[4]) + 1):
                protein[i]['IN_RIGHT_REPEAT'] = 1

            in_repeat['l_start'] = int(in_info[1])
            in_repeat['l_end'] = int(in_info[2])
            in_repeat['r_start'] = int(in_info[3])
            in_repeat['r_end'] = int(in_info[4])
            in_repeat['length'] = int(in_info[5])
            in_repeat['l_str'] = in_info[6]
            in_repeat['r_str'] = in_info[7]

            in_repeats.append(in_repeat)

        f_in.close()

    except FileNotFoundError:
        print('{0}: no IN repeats file'.format(protein_id))

    instances = []

    for i in range(3, len(protein) - 1):

        instance = {}

        for j in range(-2, 1):
            for column in columns:
                if j != 0:
                    instance['{0}_M_{1}'.format(column, abs(j))] = protein[i + j][column]
                else:
                    instance[column] = protein[i][column]
        for j in range(1, 3):
            for column in columns:
                instance['{0}_P_{1}'.format(column, j)] = protein[i + j][column]

        for predictor in predictors:
            instance[predictor] = 'O'

            for j in range(-2, 3):
                if protein[i + j][predictor] == 'D':
                    instance[predictor] = 'D'

        # dn repeats

        instance['PB_OUT_DN_REPEAT'] = 1
        instance['DN_REPEAT_IN_PB'] = 0
        instance['PB_IN_DN_REPEAT'] = 0
        instance['PB_CENTER_IN_DN_REPEAT'] = 0
        instance['PB_INTERSECT_LEFT_DN_REPEAT'] = 0
        instance['PB_INTERSECT_RIGHT_DN_REPEAT'] = 0
        instance['LEFT_EDGE_LEFT_DN_REPEAT_IN_PB'] = 0
        instance['RIGHT_EDGE_LEFT_DN_REPEAT_IN_PB'] = 0
        instance['LEFT_EDGE_RIGHT_DN_REPEAT_IN_PB'] = 0
        instance['RIGHT_EDGE_RIGHT_DN_REPEAT_IN_PB'] = 0
        instance['PB_INTERSECT_DN_HOMOREPEAT'] = 0

        for j in range(-2, 3):

            if (protein[i + j]['DN_LEFT_REPEAT'] == 1) or (protein[i + j]['DN_RIGHT_REPEAT'] == 1):
                instance['PB_OUT_DN_REPEAT'] = 0

        for dn_rep in dn_repeats:
            if (dn_rep['l_start'] >= i-2 and dn_rep['l_end'] <= i + 2) or (dn_rep['r_start'] >= i-2 and dn_rep['r_end'] <= i + 2):
                instance['DN_REPEAT_IN_PB'] = 1

            if (dn_rep['l_start'] <= i-2 and dn_rep['l_end'] >= i + 2) or (dn_rep['r_start'] <= i-2 and dn_rep['r_end'] >= i + 2):
                instance['PB_IN_DN_REPEAT'] = 1

            if (dn_rep['l_start'] <= i-2 and dn_rep['l_end'] >= i + 2):
                instance['PB_INTERSECT_LEFT_DN_REPEAT'] = 1

            if (dn_rep['r_start'] <= i-2 and dn_rep['r_end'] >= i + 2):
                instance['PB_INTERSECT_RIGHT_DN_REPEAT'] = 1

            #levi kraj ripita pocinje u pb?
            if (dn_rep['l_start'] >= i-2 and dn_rep['l_start'] <= i + 2):
                instance['LEFT_EDGE_LEFT_DN_REPEAT_IN_PB'] = 1

            #desni kraj ripita pocinje u pb?
            if (dn_rep['l_end'] >= i-2 and dn_rep['l_end'] <= i + 2):
                instance['RIGHT_EDGE_LEFT_DN_REPEAT_IN_PB'] = 1

            if (dn_rep['r_start'] >= i-2 and dn_rep['r_start'] <= i + 2):
                instance['LEFT_EDGE_RIGHT_DN_REPEAT_IN_PB'] = 1

            if (dn_rep['r_end'] >= i-2 and dn_rep['r_end'] <= i + 2):
                instance['RIGHT_EDGE_RIGHT_DN_REPEAT_IN_PB'] = 1

            if ((dn_rep['l_end'] >= i-2 and dn_rep['l_start'] <= i + 2) or (dn_rep['r_end'] >= i-2 and dn_rep['r_start'] <= i + 2)) \
                    and (len(set(dn_rep['l_str'])) == 1):
                instance['PB_INTERSECT_DN_HOMOREPEAT'] = 1

        if protein[i]['DN_LEFT_REPEAT'] == 1 or protein[i]['DN_RIGHT_REPEAT'] == 1:
            instance['PB_CENTER_IN_DN_REPEAT'] = 1

        # in repeats

        instance['PB_OUT_IN_REPEAT'] = 1
        instance['IN_REPEAT_IN_PB'] = 0
        instance['PB_IN_IN_REPEAT'] = 0
        instance['PB_CENTER_IN_IN_REPEAT'] = 0
        instance['PB_INTERSECT_LEFT_IN_REPEAT'] = 0
        instance['PB_INTERSECT_RIGHT_IN_REPEAT'] = 0
        instance['LEFT_EDGE_LEFT_IN_REPEAT_IN_PB'] = 0
        instance['RIGHT_EDGE_LEFT_IN_REPEAT_IN_PB'] = 0
        instance['LEFT_EDGE_RIGHT_IN_REPEAT_IN_PB'] = 0
        instance['RIGHT_EDGE_RIGHT_IN_REPEAT_IN_PB'] = 0
        instance['PB_INTERSECT_IN_HOMOREPEAT'] = 0

        for j in range(-2, 3):

            if (protein[i + j]['IN_LEFT_REPEAT'] == 1) or (protein[i + j]['IN_RIGHT_REPEAT'] == 1):
                instance['PB_OUT_IN_REPEAT'] = 0

        for in_rep in in_repeats:
            if (in_rep['l_start'] >= i-2 and in_rep['l_end'] <= i + 2) or (
                    in_rep['r_start'] >= i-2 and in_rep['r_end'] <= i + 2):
                instance['IN_REPEAT_IN_PB'] = 1

            if (in_rep['l_start'] <= i-2 and in_rep['l_end'] >= i + 2) or (
                    in_rep['r_start'] <= i-2 and in_rep['r_end'] >= i + 2):
                instance['PB_IN_IN_REPEAT'] = 1

            if (in_rep['l_start'] <= i-2 and in_rep['l_end'] >= i + 2):
                instance['PB_INTERSECT_LEFT_IN_REPEAT'] = 1

            if (in_rep['r_start'] <= i-2 and in_rep['r_end'] >= i + 2):
                instance['PB_INTERSECT_RIGHT_IN_REPEAT'] = 1

            if (in_rep['l_start'] >= i-2 and in_rep['l_start'] <= i+2):
                instance['LEFT_EDGE_LEFT_IN_REPEAT_IN_PB'] = 1

            if (in_rep['l_end'] >= i-2 and in_rep['l_end'] <= i + 2):
                instance['RIGHT_EDGE_LEFT_IN_REPEAT_IN_PB'] = 1

            if (in_rep['r_start'] >= i-2 and in_rep['r_start'] <= i + 2):
                instance['LEFT_EDGE_RIGHT_IN_REPEAT_IN_PB'] = 1

            if (in_rep['r_end'] >= i-2 and in_rep['r_end'] <= i + 2):
                instance['RIGHT_EDGE_RIGHT_IN_REPEAT_IN_PB'] = 1

            if ((in_rep['l_end'] >= i-2 and in_rep['l_start'] <= i + 2) or (
                    in_rep['r_end'] >= i-2 and in_rep['r_start'] <= i + 2)) \
                    and (len(set(in_rep['l_str'])) == 1):
                instance['PB_INTERSECT_IN_HOMOREPEAT'] = 1

        if protein[i]['IN_LEFT_REPEAT'] == 1 or protein[i]['IN_RIGHT_REPEAT'] == 1:
            instance['PB_CENTER_IN_IN_REPEAT'] = 1

        instances.append(instance)

    df = pd.DataFrame(instances)
    df.to_csv(df_dir + '{0}_for_pbp.csv'.format(protein_id), index=False)

proteins_use = []

path_to_data = 'example_protein_data/'

if len(sys.argv)==2:
    path_to_data=sys.argv[1]

check_dir_exists(path_to_data, 'with data')

spider3_dir = path_to_data + 'spider3/'
check_dir_exists(spider3_dir, 'spider3')

disorder_dir = path_to_data + 'disorders_masspred/'
check_dir_exists(disorder_dir, 'with MassPred disorder outputs')

pondr_dir = path_to_data + 'pondr/'
check_dir_exists(pondr_dir, 'with VLXT output')

dn_dir = path_to_data + 'dn_repeats/'
check_dir_exists(dn_dir, 'with DN repeats')

in_dir = path_to_data + 'in_repeats/'
check_dir_exists(in_dir, 'with IN repeats')

df_dir = 'data_for_pbp/'

if not os.path.exists(df_dir):
    os.makedirs(df_dir)

for file_name in os.listdir(spider3_dir):
    if file_name.endswith('.spd33'):
        prot=file_name.replace('.spd33', '')
        merge_protein_info_from_files(prot)
