from __future__ import division,print_function
import time,re,sys,os,ast,collections
import pandas as pd
import re,csv,codecs


def convert_uniprot_id_and_site_info(file_input,output_file,output_file_contain_empty_site,columns_list,pmidFile):

    #create a dic to map entrez id to uniprotkb ac
    dic_file='./entrez_reviewed_ac_mapping_new.txt'
    reviewed = {}
    with codecs.open(dic_file, 'r', encoding='utf8') as f:
        for line in f:
            line = line.strip()
            line_split=line.split('\t')
            reviewed[line_split[0]]=line_split[1]

    #add pmid filter
    pmid_filter=False
    if len(pmidFile)>0:
        pmidList = pd.read_csv(pmidFile,header=None).iloc[:,0].tolist()
        pmidList=[str(ii) for ii in pmidList]
        pmid_filter=True
    #add site filter, we only consider 6 amino acid
    amino_acid_list=['Asn', 'Arg', 'Ser', 'Thr', 'Tyr', 'Trp']
    #use this set to remove the duplicates
    row_set=set()
    row_set_oger=set()
    row_set_ncbi=set()
    row_set_ces=set()
    with open(output_file_contain_empty_site, 'w') as csvfile:
        spamwriter_oces = csv.writer(csvfile, delimiter='\t', quotechar='|', quoting=csv.QUOTE_MINIMAL)
        spamwriter_oces.writerow(columns_list)
        with open(output_file, 'w') as csvfile:
            spamwriter_o = csv.writer(csvfile, delimiter='\t', quotechar='|', quoting=csv.QUOTE_MINIMAL)


            spamwriter_o.writerow(columns_list)
            with open(file_input) as csvfile:
                spamreader = csv.reader(csvfile, delimiter='\t', quotechar='|')

                for row in spamreader:

                    col_title=row

                    break

                for row in spamreader:
                    new_row=[]
                    from_oger=False
                    for_ncbi=False
                    #row[0] is pmid
                    new_row.append(row[0])
                    #pmid filter
                    if pmid_filter and row[0] not in pmidList:
                        continue
                    #row[4] is normalized protein id
                    if row[4]=='protein':
                        continue

                    if row[5] in reviewed:
                        new_row.append(reviewed[row[5]])
                        new_row.append(row[4])
                    elif len(row[5])>1 and not row[5].isdigit():
                        new_row.append(row[5])
                        new_row.append(row[4])
                        from_oger=True
                    elif len(row[5])>=1 and row[5].isdigit():
                        new_row.append(row[5])
                        new_row.append(row[4])
                        for_ncbi=True
                    else:
                        #print('Not found: ',row[3])
                        #continue
                        new_row.append('')
                        new_row.append(row[4])

                    #deal with the site, row[5] is site string
                    new_row=tuple(new_row)
                    res=split_amino_acid_position(row[6])
                    #print('pmid:',row[0],' ', res)
                    if len(res)==0 and row[6]!='':
                        continue
                    elif len(res)==0 and row[6]=='':
                        #print(row)
                        new_row_list=list(new_row)
                        new_row_list.append('')
                        new_row_list.append('')
                        row_set_ces.add(tuple(new_row_list))
                    for aa,pos in res:
                        new_row_list=list(new_row)
                        new_row_list.append(aa)
                        new_row_list.append(pos)
                        row_set.add(tuple(new_row_list))
                        row_set_ces.add(tuple(new_row_list))
                        if from_oger:
                            row_set_oger.add(tuple(new_row_list))
                        if for_ncbi:
                            row_set_ncbi.add(tuple(new_row_list))
            for ri in row_set:
                if ri[3] in amino_acid_list:
                    rowi=[rii.encode('utf-8') for rii in list(ri) ]
                    spamwriter_o.writerow(rowi)
        for ri in row_set_ces:
            if ri[3] in amino_acid_list or ri[3]=='':
                rowi=[rii.encode('utf-8') for rii in list(ri) ]
                spamwriter_oces.writerow(rowi)
    '''
    with open(output_file.split('.')[0]+'_oger.tsv', 'w') as csvfile:
        spamwriter_oger = csv.writer(csvfile, delimiter='\t', quotechar='|', quoting=csv.QUOTE_MINIMAL)
        spamwriter_oger.writerow(columns_list)
        for ri in row_set_oger:
            spamwriter_oger.writerow(list(ri))
    with open(output_file.split('.')[0]+'_ncbi.tsv', 'w') as csvfile:
        spamwriter_ncbi = csv.writer(csvfile, delimiter='\t', quotechar='|', quoting=csv.QUOTE_MINIMAL)
        spamwriter_ncbi.writerow(columns_list)
        for ri in row_set_ncbi:
            spamwriter_ncbi.writerow(list(ri))
    '''

