#!/bin/python
'''
Copied from smcrowe1
To be modified for GEOS
Modified by klmorga5
'''
from netCDF4 import Dataset
#import matplotlib.pyplot as plt
import numpy as np
import pdb,sys,glob,calendar,os
import datetime as dt
import earth_calcs as ec
from gisscoxmunk import giss_cox_munk as gcm

tau_threshold = 0.5
pol_angle = -30*np.pi/180.
alb_lb = 0.05
alb_ub = 0.6
snr_lb = 200

osse_dir = '/discover/nobackup/projects/gmao/geos_carb/lott/OSSE_output/'
oco_nr_dir = './oco/'
oco_nr_prefix = 'oco_nature_run_sampling_'
oco_output_dir = './oco/'

oco_uncert_params = []
oco_uncert_params.append(np.loadtxt(oco_output_dir+'ocean_wco2_params.txt'))   # need 
oco_uncert_params.append(np.loadtxt(oco_output_dir+'land_wco2_params.txt'))    # need
oco_uncert_params = np.array(oco_uncert_params)

def read_file(date=[2012,1,2]):
    write_vars = ['latitude','longitude','GCHEM_TM5_diff','sat_aa','saa','sat_za','sza','band6_albedo','obs_mode']
    yr = date[0]
    mon = date[1]
    day = date[2]
    isccpfname = oco_nr_dir+oco_nr_prefix+('%4d%02d%02d%s')%(yr,mon,day,'.nc4')
    oco_vars = {}
    if os.path.exists(isccpfname):
        isccp_fid_nr = Dataset(isccpfname)
        write_vars.extend(isccp_fid_nr.groups['cld_isccp'].variables.keys())
        write_vars.remove('time')

        write_vars.extend(isccp_fid_nr.groups['bioco2'].variables.keys())
        write_vars.extend(isccp_fid_nr.groups['urbco2'].variables.keys())
        write_vars.remove('time')
        write_vars.remove('CO2_CL')
        write_vars.remove('CO2_CL') 

        write_vars.append('bio_CO2_CL')
        write_vars.append('urb_CO2_CL')
        #write_vars.append('cfrac')

        v = 'bio_CO2_CL'
        vr = isccp_fid_nr['bioco2/CO2_CL']
        oco_vars[v] = {}
        oco_vars[v]['values'] = vr[:]
        oco_vars[v]['attributes'] = {'long_name':vr.long_name, 'units':vr.units}
        v = 'urb_CO2_CL'
        vr = isccp_fid_nr['urbco2/CO2_CL']
        oco_vars[v] = {}
        oco_vars[v]['values'] = vr[:]
        oco_vars[v]['attributes'] = {'long_name':vr.long_name, 'units':vr.units}
        v = 'cfrac'
        vr = isccp_fid_nr['cld_isccp/cfrac']
        oco_vars[v] = {}
        oco_vars[v]['values'] = vr[:]
        oco_vars[v]['attributes'] = {'long_name':vr.long_name, 'units':vr.units}
        for v in isccp_fid_nr.variables.keys():
            oco_vars[v] = {}
            oco_vars[v]['values'] = isccp_fid_nr[v][:]
        for g in isccp_fid_nr.groups.keys():
            for v in isccp_fid_nr[g].variables.keys():
                if v == 'time':continue
                vr = isccp_fid_nr[g][v]
                oco_vars[v] = {}
                oco_vars[v]['values'] = vr[:]
                if v == 'dPS':continue
                oco_vars[v]['attributes'] = {'long_name':vr.long_name, 'units':vr.units}
        oco_vars['merra_albedo'] = {}
        oco_vars['merra_albedo']['values'] = (oco_vars['NIRDR']['values']*oco_vars['ALBNIRDR']['values'] + oco_vars['NIRDF']['values']*oco_vars['ALBNIRDF']['values']) / (oco_vars['NIRDR']['values'] + oco_vars['NIRDF']['values'])
        oco_vars['FRSNO']['values'][np.isnan(oco_vars['FRSNO']['values'])] = 0
        oco_vars['FRSEAICE']['values'][np.isnan(oco_vars['FRSEAICE']['values'])] = 0
        oco_vars['ASNOW_GL']['values'][np.isnan(oco_vars['ASNOW_GL']['values'])] = 0
        write_vars.append('merra_albedo')
    else:
        print('No file found for date')
    oco_vars['date'] = date[:]
    return oco_vars,write_vars

