#!/usr/bin/python
from __future__ import division
import os
import sys
import time
from lxml import etree
import lxml.etree as ET
from os import popen
from os import system
sys.stdout = os.fdopen(sys.stdout.fileno(), 'w', 0)
URL="http://polarprotpred.ttk.hu"
def msg1(line):
	sys.stderr.write("Usage:")
	sys.stderr.write("'"+str(sys.argv[0])+" -i fasta file [-m binary|categorical -o output_file]\n")
	sys.stderr.write("or\n")
	sys.stderr.write("'"+str(sys.argv[0])+"  --input fasta_file [--mode binary|categorical --output output_file]\n")
	sys.stderr.write("Example: "+str(sys.argv[0])+" -i /home/user/protein.fas -m categorical -o /home/user/protein.txt\n")
	sys.stderr.write("Output location can be omitted to display result on standard output"+"\n")
	sys.stderr.write("Default prediction mode is binary"+"\n")
	sys.stderr.write("After submission, the script will check for results every 30 sec."+"\n")
	sys.stderr.write("Status changes are displayed."+"\n")
def getargs(line):
	infile=""
	outfile=""
	mode="binary"
	for i in range (0, len(line)):
		if line[i][:2]=="-i":
			if i+1>=len(line):
				msg1(line)
				return ["","",False]
			infile=line[i+1]
		if line[i][:7]=="--input":
			if i+1>=len(line):
				msg1(line)
				return ["","",False]
			infile=line[i][line[i].find("=")+1:]
		if line[i][:2]=="-m":
			if i+1>=len(line):
				msg1(line)
				return ["","",False]
			mode=line[i+1]
		if line[i][:7]=="--mode":
			if i+1>=len(line):
				msg1(line)
				return ["","",False]
			mode=line[i][line[i].find("=")+1:]
		if line[i][:2]=="-o":
			if i+1>=len(line):
				msg1(line)
				return ["","",False]
			outfile=line[i+1]
		if line[i][:8]=="--output":
			if i+1>=len(line):
				msg1(line)
				return ["","",False]
			outfile=line[i][line[i].find("=")+1:]
	if outfile!="":
		try:
			f1=open(outfile,"w")
			f1.close()
		except IOError:
			sys.stderr.write("Invalid path, check output file directory path and permissions. Use absolute path."+"\n")
			sys.stderr.write(str(outfile)+"\n")
			return [infile,outfile,mode,False]
	try:
		f1=open(infile,"r")
		f1.close()
	except IOError:
		sys.stderr.write("Invalid path, check input file directory path and permissions. Use absolute path."+"\n")
		sys.stderr.write(str(infile)+"\n")
		return [infile,outfile,mode,False]
	return [infile,outfile,mode,True]
def read(fastafile):
	seq=""
	header=""
	seqs={}
	while 1:
		line=fastafile.readline()
		if line=="":
			break
		if line[0]==">":
			if seq!="" and header!="":
				seq=seq.replace(" ","")
				seqs[header]=seq
				seq=""
			header=line[1:].strip()
		else:
			seq=seq+line.strip()
	seq=seq.replace(" ","")
	header=header.replace('|', '_')
	seqs[header]=seq
	fastafile.close()
	if len(seqs)==0:
		sys.stderr.write("Invalid input file format, please provide a fasta file."+"\n")
	return seqs
def submit(header,proteinsequence,mode):
	jobID={}
	fail=[]
	tmp=popen("wget -qO- \"http://polarprotpred.ttk.hu/direct/id="+header+"@seq="+proteinsequence+"@mode="+mode+"\"")
	content=tmp.readlines()
	ret=content[0]
	time.sleep(1)
	return ret
def poll(header):
    tmp=popen("wget -qO- \"http://polarprotpred.ttk.hu/direct/poll="+header+"\"")
    content=tmp.readlines()
    time.sleep(1)
    return content[0]
def state(query):
	st=[0,0,0,0,0,0]
	for key in query:
		if query[key][1]=="Finished":
			st[0]+=1
		if query[key][1]=="Error":
			st[1]+=1
		if query[key][1]=="Invalid":
			st[2]+=1
		if query[key][1]=="Scheduled":
			st[3]+=1
		if query[key][1]=="Running":
			st[4]+=1
		if query[key][1]=="Invalid amino acid":
			st[5]+=1
	sys.stderr.write("Scheduled: "+str(st[3])+" jobs\n")
	sys.stderr.write("Running: "+str(st[4])+" jobs\n")
	sys.stderr.write("Finished: "+str(st[0])+" jobs\n")
	if st[5]>0:
		sys.stderr.write("Invalid amino acid: "+str(st[5])+" jobs\n")
	if st[2]>0:
		sys.stderr.write("The cluster cannot find the jobid: "+str(st[2])+" jobs\n")
	if st[1]>0:
		sys.stderr.write("PolarProtPred encountered an error while running: "+str(st[1])+" jobs\n")
def res(header):
    tmp=popen("wget -qO- \"http://polarprotpred.ttk.hu/job/results/"+header+"/tab"+"\"")
    content=tmp.readlines()
    time.sleep(1)
    return content

if len(sys.argv)<3:
	msg1(sys.argv)
elif sys.argv[1].lower()=="help" or sys.argv[1].lower()=="-h" or sys.argv[1].lower()=="-help":
	msg1(sys.argv)
else:
	[inf,outf,mode,valid]=getargs(sys.argv)
	if valid==True:
		sequences=read(open(inf,"r"))
		ID={}
		result={}
		if len(sequences)>0:
			for key in sequences:
				ID[key]=[submit(key,sequences[key],mode),""]
				if ID[key][0]=="Invalid amino acid":
					ID[key][1]="Invalid amino acid"
		if len(ID)>0:
			while 1:
				for key in ID:
					if ID[key][1]!="Finished" and ID[key][0]!="Invalid amino acid":
						ID[key][1]=poll(ID[key][0])
				state(ID)
				for key in ID:
					try:
						tmp=result[key]
					except KeyError:
						if ID[key][1]=="Finished":
							result[key]=[]
							result[key]=res(ID[key][0])
				quit=True
				for key in ID:
					if ID[key][1]=="Running" or ID[key][1]=="Scheduled":
						quit=False
				if quit==True:
					break
				time.sleep(30)
		if outf!="":
			o=open(outf,"w")
		for key in result:
			for i in range (1, len(result[key])):
				if outf=="":
					print(result[key][i])
				else:
					o.write(result[key][i]+"\n")
		if outf!="":
			o.close()
