#!/usr/bin/env python3
import sys, math

'''
get statistics :"Sensitivity","Specificity","PPV","NPV","Accuracy" of coding test result file.

coding test result file:
name    actCls    method1    method2    method3    method4    ...
ENSMUST00000133570.7_12345678    1       1    1    1    1
ENSMUST00000176016.1_12345678    1       0    1    0    1
ENSMUST00000101121.2_12345678    1       1    1    1    1
ENSMUST00000191996.1_12345678    0       0    0    0    0
ENSMUST00000141428.7_87654321    1       1    0    1    1
ENSMUST00000065383.10_87654321   1       0    0    0    1
ENSMUST00000209459.1_87654321    0       0    0    0    0
ENSMUST00000128591.1_87654321    1       1    1    1    1

'''

HEADER=["species","database","length","transcript_type","methods","TP","FP","TN","FN","Sensitivity","Specificity","PPV","NPV","Accuracy","MCC"]
HEADER1=["species","ensemble_type","methods","TP","FP","TN","FN","Sensitivity","Specificity","PPV","NPV","Accuracy","MCC"]

def file2name(filename):
    base_name = filename.split("/")[-1]
    name = base_name.split(".")[0]
    name_array = name.split("_")
    species = "_".join(name_array[:-1])
    esb_type = name_array[-1]
    return species, esb_type

def perf_measure(y_pred,y_actual):

    TP=0
    FP=0
    TN=0
    FN=0
    #y_pred = list(y_pred)
    #print("perf_measure: y_pred,y_actual")
    #print(y_pred)
    #print(y_actual)
    for i in range(len(y_pred)):
        if y_actual[i] == y_pred[i] == 1:
            TP += 1
        elif y_actual[i] == y_pred[i] == 0:
            TN += 1
        elif y_actual[i] == 1:#y_pred == 0
            FN += 1
        else:
            FP += 1
    return TP,FP,TN,FN

def perf_measures(y_preds,y_actual):
    data = []
    #y_preds = list(y_preds)
    #print("perf_mersures-ypred y_actual")
    #print(y_preds)
    #print(y_actual)
    # every method
    for i in range(len(y_preds[0])):
        data.append(perf_measure(list(map(lambda x:x[i],y_preds)),y_actual))
#    for y_pred in y_preds:
#        data.append(perf_measure(y_pred,y_actual))
    return data

def read_array(infile,has_header=True,sep="\t"):
    
    # data:{key:[[y_pred],[y_actual]]}
    # y_actual = [[prediction on gene1],[prediction on gene2],...]
    c_name = 0
    c_actual = 1
    data = [[],[]]
    try:infile=open(infile)
    except:pass
    if has_header:
        methods=infile.readline().strip().split(sep)[2:]
    for line in infile:
        line_array=line.strip().split(sep)
        y_pred = list(map(lambda x:int(x),line_array[2:]))
        y_actual = int(line_array[c_actual])
        data[0].append(y_pred)
        data[1].append(y_actual)
    return methods, data

def perf_stat(TP,FP,TN,FN):
    def div(a,b):
        if b==0:
            return "NULL"
        else:
            return "%.3f"%(float(a)/b)
        
    Sen=div(TP,(TP+FN))#float(TP)/(TP+FN)
    Spe=div(TN,(TN+FP))#float(TN)/(TN+FP)
    PPV=div(TP,(TP+FP))#float(TP)/(TP+FP)
    NPV=div(TN,(TN+FN))#float(TN)/(TN+FN)
    Acc=div((TP+TN),(TP+TN+FP+FN))#float(TP+TN)/(TP+TN+FP+FN)
    Mcc=div((TP*TN+FP*FN),(math.sqrt((TP+FP)*(TP+FN)*(TN+FP)*(TN+FN))))
    return Sen,Spe,PPV,NPV,Acc,Mcc

def perf_stats(data):
    outdata = []
    for e in data:
        outdata.append(perf_stat(*e))
    return outdata

def stat_array(infile,outfile,sep="\t",write_header=True,add_filename=False):
    methods, data = read_array(open(infile))
    if write_header:
        outfile.write(sep.join(HEADER)+"\n")
    y_pred, y_actual = data
    mes = perf_measures(y_pred,y_actual)
    i = 0
    for st in perf_stats(mes):
        method=methods[i]
        if add_filename:
            outdata = [*file2name(infile), method]
        else:
            outdata = [method]
        outdata.extend(mes[i])
        outdata.extend(st)
        outfile.write(sep.join(map(str,outdata))+"\n")
        i+=1

def stat_arrays(infiles,outfile,sep="\t"):
    outfile.write(sep.join(HEADER1)+"\n")
    for infile in infiles:
        stat_array(infile,outfile,sep,write_header=False,add_filename=True)

if __name__ == '__main__':
        
    infile=sys.argv[1]
    infiles=sys.argv[1:]
    outfile=sys.stdout
    #stat_array(infile,outfile)
    stat_arrays(infiles,outfile)