def cloud_screen(nr_dict={},write_data={}):
    oco_good = {}
    bc_aod = oco_nr_vars['BCEXTTAU']['values']*(755./550.)**(-1*oco_nr_vars['BCANGSTR']['values'])
    du_aod = oco_nr_vars['DUEXTTAU']['values']*(755./550.)**(-1*oco_nr_vars['DUANGSTR']['values'])
    oc_aod = oco_nr_vars['OCEXTTAU']['values']*(755./550.)**(-1*oco_nr_vars['OCANGSTR']['values'])
    ss_aod = oco_nr_vars['SSEXTTAU']['values']*(755./550.)**(-1*oco_nr_vars['SSANGSTR']['values'])
    su_aod = oco_nr_vars['SUEXTTAU']['values']*(755./550.)**(-1*oco_nr_vars['SUANGSTR']['values'])
    tot_aod = bc_aod + du_aod + oc_aod + ss_aod + su_aod

    cld_isccp = oco_nr_vars['cfrac']['values'][:]
    cldlow = oco_nr_vars['CLDLOW']['values'][:]
    cldmid = oco_nr_vars['CLDMID']['values'][:]
    cldhgh = oco_nr_vars['CLDHGH']['values'][:]
    cldtot = oco_nr_vars['CLDTOT']['values'][:]
    taulow = oco_nr_vars['TAULOW']['values'][:]
    taumid = oco_nr_vars['TAUMID']['values'][:]
    tauhgh = oco_nr_vars['TAUHGH']['values'][:]
    tautot = oco_nr_vars['TAUTOT']['values'][:]

    prob_bin = np.zeros((tautot.shape[0],8))
    prob_isccp = np.zeros((tautot.shape[0],1))
    tau_bin = prob_bin.copy()

    #Clear sky
    prob_bin[:,0] = (1-cldlow)*(1-cldmid)*(1-cldhgh)
    
    #Two cloudy layers
    prob_bin[:,1] = (1-cldlow)*(cldmid)*(cldhgh)
    prob_bin[:,2] = (cldlow)*(1-cldmid)*(cldhgh)
    prob_bin[:,3] = (cldlow)*(cldmid)*(1-cldhgh)
    tau_bin[:,1] = taumid + tauhgh
    tau_bin[:,2] = taulow + tauhgh
    tau_bin[:,3] = taulow + taumid
    prob_isccp[:,0] = 1.0 - cld_isccp[:]

    #Three cloudy layers
    prob_bin[:,4] = (cldlow)*(cldmid)*(cldhgh)
    tau_bin[:,4] = taulow+taumid+tauhgh

    #One cloudy layer
    prob_bin[:,5] = cldlow - prob_bin[:,2] - prob_bin[:,3] - prob_bin[:,4]
    prob_bin[:,6] = cldmid - prob_bin[:,1] - prob_bin[:,3] - prob_bin[:,4]
    prob_bin[:,7] = cldhgh - prob_bin[:,1] - prob_bin[:,2] - prob_bin[:,4]
    tau_bin[:,5] = taulow[:]
    tau_bin[:,6] = taumid[:]
    tau_bin[:,7] = tauhgh[:]
    
    #airmass
    sza = oco_nr_vars['sza']['values'][:]
    za = oco_nr_vars['sat_za']['values'][:]
    airmass = 1.0/np.cos(np.pi/180.*sza) + 1.0/np.cos(np.pi/180.*za)

#    total_tau = tau_bin + tot_aod[:,None]
    total_tau = tot_aod[:,None]
    #good_screen = (airmass[:]*((prob_bin > 0)*tau_bin + tot_aod[:,None]).sum(1) < tau_threshold)
#    good_screen = (airmass[:]*(tau_bin[:,0] + tot_aod[:]) < tau_threshold)*(prob_bin[:,0] > 0.1) 
    good_screen = (airmass[:]*(tot_aod[:])<tau_threshold)*(prob_isccp[:,0]>0.5)
    #total_prob = (prob_bin*good_screen[:,None]).sum(1)
#    total_prob = prob_bin[:,0]*good_screen[:] 
    total_prob = prob_isccp[:,0]*good_screen[:]
    #total_footprints = (prob_bin*good_screen[:,None]*oco_nr_vars['n_footprints']['values'][:,None]).sum(1)
    total_footprints = prob_isccp[:,0]*good_screen[:]*oco_nr_vars['n_footprints']['values'][:]    

    oco_nr_vars['prob_bin'] = {}
    oco_nr_vars['total_od_bin'] = {}
    oco_nr_vars['isccp_screened_footprints'] = {}
    oco_nr_vars['od_screen'] = {}
    oco_nr_vars['airmass'] = {}
    oco_nr_vars['n_screened_footprints'] = {}

    oco_nr_vars['prob_bin']['values'] = prob_isccp
    oco_nr_vars['total_od_bin']['values'] = total_tau
    oco_nr_vars['n_screened_footprints']['values'] = total_footprints[:]
    oco_nr_vars['isccp_screened_footprints']['values'] = (prob_isccp*good_screen[:,None]*oco_nr_vars['n_footprints']['values'][:,None]).astype(int)
    oco_nr_vars['od_screen']['values'] = good_screen
