from sklearn.preprocessing import OneHotEncoder
import pandas as pd
from keras.models import load_model
import sys
import os
import numpy as np

"""

Call with arguments:  path_to_model   path_to_input_files
 
Format of name of input files:
- spider3: protein_id.spd33
- disorders info: protein_id.disorder.txt
- pondr disorder: protein_id.txt
- dn repeats: protein_id.dn.txt
- in repeats: protein_id.in.txt
"""

def check_dir_exists(dir_path, message):

    if not os.path.exists(dir_path):
        print('Directory {0} does not exist'.format(message))
        exit(0)

#function for preparation input data
def merge_protein_info(protein_id):

    # variables for preprocessing
    ss_aa_encoder = OneHotEncoder(categories=[['C', 'E', 'H'],
                                              ['A', 'C', 'D', 'E', 'F', 'G', 'H', 'I', 'K', 'L', 'M', 'N', 'P', 'Q',
                                               'R',
                                               'S', 'T', 'V', 'W', 'Y']])

    max_acc = {'A': 106.0, 'R': 248.0, 'N': 157.0, 'D': 163.0, 'C': 135.0, 'Q': 198.0, 'E': 194.0, 'G': 84.0,
               'H': 184.0, 'I': 169.0,
               'L': 164.0, 'K': 205.0, 'M': 188.0, 'F': 197.0, 'P': 136.0, 'S': 130.0, 'T': 142.0, 'W': 227.0,
               'Y': 222.0, 'V': 142.0}


    numeric_columns = [ 'phi', 'psi', 'norm_acc', 'dn_repeat', 'in_repeat','vlxt','disembl_hl','disembl_lc', 'disembl_r465',
                        'espritz_disprot', 'espritz_nmr', 'espritz_xray', 'globplot', 'iupred2a_l', 'iupred2a_s','isunstruct','ronn']

    final_columns_order = ['phi', 'psi', 'norm_acc', 'vlxt', 'disembl_hl', 'disembl_lc', 'disembl_r465',
                           'espritz_disprot', 'espritz_nmr', 'espritz_xray', 'globplot', 'iupred2a_l', 'iupred2a_s',
                           'isunstruct', 'ronn', 'dn_repeat', 'in_repeat', 'aa_A', 'aa_C', 'aa_D', 'aa_E', 'aa_F',
                           'aa_G', 'aa_H', 'aa_I', 'aa_K', 'aa_L', 'aa_M', 'aa_N', 'aa_P', 'aa_Q', 'aa_R', 'aa_S',
                           'aa_T', 'aa_V', 'aa_W', 'aa_Y', 'ss_C', 'ss_E', 'ss_H']

    protein={}
    max_length=1754

    global spider3_dir
    global disorder_dir
    global pondr_dir
    global in_dir
    global dn_dir

    #processing spider3 data
    f_spd33 = open(spider3_dir + '/' + '{}.spd33'.format(protein_id), 'r')

    for line in f_spd33:
        if not line.startswith('#'):
            aa_info=line.split()
            protein[int(aa_info[0])]={

                #get info from spider3
                'aa':aa_info[1], 'ss':aa_info[2],
                'norm_acc': float(aa_info[3])/max_acc[aa_info[1].strip()] if float(aa_info[3])/max_acc[aa_info[1].strip()] < 1.0 else 1.0,
                'phi':(float(aa_info[4])+180.0)/360.0, 'psi':(float(aa_info[5])+180.0)/360.0,

                #set default values for rest of the attributes
                'vlxt':0, 'disembl_hl':0, 'disembl_lc':0, 'disembl_r465':0, 'espritz_disprot':0,
                'espritz_nmr':0, 'espritz_xray':0, 'globplot':0, 'iupred2a_l':0, 'iupred2a_s':0,
                'isunstruct':0, 'ronn':0, 'dn_repeat':0, 'in_repeat':0}



    f_spd33.close()

    #disorder info
    f_disorder = open(disorder_dir + '/' + '{0}.disorder.txt'.format(protein_id), 'r')

    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('-', '_').lower()

        if pred!='vsl2b':
            for i in range(start, end+1):
                protein[i][pred]=1 if dis_info[5].strip()=='D' else 0

    f_disorder.close()

    #pondr
    f_pondr = open(pondr_dir + '/' + '{0}.txt'.format (protein_id), 'r')
    read_aa=False
    for line in f_pondr:

        if read_aa:
            pondr_info = line.strip().split()
            protein[int(pondr_info[0])]['vlxt']= 1 if float(pondr_info[2])>0.5 else 0

        if line.startswith('Num Res'):
            read_aa=True

    f_pondr.close()

    #dn repeats
    try:
        f_dn = open(dn_dir + '/' + '{0}.dn.txt'.format (protein_id), 'r')

        for line in f_dn:
            dn_info=line.strip().split(',')
            #left repeat
            for i in range(int(dn_info[1]), int(dn_info[2])+1):
                protein[i]['dn_repeat']=1
            # right repeat
            for i in range(int(dn_info[3]), int(dn_info[4]) + 1):
                protein[i]['dn_repeat'] = 1

        f_dn.close()

    except FileNotFoundError:
        pass

    # in repeats
    try:
        f_in = open(in_dir + '/' + '{0}.in.txt'.format (protein_id), 'r')

        for line in f_in:
            in_info = line.strip ().split (',')
            # left repeat
            for i in range (int (in_info[1]), int (in_info[2]) + 1):
                protein[i]['in_repeat'] = 1
            # right repeat
            for i in range (int (in_info[3]), int (in_info[4]) + 1):
                protein[i]['in_repeat'] = 1

        f_in.close ()
    except FileNotFoundError:
        pass


    df = pd.DataFrame(protein.values())

    #preparation of categorical attributes
    ss_aa_encoder.fit(df[['ss','aa']])
    dummy_names= ['ss_' + x for x in ss_aa_encoder.categories_[0]] + ['aa_' + x for x in ss_aa_encoder.categories_[1]]
    dummydf = pd.DataFrame(ss_aa_encoder.transform(df[['ss', 'aa']]).toarray(), columns = dummy_names)

    #merging attributes
    df_clean = pd.concat([df, dummydf], axis=1)
    df_clean.drop(['aa', 'ss'], axis=1, inplace=True)
    df_clean[numeric_columns] = df_clean[numeric_columns].apply(pd.to_numeric, downcast='float')
    inst_num, col = df_clean.shape

    #empty records
    empty_record = {}
    for col in final_columns_order:
        empty_record[col]=0

    df_clean = df_clean.append([empty_record]*(max_length-inst_num))
    df_clean.reset_index(inplace=True, drop=True)

    #preparing input in required format
    df_clean['id']=1
    data=np.array(df_clean.groupby('id')[final_columns_order].apply(lambda x: x.values.tolist()).tolist())

    seq=''.join([protein[i]['aa'] for i in range(1, len(protein)+1)])

    return (data, seq)

