from subprocess import call, check_call
import subprocess
import os
import numpy as np
import matplotlib.pyplot as plt


vname = '1a34'
n_mpi_proc = 8
n_modes = 30
n_prot = 60
min_subdivisions = 8

# TODO: downloads virus

# TODO: check the file exists

# TODO: coarse graining

# create directory structure and unpack virus

vdb_file_name = vname+'_full.vdb'

os.system('mkdir '+vname)

os.system("gunzip -c " + vdb_file_name + '.gz > '+vname+'/'+vdb_file_name)
#check_call(['gunzip',vdb_file_name+'.gz'])
#check_call(['mv',vdb_file_name,vname])
os.chdir(vname)

os.system('mkdir pdb_files')
pdb_file_name = vname+'_ca.pdb'
#proc = subprocess.Popen(['grep','" CA"',vdb_file_name], 0, None, subprocess.PIPE, subprocess.PIPE, None)
#f = open('pdb_files/'+pdb_file_name)
#for l in proc.stdout.readlines():
#    f.write(l)
os.system('grep " CA" '+vdb_file_name+' > pdb_files/'+pdb_file_name)

subprocess.call(['ln','-s','pdb_files/'+pdb_file_name])
betagm_params = 'Int_Range 7.5 \n'+ \
                'V_PEPT    2.0 \n'+ \
                'V_CACA    1.0 \n'+ \
                'V_CACB    1.0 \n'+ \
                'V_CBCB    1.0 \n'+ \
                'DUMP_EIGENVALUES                  1 \n'+ \
                'N_TOP_SLOW_EIGENVECT              '+ str(n_modes+10)+ '\n'+ \
                'DUMP_FULL_COVMAT                  1 \n'+ \
                'DUMP_REDUCED_COVMAT               1 \n'+ \
                'DUMP_NORMALISED_REDUCED_COVMAT    1 \n'+ \
                'DUMP_MEAN_SQUARE_DISPL            1 \n'

f=open('PARAMS_BETAGM.DAT','w')
f.write(betagm_params)
f.close()

f=open('PDB_FILES.TXT','w')
f.write(pdb_file_name)
f.close()

check_call(['betaGM-lowmem-slepc'])
check_call(['PiSQRD++','-f',pdb_file_name,'-j',str(n_mpi_proc),'-m',str(n_modes), 
            '--cycle',str(min_subdivisions),str(n_prot),'1'])

#ANALYSIS

# a. Size of domains

ntypes_min_jump = 5
ntypes_limit = 5
ntypes_cutoff = 0.03

ntypes = [1 for i in range(0,n_prot)]

for i in range(min_subdivisions,n_prot):
    sizes = [0 for j in range(0,i)]
    fname = pdb_file_name + '_' + str(i) + 'domains.dat'
    f = open(fname,'r')
    for l in f.readlines():
        if(len(l.split()) == 0):
            continue
        idx = int(l.split()[0])
        dom = int(l.split()[1])
        sizes[dom] += 1
    f.close()
    f = open(pdb_file_name + '_' + str(i) + 'domains_sizes.dat','w')
    for j in range(0,i):
        f.write(str(j) + ' ' + str(sizes[j]) + '\n')
    f.close()
    ave_size = np.mean(sizes)
    sizes = sorted(sizes)
    tile_start = sizes[0]
    for j in range(1,i):
        if ( (sizes[j]-tile_start) > ave_size*ntypes_cutoff ) and \
           ( (sizes[j]-tile_start) > ntypes_min_jump ):
            ntypes[i] +=1
            tile_start = sizes[j]
        if ntypes[i] >= ntypes_limit:
            break

f = open('nTypes.dat','w')
for i in range(min_subdivisions,n_prot):
    f.write(str(i) + ' ' + str(ntypes[i]) + '\n')
f.close()


# b. Integrity

integrities = [0 for i in range(0,n_prot)]
for i in range(min_subdivisions,n_prot):
    proc = subprocess.Popen(
            ['integrity-score3',pdb_file_name,pdb_file_name+'_'+str(i)+'domains.dat'], 
            0, None, subprocess.PIPE, subprocess.PIPE, None)
    integrities[i] = float(proc.stdout.read().strip())
    
f= open('integrity.dat','w')
for i in range(min_subdivisions,n_prot):
    f.write(str(i) + ' ' + str(integrities[i]) + '\n')
f.close()

# c. Strain

strains = [0 for i in range(0,n_prot)]

for i in range(min_subdivisions,n_prot):
    f = open(pdb_file_name + '_' + str(i) + 'domains_energies.dat','r')
    for l in f.readlines():
        if(len(l.split()) == 0):
            continue
        idx = int(l.split()[0])
        strains[i] += float(l.split()[1])
    f.close()

f= open('strain.dat','w')
for i in range(min_subdivisions,n_prot):
    f.write(str(i) + ' ' + str(strains[i]) + '\n')
f.close()



xx = range(min_subdivisions,n_prot)


g=plt.subplot(3, 1, 1, yscale ='log')
plt.plot(xx, strains[min_subdivisions:n_prot], 'ko-')
plt.ylim([min(strains[min_subdivisions:n_prot]),max(strains[min_subdivisions:n_prot])])
plt.ylabel('Strain')

plt.subplot(3, 1, 2)
plt.plot(xx, integrities[min_subdivisions:n_prot], 'ko-')
plt.ylabel('Integrity')

plt.subplot(3, 1, 3)
plt.ylim([0,6])
plt.plot(xx, ntypes[min_subdivisions:n_prot], 'ko-')
plt.xlabel('Number of domains(Q)')
plt.ylabel('# of types')

plt.show()