#    oco_nr_vars['airmass']['values'] = airmass
    write_vars.append('n_screened_footprints')
    return oco_nr_vars,write_vars

def uncert_param(snr=0.,p=[]):
    return p[0]/(1.+p[1]*snr**p[2])

def uncert_param_inverse(uncert=0,p=[]):
    return ((p[0]-uncert)/(p[1]*uncert))**(1./p[2])

def calc_co2_uncert(oco_nr_vars,write_vars):
    fsun = 2073 #nW cm sr^(-1)  cm^(-2)
    Cphot = 6.69656e-3
    Cback = 6.02299e-3
    Cratio = (Cback/Cphot)**2
    #original mms in photons/sec/m^2/micron/sr
    #conversion: photons per sec -> Joules per sec, Energy = h*c/lambda
    # area units: 1e4 cm^2 / m^2
    # 1e6 microns per m -> 1e4 microns per cm
    # Convert W to nW: 1e3
    conv_fac = 1e9*6.62606957e-34*3e8/1.6e-6*1e-4/(1./1.6)**2*1e-4
    mms_oco_units = 2.45e20
    mms = mms_oco_units * conv_fac
    mms100 = mms/100.
    N0 = mms100*Cback
    N1 = mms100*Cphot**2
    co2_uncorr_uncert = np.nan*np.zeros(len(oco_nr_vars['latitude']['values']))
    co2_corr_uncert = co2_uncorr_uncert.copy()
    effective_snr = co2_uncorr_uncert.copy()
    good_inds = np.where(oco_nr_vars['od_screen']['values'])[0]
    oco_nr_vars['airmass']['values'] = 1./np.cos(np.pi/180.*oco_nr_vars['sza']['values']) + 1./np.cos(np.pi/180*oco_nr_vars['sat_za']['values'])
    wspd = np.sqrt(oco_nr_vars['U10M']['values'][:]**2 + oco_nr_vars['V10M']['values'][:]**2)
    cox_alb = gcm(wspd,1.31,oco_nr_vars['sza']['values'][:],oco_nr_vars['sat_za']['values'][:],oco_nr_vars['saa']['values'][:]-oco_nr_vars['sat_aa']['values'][:]-180.)[:,0,:]
    mod_alb = oco_nr_vars['band6_albedo']['values'][:]
    oco_nr_vars['albedo'] = {}
    oco_nr_vars['albedo']['values'] = np.nan*np.zeros(len(oco_nr_vars['latitude']['values']))
    for i_ob in good_inds:
        if np.isnan(mod_alb[i_ob]): #invalid land surface - use the ocean parameterization
            p = oco_uncert_params[0]
            alb = cox_alb[i_ob,0] + cox_alb[i_ob,1]*np.cos(2*pol_angle)+cox_alb[i_ob,2]*np.sin(2*pol_angle)
        else:
            p = oco_uncert_params[1]
            alb = mod_alb[i_ob] 
        oco_nr_vars['albedo']['values'][i_ob] = alb 
        S = 0.5*fsun*alb*np.cos(oco_nr_vars['sza']['values'][i_ob]*np.pi/180.)*np.exp(-oco_nr_vars['total_od_bin']['values'][i_ob,:]*oco_nr_vars['airmass']['values'][i_ob])
        #Noise parameterization from Chris O'Dell
        N = np.sqrt(N0**2+N1*S)#mms100 * np.sqrt( abs(S)/mms100*Cphot**2 + Cback**2 )
        corr_uncert = 0
        uncorr_uncert = 0
        good_ob = False
        for i_bin in range(len(S)):
            snr = np.min((S[i_bin]/N[i_bin],10000))
            
            if snr <= snr_lb: continue
            if oco_nr_vars['isccp_screened_footprints']['values'][i_ob,i_bin] <= 0: continue

            good_ob = True
            uncert = uncert_param(snr,p)
            corr_uncert = np.nanmax((uncert,corr_uncert))
            uncorr_uncert +=  uncert**2/oco_nr_vars['isccp_screened_footprints']['values'][i_ob,i_bin]
        # screen out low albedos
        if alb < alb_lb: good_ob = False
        if alb > alb_ub: good_ob = False
        # screen out snow
        if oco_nr_vars['FRSNO']['values'][i_ob] > 0.05: good_ob = False
        # screen out sea ice
        if oco_nr_vars['FRSEAICE']['values'][i_ob] > 0.05: good_ob = False
        # screen out glacial ice
        if oco_nr_vars['ASNOW_GL']['values'][i_ob] > 0.05: good_ob = False
        # screen out ocean nadir observations
 #       if oco_nr_vars['obs_mode']['values'][i_ob] > 0.5 and np.isnan(mod_alb[i_ob]): good_ob = False

        if good_ob: 
            co2_uncorr_uncert[i_ob] = np.sqrt(uncorr_uncert)
            co2_corr_uncert[i_ob] = corr_uncert
            effective_snr[i_ob] = uncert_param_inverse(corr_uncert,p)
            #if corr_uncert < 0.1: pdb.set_trace()
    oco_nr_vars['co2_uncorrelated_random_error'] = {}
    oco_nr_vars['co2_correlated_random_error'] = {}
    oco_nr_vars['effective_snr'] = {}
    oco_nr_vars['od_snr_screen'] = {}

    oco_nr_vars['effective_snr']['values'] = effective_snr[:]
    oco_nr_vars['co2_uncorrelated_random_error']['values'] = co2_uncorr_uncert
    oco_nr_vars['co2_correlated_random_error']['values'] = co2_corr_uncert
    oco_nr_vars['od_snr_screen']['values'] = 1-np.isnan(co2_corr_uncert)
    write_vars.extend(['co2_uncorrelated_random_error','effective_snr','co2_correlated_random_error'])
    return oco_nr_vars,write_vars