def convert_uniprot_id_and_site_info_for_test(file_input,output_file,output_file_contain_empty_site,columns_list,pmidFile):

    #create a dic to map entrez id to uniprotkb ac
    dic_file='/home/psu/Documents/biore/data/database/UniProt/entrez_reviewed_ac_mapping_new.txt'
    reviewed = {}
    with codecs.open(dic_file, 'r', encoding='utf8') as f:
        for line in f:
            line = line.strip()
            line_split=line.split('\t')
            reviewed[line_split[0]]=line_split[1]

    #add pmid filter
    pmid_filter=False
    if len(pmidFile)>0:
        pmidList = pd.read_csv(pmidFile,header=None).iloc[:,0].tolist()
        pmidList=[str(ii) for ii in pmidList]
        pmid_filter=True
    #use this set to remove the duplicates
    row_set=set()

    with open(output_file_contain_empty_site, 'w') as csvfile:
        spamwriter_oces = csv.writer(csvfile, delimiter='\t', quotechar='|', quoting=csv.QUOTE_MINIMAL)

        with open(output_file, 'w') as csvfile:
            spamwriter_o = csv.writer(csvfile, delimiter='\t', quotechar='|', quoting=csv.QUOTE_MINIMAL)



            with open(file_input) as csvfile:
                spamreader = csv.reader(csvfile, delimiter='\t', quotechar='|')

                for row in spamreader:

                    col_title=row
                    col_title.append("Final_Output")
                    break

                spamwriter_o.writerow(col_title)
                for row in spamreader:

                    #pmid filter
                    if pmid_filter and row[0] not in pmidList:
                        continue

                    #deal with the site, row[5] is site string
                    res=split_amino_acid_position(row[6])
                    #skip sites
                    if len(res)==0 and row[6]!='':
                        continue

                    #skip empty site
                    if row[6]=='':
                        continue

                    if row[5] in reviewed and len(res)>0:
                        row.append('Yes')
                        row_set.add(tuple(row))
                    else:
                        row.append('No')
                        row_set.add(tuple(row))

            for ri in row_set:
                spamwriter_o.writerow(list(ri))



def split_amino_acid_position(site_string):

    amino_acid_full=['Alanine',  'Arginine',  'Asparagine',  'Aspartic acid',  'Cysteine',  'Glutamic acid',  'Glutamine',  'Glycine',  'Histidine',  'Hydroxyproline',  'Isoleucine',  'Leucine',  'Lysine',  'Methionine',  'Phenylalanine',  'Proline',  'Pyroglutamatic',  'Serine',  'Threonine',  'Tryptophan',  'Tyrosine',  'Valine',]
    amino_acid_shrot=['Ala', 'Arg', 'Asn', 'Asp', 'Cys', 'Glu', 'Gln', 'Gly', 'His', 'Hyp', 'Ile', 'Leu', 'Lys', 'Met', 'Phe', 'Pro', 'Glp', 'Ser', 'Thr', 'Trp', 'Tyr', 'Val',]
    amino_acid_one=['A', 'R', 'N', 'D', 'C', 'E', 'Q', 'G', 'H', 'O', 'I', 'L', 'K', 'M', 'F', 'P', 'U', 'S', 'T', 'W', 'Y', 'V', ]

    amino_acid_full_lower=[i.lower() for i in amino_acid_full]
    amino_acid_one_lower=[i.lower() for i in amino_acid_one]
    amino_acid_shrot_lower=[i.lower() for i in amino_acid_shrot]

    if '\'s' in site_string:
        site_string=site_string.replace('\'s','')

    site_pattern='([a-zA-Z]+)\s?-?\(?(\d+)'
    number_pattern='\D*(\d+)\D*'
    result=re.findall(site_pattern,site_string)
    number_result=re.findall(number_pattern,site_string)

    res=[]
    if result:
        #print(result)
        for aa,pos in result:
            aa=aa.strip()
            if len(aa)==3 and aa.lower() in amino_acid_shrot_lower:
                res.append((aa,pos))
            elif len(aa)==1:
                if aa.lower() in amino_acid_one_lower:
                    res.append((amino_acid_shrot[amino_acid_one_lower.index(aa.lower())],pos))

            elif len(aa)>3:
                sign_found=-1
                for aai in amino_acid_full_lower:
                    if aai in aa.lower():
                        sign_found=1
                        res.append((amino_acid_shrot[amino_acid_full_lower.index(aai)],pos))


    #more position than amino acid
    res_extra=[]
    if len(res)<len(number_result):

        for ri in res:
            for ni in number_result:
                res_extra.append((ri[0],ni))
    return list(set(res+res_extra))


