# -*- coding: utf-8 -*-
"""
Created on Tue Mar 23 14:50:33 2021

@author: jhoffman
"""
import sys
import numpy as np
from datetime import date
from calendar import monthrange
import glob
from PIL import Image, ImageOps, ImageChops, ImageEnhance
from matplotlib.colors import ListedColormap
import matplotlib.pyplot as plt
import os
import xarray as xr
from data import *
from model import *
from tensorflow.keras.models import load_model
from datetime import datetime
import h5py

n = len(sys.argv)
print("Total arguments passed:", n)
if n == 5 :
   print(sys.argv[1])
   print(sys.argv[2])
   print(sys.argv[3])
   print(sys.argv[4])

   yy=int(sys.argv[1])
   mm=int(sys.argv[2])
   dd=int(sys.argv[3])
   s=int(sys.argv[4])
else :
   quit()
print(yy,mm,dd,s)

now = datetime.now()
current_time = now.strftime("%H:%M:%S")
print("Start Time =", current_time)


xsize = 9024
ysize = 9024
xsize=xsize-2000+1
ysize=ysize-2000+1
blocksize=512

debug=0

land=Image.open("land.mask.png")
land=np.array(land)
im = Image.fromarray(land)
land=ImageOps.grayscale(im)
land=np.array(land)
land=land > 0

water=Image.open("water.mask.png")
water=np.array(water)
im = Image.fromarray(water)
water=ImageOps.grayscale(im)
water=np.array(water)
water=water > 0

det_thresh=50
if s == 0:
   sat='VIIRS'
   sn='[NJ][P0][P1]'
   det_thresh=45
   model = load_model("unet_lead.VIIRS.h5")
if s == 1:
   sat='MODIS'
   sn='M*D'
   det_thresh=32
   model = load_model("unet_lead.MODIS.h5")
s=str(s)            
doy=date(int(yy),int(mm),int(dd)).timetuple().tm_yday
mm=str(mm).zfill(2)
dd=str(dd).zfill(2)
doy=str(doy).zfill(3)
yy=str(yy).zfill(4)
fn= glob.glob('/apollo/cloud/scratch/jhoffman/leads/'+s+'/'+yy+'/'+doy+'/'+yy+doy+'*'+sn+'.BT.png')
fn=sorted(fn)
sum_result=np.zeros((xsize,ysize))
max_result=np.zeros((xsize,ysize))
total_result=np.zeros((xsize,ysize))
sum_resultc=np.zeros((xsize,ysize))
sum_resultc.fill(1)
sum_resultt=np.zeros((xsize,ysize))
hhnew=100
hhi='00'
if len(fn) == 0:
   print("no valid files")
   end = datetime.now()
   current_time = end.strftime("%H:%M:%S")
   print("End Time =", current_time,end-now)