def calculate_somkuti_bias(gc_nr_vars,write_vars):

    oco_nr_vars['bias_tccon'] = {}
    land_bias = (-0.24068966126886712 +
                np.cos(np.deg2rad(oco_nr_vars['latitude']['values'])) * 1.3927042908594287 +
                oco_nr_vars['airmass']['values'] * -0.5090796401130635 +
                oco_nr_vars['SSEXTTAU']['values'] * 4.690745085341598 +
                oco_nr_vars['DUEXTTAU']['values'] * 2.7965163912126467 +
                oco_nr_vars['BCEXTTAU']['values'] * -3.6422052806676786 +
                oco_nr_vars['OCEXTTAU']['values'] * 2.2936700195759476 +
                oco_nr_vars['SUEXTTAU']['values'] * 3.2259143583137497 +
                oco_nr_vars['erai_alt']['values'] / 9.81 * -0.0003210222971355435 +
                oco_nr_vars['refl_mean_NIR']['values'] * 1e-4 * -1.3685630955512533 +
                oco_nr_vars['refl_mean_SWIR1']['values'] * 1e-4 * -1.802772803741406 +
                oco_nr_vars['refl_mean_SWIR2']['values'] * 1e-4 * 6.4040986156361654 +
                oco_nr_vars['refl_mean_SWIR3']['values'] * 1e-4 * -4.089139029908539)
    ocean_bias = (
                0.7316293168676731 +
                np.cos(np.deg2rad(oco_nr_vars['latitude']['values'])) * 1.1701116517231982 +
                oco_nr_vars['airmass']['values'] * -0.7587014045893425 +
                oco_nr_vars['SSEXTTAU']['values'] * 0.6474496013696902 +
                oco_nr_vars['DUEXTTAU']['values'] * -2.7989584176701277 +
                oco_nr_vars['BCEXTTAU']['values'] * -74.02110924140642 +
                oco_nr_vars['OCEXTTAU']['values'] * 4.524531058809314 +
                oco_nr_vars['SUEXTTAU']['values'] * 2.4807031482717443
                )
    oco_nr_vars['bias_tccon']['values'] = np.zeros(land_bias.shape)
    inds = np.where(np.isnan(oco_nr_vars['band6_albedo']['values']))[0]
    oco_nr_vars['bias_tccon']['values'][inds] = ocean_bias[inds].copy()
    inds = np.where(~np.isnan(oco_nr_vars['band6_albedo']['values']))[0]
    oco_nr_vars['bias_tccon']['values'][inds] = land_bias[inds].copy()

    oco_nr_vars['bias_models'] = {}
    ocean_bias = (
                -18.791115953392456 +
                oco_nr_vars['latitude']['values'] * -0.028878331191302477 +
                oco_nr_vars['airmass']['values'] * 0.109069121497231 +
                (oco_nr_vars['latitude']['values'] * oco_nr_vars['airmass']['values']) * 0.0076211300239649934 +
                np.cos(np.deg2rad(oco_nr_vars['latitude']['values'])) * -0.554626445951762 +
                np.exp(oco_nr_vars['erai_lnsp']['values']) * 0.00018358886399796607 +
                oco_nr_vars['SSEXTTAU']['values'] * -1.1720841547234304 +
                oco_nr_vars['DUEXTTAU']['values'] * 3.643499446751121 +
                oco_nr_vars['BCEXTTAU']['values'] * 28.18064639469636 +
                oco_nr_vars['OCEXTTAU']['values'] * 2.664854428557407 +
                oco_nr_vars['SUEXTTAU']['values'] * -4.376619932223321
                )
    land_bias = (
                5.129460780008606 +
                np.exp(oco_nr_vars['erai_lnsp']['values']) * -4.241481860668963e-05 +
                oco_nr_vars['latitude']['values'] * -0.016490869739507984 +
                oco_nr_vars['airmass']['values'] * -0.40931429057161106 +
                (oco_nr_vars['latitude']['values'] * oco_nr_vars['airmass']['values']) * 0.00712806124413639 +
                (np.cos(np.deg2rad(oco_nr_vars['latitude']['values']))) * 0.2571824292565922 +
                oco_nr_vars['DUEXTTAU']['values'] * 3.090414441405752 +
                oco_nr_vars['BCEXTTAU']['values'] * 10.171439235544444 +
                oco_nr_vars['OCEXTTAU']['values'] * 1.1138280975144368 +
                oco_nr_vars['SUEXTTAU']['values'] * -4.265346292156446 +
                oco_nr_vars['SSEXTTAU']['values'] * 5.17071894079167 +
                oco_nr_vars['erai_alt']['values'] / 9.81 * -0.00037355592731388787  +
                oco_nr_vars['refl_mean_NIR']['values'] * 1e-4 * -0.1973228967627258 +
                oco_nr_vars['refl_mean_SWIR3']['values'] * 1e-4 * -1.911111234840418
                )

    oco_nr_vars['bias_models']['values'] = np.zeros(land_bias.shape)
    inds = np.where(np.isnan(oco_nr_vars['band6_albedo']['values']))[0]
    oco_nr_vars['bias_models']['values'][inds] = ocean_bias[inds].copy()
    inds = np.where(~np.isnan(oco_nr_vars['band6_albedo']['values']))[0]
    oco_nr_vars['bias_models']['values'][inds] = land_bias[inds].copy()
    write_vars.append('bias_models')
    write_vars.append('bias_tccon')
    return oco_nr_vars,write_vars

