from darepype.drp import DataFits # pipeline data object class
from darepype.drp.stepmiparent import StepMIParent # pipestep Multi-Input parent
from darepype.tools.steploadaux import StepLoadAux #pipestep steploadaux object class
from astropy.io import fits #package to recognize FITS files
import numpy as np
import logging
from astropy.stats import mad_std # The median absolute deviation, a more robust estimator than std
from astropy.time import Time
import os


class StepMasterFlatCCD(StepLoadAux, StepMIParent):

	def __init__(self):
		
		super(StepMasterFlatCCD, self).__init__()
		
		self.flatloaded = False
		self.flat = None
		self.flatname = ''

		self.darkloaded = False
		self.dark = None
		self.darkname = ''

		self.biasloaded = False
		self.bias = None
		self.biasname = ''

		self.log.debug('Init: done')
	
	def setup(self):
	
		self.name = 'masterflatccd'
		self.procname = 'MFLAT'

		self.log = logging.getLogger('pipe.step.%s' % self.name)
	
		self.paramlist = []

		self.paramlist.append(['combinemethod', 'median',
				       'Specifies how the files should be combined - options are median, average, sum'])
		self.paramlist.append(['outputfolder', '',
			               'Output directory location - default is folder with input files'])
		self.paramlist.append(['hotpxlim', 99.5, 'Hot pixel limit percentile'])
		self.paramlist.append(['outputfolder2', '', 'Alternate output directory path'])
		self.paramlist.append(['gainpcntlim', 0.3, 'gain quality threshold'])
		self.paramlist.append(['dstdpcntlim', 10.0, 'dstd quality threshold'])
		self.paramlist.append(['numfilelim', 8, 'Minimum number of input files'])
		self.paramlist.append(['print_switch', False, 'Set True to turn on print statements'])

		self.loadauxsetup('mdark')
		self.loadauxsetup('mbias')

	def timesortHDR(self, datalist, date_key = 'date-obs'):
		''' hbgvbhj
                '''
	
		date_obs = []		
		for d in datalist:
			if '_bin1L' in d.filename:
				head = d.getheader(d.imgnames[1])
				date_obs.append(head[date_key])
			else:
				head = d.getheader()
				date_obs.append(head[date_key])
		t = Time(date_obs, format='isot', scale='utc')
		tsort = np.argsort(t)
		tfiles = []
		utime = []
		for i in tsort:
			tfiles.append(datalist[i])
			utime.append(t[i].unix)
		return tfiles, utime

	def read_oneDF(whichfiles, whichfile, whichpath, convert_raw=False):
	
		fitsfilename = os.path.join(whichpath, whichfiles[whichfile])
		if 'bin1' in fitsfilename and '_RAW.fit' in whichfiles[whichfile]:
			ddf = DataFits(config=config)
			ddf.load(fitsfilename)
			df = DataFits(config=config)
			df.header = ddf.getheader(ddf.imgnames[1]).copy()
			del df.header['xtension']
			df.header.insert(0, ('simple', True, 'file does conform to FITS standard'))
			df.imageset(ddf.imageget(ddf.imgnames[1]))
		else:
			df = DataFits(config=config)
			df.load(fitsfilename)

					
	def run(self):
	
		pt = self.getarg('print_switch')

		hotpxlim = self.getarg('hotpxlim')
		if pt: print('hotpxlim =', hotpxlim)
		if pt: print('')
		

		flatpath = "/data/images/Temp/Users/levi/MFLAT Step/RAW FLATS"		
		flatfiles = [f for f in os.listdir(flatpath) if '.fit' in f and 'g-band' in f and 'RAW.fit' in f and '._' not in f]
		flatfiles, utimeH = self.timesortHDR(self.datain, date_key = 'date-obs')
		
		
		rows, cols = flat.shape[0], flat.shape[1]

		
		numflats = len(flatfiles)
		flatheadlist, flatstats = [[], []], [[]]
		flatimage = np.zeros((numflats, rows, cols))

		
		flatbaseheader = flatheadlist[0]


		flatexptimes = []
		imedian = flatstats['median']
		for j in range(len(flatfiles)):
			exptime = flatheadlist[j].header['exptime']
			flatexptimes.append(exptime)
		xtimemin = np.min(flatexptimes)
		xtimemax = np.max(flatexptimes)
		if pt: print('xtimemin =', xtimemin, 'xtimemax =', xtimemax)



		darkimage = np.zeros_like(flatimage)
		for j in range(flatimage.shape[0]):
			darkimage[j] = ((dark - bias) * (flatexptimes[j])/darkexptime) + bias

		flatimageDS = flatimage - darkimage


		flatimageDSN = np.zeros_like(flatimageDS)
		flatmediansDS = np.zeros((numflats))
	
		for j in range(numflats):
			flatimageDS[j][hotpix] = np.nan
			flatmediansDS[j] = np.nanmedian(flatimageDS[j])
			flatimageDSN[j] = flatimageDS[j] / flatmediansDS[j]


		flat = np.nanmedian(flatimageDSN, axis=0)
		flatmedian, flatmean, flatstd, flatmadstd = [0], [0], [0], [0]

		flatmedian = np.nanmedian(flat)
		flatmean = np.nanmean(flat)
		flatstd = np.nanstd(flat)
		flatmadstd = mad_std(flat, ignore_nan=True)

		if pt: print('Median flat median =', flatmedian )
		if pt: print('Mean flat median =', flatmean )
		if pt: print('Median flat std =', flatstd)
		if pt: print('Median flat mad_std =', flatmadstd)
		if pt: print('Shape of flatimage:', flatimage.shape)
		if pt: print('Shape of flat:', flat.shape)
		if pt: print('')

		if pt: print('List of medians of dark-subtracted flat images')
		for i in range(numflats):
			if pt: print('{:<3}{:<10.2f}'.format(i, flatmediansDS[i]))
		if pt: print('')


		medmin, medmax = [0], [0]
		medmin = np.nanmin(flatmediansDS)
		medmax = np.nanmax(flatmediansDS)
		if pt: print('minimum medians =', medmin)
		if pt: print('maximum medians =', medmax)
		if pt: print('')


		mflat = np.zeros_like(flat)
		mflat = flat / flatmedian
		mflatmedian = np.nanmedian(mflat)
		mflatmadstd = mad_std(mflat, ignore_nane=True)

		if pt: print('mflatmedian, mflatmadstd =', mflatmedian, mflatmadstd)
		if pt: print('')


		vmx = mflatmadstd * 2.0
		vmn = - mflatmadstd * 2.0
		difimage = np.zeros_like(flatimageDS)
		difmadstd = np.zeros((numflats))
		difstd = np.zeros((numflats))
		difmean = np.zeros((numflats))
		difmedian = np.zeros((numflats))
		for i in range(flatimageDS.shape[0]):
			difimage[i] = flatimageDSN[i]-mflat
			difmadstd[i] = mad_std(difimage[i], ignore_nan=True)
			difstd[i] = np.nanstd(difimage[i])
			difmean[i] = np.nanmean(difimage[i])
			difmedian[i] = np.nanmedian(difimage[i])
		dstdmean, dstdmax, dstdmin = np.nanmean(difstd), np.nanmax(difstd), np.nanmin(difstd)
		dstdpcnt = (dstdmax - dstdmin)*100/dstdmean


		self.dataout = DataFits(config=self.config)
		self.dataout.header = flatbaseheader.copy()
		self.dataout.image = mflat


		infolder = os.path.split(flatfiles.filename)
		lfname = os.path.split((len(flatfiles)-1).filename)
		ffname = os.path.split(flatfiles[0].filename)
		if pt: print('lfname =', lfname)
		if pt: print('ffname =', ffname)
		lf = lfname.split('_')
		ff = ffname.split('_')
		flatname = 'mflat_'+ff[1]+'_'+ff[3]+'DR'+'_'+ff[4]+'_'+ff[5]+'-'+lf[5]+'_'+ff[6]+'_'+ff[7]+'_'+'XXX.fits'
		if pt: print('flatname =', flatname)
		if pt: print('')



		gainpcntlim = self.getarg('gainpcntlim')
		dstdpcntlim = self.getarg('dstdpcntlim')
		numfilelim = self.getarg('numfilelim')
		quality = gainpcnt < gainpcntlim and dstdpcnt < dstdpcntlim and numflats >= numfilelim
		outputfolder = self.getarg('outputfolder')
		outputfolder2 = self.getarg('outputfolder2')
		if (outputfolder != '') and (quality == True):
			outputfolder = os.path.expandvars(outputfolder)
			self.dataout.filename = os.path.join(outputfolder, flatname)
		elif (outputfolder != '') and (quality == False):
			outputfolder = os.path.expandvars(outputfolder2)
			self.dataout.filename = os.path.join(outputfolder2, flatname)
		else:
			self.dataout.filename = os.path.join(infolder, flatname)
		if pt: print('outputfolder = ', outputfolder)
		if pt: print('outputfolder2 = ', outputfolder2)
		if pt: print('dstdpcntlim = ', dstdpcntlim)
		if pt: print('gainpcntlim = ', gainpcntlim)
		if pt: print('numfilelim = ', numfilelim)


		imedianH, imadH, imeanH, istdH = flatstats['median'], flatstats['mad'], flatstats['mean'], flatstats['std']


		utime = np.asarray(utimeH)
		etime = utime - utime[0]

		ambient = np.zeros((numflats))
		primary = np.zeros((numflats))
		secondar = np.zeros((numflats))
		dewtem1 = np.zeros((numflats))
		for i in range(numflats):
			ambient[i] = flatheadlist[i].header['ambient']
			primary[i] = flatheadlist[i].header['primary']
			secondar[i] = flatheadlist[i].header['secondar']
			dewtem1[i] = flatheadlist[i].header['dewtem1']
	
		index = np.arange(numflats)


		IDs = []
		for i in range(numflats):
			fi = flatfiles[i].split('_')
			IDS.append(fi[4]+'_'+fi[5])
		fileIDs = np.asarray(IDs)


		tcols = []
		tcols.append(fits.Column(name='index', format='I', array=index))
		tcols.append(fits.Column(name='fileID', format='20A', array=fileIDs))
		tcols.append(fits.Column(name='median', format='D', array=imedianH, unit='ADU'))
		tcols.append(fits.Column(name='mean', format='D', array=imeanH, unit='ADU'))
		tcols.append(fits.Column(name='std', format='D', array=istdH, unit='ADU'))
		tcols.append(fits.Column(name='mad', format='D', array=imadH, unit='ADU'))
		tcols.append(fits.Column(name='dmedian', format='D', array=difmedian, unit='ADU'))
		tcols.append(fits.Column(name='dmean', format='D', array=difmean, unit='ADU'))
		tcols.append(fits.Column(name='dstd', format='D', array=difstd, unit='ADU'))
		tcols.append(fits.Column(name='dmad', format='D', array=difmadstd, unit='ADU'))
		tcols.append(fits.Column(name='ambient', format='D', array=ambient, unit='C'))
		tcols.append(fits.Column(name='primary', format='D', array=primary, unit='C'))
		tcols.append(fits.Column(name='secondar', format='D', array=secondar, unit='C'))
		tcols.append(fits.Column(name='dewtem1', format='D', array=dewtem1, unit='C'))
		tcols.append(fits.Column(name='elapsed time', format='D', array=etime, unit='seconds'))


		c = fits.ColDefs(tcols)
		table = fits.BinTableHDU.from_columns(c)
		tabhead = table.header
		self.dataout.tableset(table.data, tablename = 'table', tableheader=tabhead)


		self.dataout.header['notes'] = '1st HDU: 2D mflat image'
		self.dataout.header['notes3'] = 'Table HDU: statistical and environmental data'
		self.dataout.header['imagetyp'] = 'MFLAT'
		self.dataout.header['bzero'] = 0.0
		self.dataout.header['ambient'] = np.nanmean(ambient)
		self.dataout.header['primary'] = np.nanmean(primary)
		self.dataout.header['secondar'] = np.nanmean(secondar)
		self.dataout.header['dewtem1'] = np.nanmean(dewtem1)        
        
		self.dataout.setheadval('xtimemin', xtimemin, 'Minimum exposure in the set of flats')
		self.dataout.setheadval('xtimemax', xtimemax, 'Maximum exposure in the set of flats')
		self.dataout.setheadval('medminH', medmin, 'Minimum high gain median in the set of flats')
		self.dataout.setheadval('medmaxH', medmax, 'Maximum high gain median in the set of flats')
		self.dataout.setheadval('dstdmean', dstdmean, 'Mean std in the set of difference images')
		self.dataout.setheadval('dstdpcnt', dstdpcnt, '(dstdmax-dstdmin)*100.0/dstdmean')
		self.dataout.setheadval('gainmean', grat_mean, 'Mean gain ratio')
		self.dataout.setheadval('gainpcnt', gainpcnt, '(gainmax-gainmin)*100.0/gainmean')
		self.dataout.setheadval('numfiles', numflats, 'Number of RAW exposures in input datasets')
		self.dataout.setheadval('hotpxlim', hotpxlim, 'Upper limit percentile for unmasked dark current')

		
		self.dataout.setheadval('HISTORY','HDR Master Flat: made from %d x 2 files' % numflats)
		if pt: print(self.dataout.header)

if __name__ == '__main__':
	StepMasterFlatCCD().execute()




































				


