for i in range(len(fn)): 
   print(sat,mm,dd,yy)
   hh=fn[i].split('.')[-4]
   print(i, fn[i], hh)
   if hh == '24':
      continue
   if int(hh) >= hhnew:
      hhnew=(int(hhi)+1)*100
      hhi=int(hhi)+1
      hhi=str(hhi)

   testdir='/apollo/cloud/scratch/jhoffman/leads/'+s+'/'+yy+'/'+doy+'/'+sat+'.'+hh+'/test'
   resultdir='/apollo/cloud/scratch/jhoffman/leads/'+s+'/'+yy+'/'+doy+'/'+sat+'.'+hh+'/result'
   if not os.path.exists(testdir):
      os.makedirs(testdir)
   if not os.path.exists(resultdir):
      os.makedirs(resultdir)
   if os.path.exists(testdir+"/0.png"):
      for fl in glob.glob(testdir+"/*.png"):
         os.remove(fl)
   if os.path.exists(resultdir+"/0_predict.png"):
      for fl in glob.glob(resultdir+"/*.png"):
         os.remove(fl)
      
   leadpic=Image.open(fn[i])
   BT=np.array(leadpic).astype(np.uint8)
   BT=(np.array(BT))*(np.array(BT < 255))
   im = Image.fromarray(BT)
   im=ImageOps.grayscale(im)
   enhancer =ImageEnhance.Contrast(im)
   im.save('/apollo/cloud/scratch/jhoffman/leads/'+s+'/'+yy+'/'+doy+'/'+sat+'.'+hh+'/'+yy+doy+"python.test.png")
   BT = np.asarray(im)
   c=0
   blockbuffer=50
   halfbuffer=24
   xi=int(xsize/(blocksize-blockbuffer))
   yi=int(ysize/(blocksize-blockbuffer))
   for ii in range(xi+1): 
      for jj in range(yi+1):
         istep0=(ii*blocksize)-(ii*blockbuffer)
         jstep0=jj*blocksize-(jj*blockbuffer)

         istep1=istep0+blocksize-1
         jstep1=jstep0+blocksize-1
         if istep1 > xsize-1:
            istep1=xsize-1
         if jstep1 > ysize-1:
            jstep1=ysize-1
         iimg=np.zeros((blocksize,blocksize))
         iimg=BT[istep0:istep1,jstep0:jstep1]       
         landtest=np.array(land[istep0:istep1,jstep0:jstep1]) 
         if np.amax(iimg) == 0 or np.amin(iimg) == 255 or np.amax(iimg) == np.amin(iimg):
            continue
         img = Image.fromarray(iimg)
         img=ImageOps.grayscale(img)
         enhancer =ImageEnhance.Contrast(img)
         img = enhancer.enhance(3)
         img.save(testdir+"/"+str(c)+".png")
         c+=1
   print(c)
   if c == 0 :
      continue
   testLeads = testGenerator(testdir)
   results = model.predict_generator(testLeads,c)
   saveResult(resultdir,results)
   c=0
   result=np.zeros((xsize,ysize))
   resultc=np.zeros((xsize,ysize))
   resultt=np.zeros((xsize,ysize))
   ultc=np.zeros((xsize,ysize))
   ultt=np.zeros((xsize,ysize))
   if os.path.exists(resultdir+"/0_predict.png"):
      for ii in range(xi+1):
         for jj in range(yi+1):
            istep0=(ii*blocksize)-(ii*blockbuffer)
            jstep0=jj*blocksize-(jj*blockbuffer)
            istep1=istep0+blocksize-1
            jstep1=jstep0+blocksize-1
            if istep1 > xsize-1:
               istep1=xsize-1
            if jstep1 > ysize-1:
               jstep1=ysize-1
            iimg=np.zeros((blocksize,blocksize))
            iimg=BT[istep0:istep1,jstep0:jstep1]   
            if np.amax(iimg) == 0 or np.amin(iimg) == 255 or np.amax(iimg) == np.amin(iimg):
               continue
            BTi=Image.open(testdir+'/'+str(c)+'.png')
            BTi=np.array(BTi)
            im = Image.fromarray(BTi)
            BTi=ImageOps.grayscale(im)
            BTi=np.array(BTi)
            pic=Image.open(resultdir+'/'+str(c)+'_predict.png')
            tresult=np.array(pic)
            im = Image.fromarray(tresult)
            tresult =ImageOps.grayscale(im)
            tresult =tresult.resize((blocksize,blocksize), Image.BILINEAR)
            tresult=np.array(tresult)
               
            result[istep0+(halfbuffer):istep1-(halfbuffer),jstep0+(halfbuffer):jstep1-(halfbuffer)]=tresult[(halfbuffer):istep1-istep0-(halfbuffer),(halfbuffer):jstep1-jstep0-(halfbuffer)]
            resultt[istep0+(halfbuffer):istep1-(halfbuffer),jstep0+(halfbuffer):jstep1-(halfbuffer)]=BTi[(halfbuffer):istep1-istep0-(halfbuffer),(halfbuffer):jstep1-jstep0-(halfbuffer)]
            resultc[istep0+(halfbuffer):istep1-(halfbuffer),jstep0+(halfbuffer):jstep1-(halfbuffer)]=(BT[istep0+(halfbuffer):istep1-(halfbuffer),jstep0+(halfbuffer):jstep1-(halfbuffer)] > 0)*(BT[istep0+(halfbuffer):istep1-(halfbuffer),jstep0+(halfbuffer):jstep1-(halfbuffer)] < 255)
            c+=1
      result = np.flipud(result)
      resultt = np.flipud(resultt)
      resultc = np.flipud(resultc)

      sum_result=np.array(sum_result) + np.array((result > det_thresh))
      max_result=np.array(np.where(result > max_result, result, max_result))
      total_result=np.array(np.where(result > 0, result, total_result))
      sum_resultt=np.array(np.where(resultt > sum_resultt, resultt, sum_resultt))

      sum_resultc=np.array(sum_resultc)+ np.array(resultc)
       
      if debug == 1: 
         img = Image.fromarray(result)
         img=ImageOps.grayscale(img)
         img.save("/apollo/cloud/scratch/jhoffman/leads/"+s+'/'+yy+'/'+doy+"/"+yy+doy+"."+hh+"."+sat+".test.result.png")              

         img = Image.fromarray(resultt)
         img=ImageOps.grayscale(img)
         img.save("/apollo/cloud/scratch/jhoffman/leads/"+s+'/'+yy+'/'+doy+"/"+yy+doy+"."+hh+"."+sat+".test.result.t.png")              


         result=np.array(result)
         mask=np.array(2*(result > det_thresh-25))+np.array((result > det_thresh))+np.array((result > det_thresh+25))
         cmap=ListedColormap(['#323232','#ff3300','#ff0000','#00ff00','#ffffff']) 
         plt.imsave("/apollo/cloud/scratch/jhoffman/leads/"+s+'/'+yy+'/'+doy+"/"+yy+doy+"."+sat+"."+hh+".mask.result.png",mask,format="png",cmap=cmap)
                   
         print("Sum max ", np.amax(sum_result))
         sum2=np.array(1*(sum_result == 1))+np.array(2*(sum_result == 2))+np.array(3*(sum_result >= 3))

         mask=np.array((4*land)+sum2)
         cdict1=ListedColormap(['#323232','#ff0000','#00ff00','#ffffff','#643200','#643200','#643200','#643200','#643200'])
         plt.imsave("/apollo/cloud/scratch/jhoffman/leads/"+s+'/'+yy+'/'+doy+"/"+yy+doy+sat+"."+hh+".composit.mask.result.png",mask,format="png",cmap=cdict1)

         sum_resultti=np.array(sum_resultt)/(np.array(sum_resultc))
         sum_resultti=(255./np.amax(sum_resultti))*sum_resultti
         img = Image.fromarray(sum_resultti)
         img=ImageOps.grayscale(img)
         img.save("/apollo/cloud/scratch/jhoffman/leads/"+s+'/'+yy+'/'+doy+"/"+yy+doy+"."+hhi+"."+sat+".sum.result.t.png")              

         sum_resultci=(255./np.amax(sum_resultc))*sum_resultc
         img = Image.fromarray(sum_resultci)
         img=ImageOps.grayscale(img)
         img.save("/apollo/cloud/scratch/jhoffman/leads/"+s+'/'+yy+'/'+doy+"/"+yy+doy+"."+hhi+"."+sat+".sum.c.png")              

      tstep = datetime.now()
      current_time = tstep.strftime("%H:%M:%S")
      print("Time =", current_time,tstep-now)