def write_file(oco_nr_vars,write_vars):
    date = oco_nr_vars['date'][:]
    scrn = np.where(oco_nr_vars['od_snr_screen']['values'])[0]
    if len(scrn) == 0: return
    fname = oco_output_dir+'oco_cloud_screened_samples_'+('%4d%02d%02d')%(date[0],date[1],date[2])+'__ISCCP_bias.nc4'
#    fname = oco_output_dir+'NOOCEAN_oco_cloud_screened_samples_'+('%4d%02d%02d')%(date[0],date[1],date[2])+'__ISCCP.nc4'
    fid_out = Dataset(fname,'w')
    fid_out.createDimension('n_obs',0)
    fid_out.createDimension('idate',6)
    for var_name in write_vars:
        if var_name == 'time':
            v = fid_out.createVariable(var_name,'i2',('n_obs','idate'))
        else:
            v = fid_out.createVariable(var_name,oco_nr_vars[var_name]['values'].dtype,'n_obs')
        v[:] = oco_nr_vars[var_name]['values'][scrn]
        if 'attributes' in oco_nr_vars[var_name].keys():
            v.long_name = oco_nr_vars[var_name]['attributes']['long_name']
            v.units = oco_nr_vars[var_name]['attributes']['units']
    fid_out.close()

d = [int(t) for t in sys.argv[1].split(',')]
sdate = dt.datetime(d[0],d[1],d[2])
d = [int(t) for t in sys.argv[2].split(',')]
edate = dt.datetime(d[0],d[1],d[2])
cdate = sdate
while cdate < edate:
    print(cdate.strftime('%Y-%m-%d'))
    oco_nr_vars,write_vars = read_file(date=cdate.timetuple()[:3])
    if 'latitude' not in oco_nr_vars.keys():
        cdate += dt.timedelta(days=1)
        continue
    oco_nr_vars,write_vars = cloud_screen(oco_nr_vars,write_vars)
    oco_nr_vars,write_vars = calc_co2_uncert(oco_nr_vars,write_vars)
    oco_nr_vars,write_vars = calculate_somkuti_bias(oco_nr_vars,write_vars)
    write_file(oco_nr_vars,write_vars)
    cdate += dt.timedelta(days=1)
