import sys
import os

#
# Author: Tal Ronnen Oron
# Annotate chromosomal locations (peaks) with closely located genes.
#

def setup():

     if len(sys.argv) < 6:
        print("this script expects the following command line arguments in the order specified below:")
        print("1. peaks file name")
        print("2. gene location file name")
        print("3. chrmosome position in peak file [0 based]")
        print("4. annotated peaks output file name")
        print("5. annotated peaks output format [FLAT/CONDENSED]")
        print("6. maximum distance between a gene and a peak")
        sys.exit()

     peakFN = sys.argv[1].strip()
     try:
        peakf = open(peakFN, "r")
     except:
        print("peak file %s does not exist or cannot be open" % peakFN)
        sys.exit(1)
     geneLocFN = sys.argv[2].strip()
     try:
        geneLocf = open(geneLocFN, "r")
     except:
        print("gene annotation file %s does not exist or cannot be open" % geneLocFN)
        sys.exit(1)

     chromPos = sys.argv[3].strip()
     if chromPos != '0' and chromPos != '1':
        print("Chromosome position in peak file needs to be 0/r", chromPos)
        sys.exit()
     chromPos = int(chromPos)

     outFN = sys.argv[4].strip()
     if os.path.exists(outFN):
        print("output file %s already exists" % outFN)
        overwrite = raw_input("do you want to overwrite it's content? [y|n]")
        if not (overwrite == 'y'):
           sys.exit()
     outf = open(outFN, "w")

     OUTFRM = sys.argv[5].strip()
     if OUTFRM != 'FLAT' and OUTFRM != 'CONDENSED':
        print("output format (fifth argument) should be the string 'FLAT' or 'CONDENSED' and not ", OUTFRM)
        sys.exit()
     
     MAX_DIST = sys.argv[6].strip()
     try:
         MAX_DIST = int(chromPos)
     except:
         print("maximum distance between a peak and a gene (sixth argument) needs to be an integer ", MAX_DIST)
         sys.exit()

     return peakf, geneLocf, chromPos, outf, OUTFRM, MAX_DIST
#
# reads gene annotation file and returns a dictionary that stores for every chromosome it's genes annotations 
#
def get_gene_locations(geneLocf):
  
    genes = {}
    lines = geneLocf.readlines()
    for line in lines:
       sLine = line.strip().split('\t')
       chrom = sLine[1]
       if chrom in genes:
          genes[chrom].append(sLine)
       else:
          genes[chrom] = [sLine]
    return genes


def get_dist(pstart, pend, gstart, gend):

    if (pstart >= gstart and pstart <= gend) or (pend >= gstart and pend <= gend):
        dist = 0
    elif pstart < gstart:
        dist = gstart - pstart
    elif pstart > gend:
        dist = pstart - gend
    else:
        print("Error in finding distance for peak: %s %d %d" % (pchrom, pstart, pend))
        print("inspecting gene %s chrom %s start %d end %d" % (gene, gchrom, gstart, gend))
        sys.exit()
    return dist

def get_closest_gene(index, chrom, pstart, pend):
  
    gene = genes[chrom][index][0]
    gstart = int(genes[chrom][index][2])
    gend = int(genes[chrom][index][3])
    dist = get_dist(pstart, pend, gstart, gend)
    closest_dist = dist
    while dist <= closest_dist:
        closest_dist = dist
        closest_gene = gene
        index += 1
        if index < len(genes[chrom]):
           dist = get_dist(pstart, pend, gstart, gend)
           gene = genes[chrom][index][0]
           gstart = int(genes[chrom][index][2])
           gend = int(genes[chrom][index][3])
        else:
           break
    return index-1, closest_gene, closest_dist   

def get_close_genes(chrom, pstart, pend):


    close_genes = []
    # get distance from the peak and the first gene on the genome
    gene_start_index = 0

    gstart = int(genes[chrom][gene_start_index][2])
    gend = int(genes[chrom][gene_start_index][2])  
    dist = get_dist(pstart, pend, gstart, gend)
    while dist > MAX_DIST and gene_start_index+1 < len(genes[chrom]):
         gene_start_index += 1
         gstart = int(genes[chrom][gene_start_index][2]) 
         gend = int(genes[chrom][gene_start_index][3])  
         dist = get_dist(pstart, pend, gstart, gend)

    index = gene_start_index    
    gene = genes[chrom][index][0]
    gstart = int(genes[chrom][index][2])
    gend = int(genes[chrom][index][3])
    dist = get_dist(pstart, pend, gstart, gend)

    while dist <= MAX_DIST:
        close_genes.append([gene, gstart, gend, dist])
        index += 1
        if index < len(genes[chrom]):
           gene = genes[chrom][index][0]
           gstart = int(genes[chrom][index][2])
           gend = int(genes[chrom][index][3])
           dist = get_dist(pstart, pend, gstart, gend)
        else:
           break
    return close_genes

peakf, geneLocf, chromPos, outf, OUTFRM, MAX_DIST = setup()

genes = get_gene_locations(geneLocf)

for chrom in genes:
   print("chrom %s %d " % (chrom, len(genes[chrom])))

header = peakf.readline()  
if OUTFRM == 'FLAT':
     outf.write("%s\tclosest gene\tstart\tend\tdist\n" % header.strip())
else:
     outf.write("%s\tclosest gene\tstart\tend\tdist\tother close genes\n" % header.strip())

pline = peakf.readline()
index = 0
last_pchrom = None

while pline:
    sPline = pline.strip().split('\t')
    pchrom, pstart, pend = sPline[chromPos:chromPos+3]  
    print(pchrom, pstart, pend)
    if pchrom not in genes:
        print("Warning: Chromosome in peaks file is not in gene annotation file: %s" % pchrom)
        pline = peakf.readline()
        continue
    close_genes = get_close_genes(pchrom, int(pstart), int(pend))
    if len(close_genes):
       close_genes = sorted(close_genes, key=lambda x: x[3], reverse=False)
       if OUTFRM == 'FLAT':
            if len(close_genes):
                for i in range(0, len(close_genes)):
                   outf.write("%s\t" % (pline.strip()))
                   outf.write("%s\t%d\t%d\t%d\n" % (close_genes[i][0], close_genes[i][1], close_genes[i][2], close_genes[i][3]))
            else:
                outf.write("%s" % (pline))
       else:
            outf.write("%s\t" % (pline.strip()))
            if len(close_genes):
               gene = close_genes[0][0]
               dist = close_genes[0][3]
               outf.write("%s\t%d\t%d\t%d\t" % (gene, close_genes[0][1], close_genes[0][2], dist))
               next_close_genes = [] 
               for i in range(1, len(close_genes)):
                   next_close_gene = ','.join(map(str, close_genes[i]))
                   next_close_genes.append(next_close_gene)
               outf.write("%s\n" % ('*'.join(next_close_genes)))
            else:
               outf.write("\n")
       for i in range(0, len(close_genes)):
            gene = close_genes[i][0]
            dist = close_genes[i][3]
    else:
       outf.write("%s\t-\t-\t-\t-\n" % (pline.strip()))

    last_pchrom = pchrom
    pline = peakf.readline()