sum2=np.array(1*(sum_result >= 3))
mask=np.array((2*land)+np.array(sum2*(water == 1)))

cdict1=ListedColormap(['#000000','#ffffff','#ffffff','#323232'])
plt.imsave("case/"+yy+doy+sat+".mask.result.png",mask,format="png",cmap=cdict1)

if debug == 1: 
   sum_resultt=np.array(sum_resultt)/(np.array(sum_resultc))
   img = Image.fromarray(sum_resultt)
   img=ImageOps.grayscale(img)
   img.save("case/"+yy+doy+"."+sat+".mean.BT.png")              

sum_result=np.array(water)*np.array(sum_result)
max_result=np.array(water)*np.array(max_result)
total_result=np.array(water)*np.array(total_result)
sum_resultt_result=np.array(water)*np.array(sum_resultt)
sum_resultc=np.array(sum_resultc-1)
sum_resultc=np.array(water)*np.array(sum_resultc)
sum_result=np.array(sum_result)
sum_result=sum_result.astype(int)
sum_resultc=sum_resultc.astype(int)


outfn="case/"+yy+doy+sat+".result.h5"
hf = h5py.File(outfn, 'w')
hf.create_dataset('lead_count', data=sum_result, compression="gzip", compression_opts=9)
hf.create_dataset('coverage_count', data=sum_resultc, compression="gzip", compression_opts=9)
hf.create_dataset('max_result', data=max_result, compression="gzip", compression_opts=9)
hf.create_dataset('mask', data=mask, compression="gzip", compression_opts=9)
hf.close()

end = datetime.now()
current_time = end.strftime("%H:%M:%S")
print("End Time =", current_time,end-now)