#function for PBs prediction
def predict_pb(model, data):

    pred_pb = model.predict(data)
    d1, pred_shape_len, t = data.shape

    for i in range(0, d1):

        predicted_sentence = ''

        for l in range(0, pred_shape_len):

            predicted_token_index = np.argmax(pred_pb[i, l, :])

            if predicted_token_index == 0:
                pred_char = '?'
            elif predicted_token_index == 1:
                pred_char = 'Z'
            else:
                pred_char = chr(predicted_token_index - 2 + ord('a'))

            predicted_sentence += pred_char

    return predicted_sentence

proteins_use = []

model50=load_model('models/seq_model50')
model70=load_model('models/seq_model_subsets')

path_to_data = '../examples_of_data_one_chain/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')

for file_name in os.listdir(spider3_dir):
    if file_name.endswith('.spd33'):
        prot=file_name.replace('.spd33', '')
        data, sequence= merge_protein_info(prot)

        result = predict_pb(model50, data)
        print('>ALG', 'seq50:50')
        print('>PROT_ID', prot)
        print('>AA', sequence)
        print('>PB', result[:len(sequence)])
        print()

        result = predict_pb(model70, data)
        print('>ALG', 'seq_subsets')
        print('>PROT_ID', prot)
        print('>AA', sequence)
        print('>PB', result[:len(sequence)])
        print()