def substract_file(file1,file2,file3):

    with open(file3, 'w') as csvfile:
        spamwriter_o = csv.writer(csvfile, delimiter='\t', quotechar='|', quoting=csv.QUOTE_MINIMAL)


        row_list1=[]
        row_list2=[]
        with open(file1) as csvfile:
            spamreader = csv.reader(csvfile, delimiter='\t', quotechar='|')

            for row in spamreader:

                col_title=row
                #print(col_title)
                break
            for row in spamreader:
                #row=[ri.decode('utf-8') for ri in row ]
                row_list1.append(row)
        with open(file2) as csvfile:
            spamreader = csv.reader(csvfile, delimiter='\t', quotechar='|')

            for row in spamreader:
                col_title=row
                #print(col_title)
                break
            for row in spamreader:
                #row=[ri.decode('utf-8') for ri in row ]
                row_list2.append(row)
        #write column names
        spamwriter_o.writerow(col_title)
        for ri in row_list1:
            if ri not in row_list2:
                #ri=[rri.encode('utf-8') for rri in ri ]
                spamwriter_o.writerow(ri)


if __name__ == '__main__':
    #test

    task_name='glygen_full_length_101820'

    file_input=task_name+'_site_fusion_final.tsv'
    output_file=task_name+'_site_fusion_output.tsv'
    output_file_contain_empty_site=task_name+'_site_fusion_output_contain_empty_site.tsv'
    columns_list=['pmid','uniprotkb_ac','protein','amino_acid','site_position']
    convert_uniprot_id_and_site_info(file_input,output_file,output_file_contain_empty_site,columns_list,'')

    '''
    task_name='glygen_testset_082620'
    file_input=task_name+'_final.tsv'
    output_file=task_name+'_output.tsv'
    output_file_contain_empty_site=task_name+'_output_contain_empty_site.tsv'
    columns_list=['pmid','uniprotkb_ac','amino_acid','site_position']
    convert_uniprot_id_and_site_info(file_input,output_file,output_file_contain_empty_site,columns_list,'')


    file_input=task_name+'_plus_site_fusion_final.tsv'
    output_file=task_name+'_plus_site_fusion_output.tsv'
    output_file_contain_empty_site=task_name+'_plus_site_fusion_output_contain_empty_site.tsv'
    columns_list=['pmid','uniprotkb_ac','amino_acid','site_position']
    convert_uniprot_id_and_site_info(file_input,output_file,output_file_contain_empty_site,columns_list,'')

    file_input=task_name+'_plus_site_fusion_final.tsv'
    output_file=task_name+'_plus_site_fusion_output_glyconnect.tsv'
    output_file_contain_empty_site=task_name+'_plus_site_fusion_output_contain_empty_site_glyconnect.tsv'
    columns_list=['pmid','uniprotkb_ac','amino_acid','site_position']
    convert_uniprot_id_and_site_info(file_input,output_file,output_file_contain_empty_site,columns_list,'glyconnect.txt')

    
    file1=task_name+'_plus_site_fusion_output.tsv'
    file2=task_name+'_plus_site_fusion_output_.tsv'
    file3=task_name+'_protein_in_titile_or_first_sentence.tsv'
    substract_file(file1,file2,file3)
    '''