#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Created on Thu Jul 23 11:51:33 2020

@author: jabadgeley
"""

# Import packages
import numpy as np
import xarray as xr
import os
import glob
#import sys
#sys.path.append('<path_to_scripts>')
from model_data import model_data
from constants import (latent_heat_vaporization as Lv,
                       adiabatic_lapse_rate_dry as gamma_d,
                       specific_heat_dry_air as cp,
                       gas_constant_dry_air as Rd,
                       epsilon)



def gather_variables(model_group, model_name, experiment, variables,
                     subexperiment, frequency):
    
    if (model_group == 'PMIP3') or (model_group == 'PMIP4'):
        
        name = f'{model_group} {model_name} {experiment} {subexperiment} {frequency}'
        
        path = f'{model_group}'
        
        full_path = os.path.join(path, model_name, experiment, subexperiment,
                                 frequency) 
        vars_2D_options = ['orog']
        vars_3D_options = ['tas','pr','ps','hurs','huss','hfls']
        vars_4D_options = ['hus','hur','zg']
        
        vars_2D_list = list(set(variables) & set(vars_2D_options))
        vars_3D_list = list(set(variables) & set(vars_3D_options))
        vars_4D_list = list(set(variables) & set(vars_4D_options))
        
        vars_0D = {}
        vars_1D = {}
        vars_2D = {}
        vars_3D = {}
        vars_4D = {}
        
        for v2 in vars_2D_list:
            try:
                nc_path = glob.glob(full_path+'/'+ v2+'_*.nc')[0]
                vars_2D[v2] = xr.open_dataset(nc_path)[v2]
                
            except:
                print(f'Variable "{v2}" does not have an associated file.')
                
        for v3 in vars_3D_list:
            try:
                nc_path = glob.glob(full_path+'/'+ v3+'_*.nc')[0]
                vars_3D[v3] = xr.open_dataset(nc_path)[v3]
                
            except:
                print(f'Variable "{v3}" does not have an associated file.')
                
        for v4 in vars_4D_list:
            try:
                nc_path = glob.glob(full_path+'/'+ v4+'_*.nc')[0]
                vars_4D[v4] = xr.open_dataset(nc_path)[v4]
                
            except:
                print(f'Variable "{v4}" does not have an associated file.')

        MD_obj = model_data(name, vars_0D, vars_1D, vars_2D, vars_3D, vars_4D)
        
    elif model_group == 'New_Model_Output':
        name = f'{model_group} {experiment} {frequency}'
        
        ## Find the file with the correct frequency... nc_path = ...
        nc_path = f'{experiment}.cam.h0.{frequency}-climatology.nc'
        
        nc = xr.open_dataset(nc_path)
        
        nc = nc.rename({'lev':'plev'})
        
        vars_0D_options = ['P0']
        vars_1D_options = ['hyam','hybm']
        vars_3D_options = ['PS']
        vars_4D_options = ['RELHUM','T','U','V','Z3','Q']
    
        vars_0D_list = list(set(variables) & set(vars_0D_options))
        vars_1D_list = list(set(variables) & set(vars_1D_options))
        vars_3D_list = list(set(variables) & set(vars_3D_options))
        vars_4D_list = list(set(variables) & set(vars_4D_options))
        
        vars_0D = {}
        vars_1D = {}
        vars_2D = {}
        vars_3D = {}
        vars_4D = {}
        
        for v0 in vars_0D_list:
            try:
                vars_0D[v0] = nc[v0]
                
            except:
                print(f'Variable "{v0}" is not in the file.')
        
        for v1 in vars_1D_list:
            try:
                vars_1D[v1] = nc[v1]
                
            except:
                print(f'Variable "{v1}" is not in the file.')
                
        for v3 in vars_3D_list:
            try:
                vars_3D[v3] = nc[v3]
                
            except:
                print(f'Variable "{v3}" is not in the file.')
                
        for v4 in vars_4D_list:
            try:
                vars_4D[v4] = nc[v4]
                
            except:
                print(f'Variable "{v4}" is not in the file.')
        
        MD_obj = model_data(name, vars_0D, vars_1D, vars_2D, vars_3D, vars_4D)
        
    else:
        return print(model_group + ' is not currently an option.')
    
    return MD_obj



def Prep_MDobj(model_group, model_name, experiment, subexperiment, 
               frequency, variables, regrid_data=False, select_lats=True,
               lons_to_360=False, time_average=True):
    
    MDobj = gather_variables(model_group = model_group, 
                             model_name = model_name, 
                             experiment = experiment, 
                             variables = variables, 
                             subexperiment = subexperiment, 
                             frequency = frequency)
        
    if time_average:
        MDobj.average_time(average_type='monthly_weights')
    
    if regrid_data:
        MDobj.lons_to_360()
        MDobj.sortby_lons()
        MDobj.regrid_to_targetlist(regrid_data['lat'],regrid_data['lon'])
    
    if select_lats:
        MDobj.select_lats((-95,-60))
        
    if lons_to_360:
        MDobj.lons_to_360()
        
    
    return MDobj



def Get_MDobjs(model_group, model_name, experiment_recent,
               experiment_past, subexperiment_recent, 
               subexperiment_past, frequency, regrid_data=False, 
               select_lats=True, lons_to_360=True, variables='default'):
    
    if type(variables) is list:
        input_variables = variables
    else:
        if (model_group == 'PMIP3') | (model_group == 'PMIP4'):
            input_variables = ['tas','orog']
        elif model_group == 'New_Model_Output':
            input_variables = ['T','Z3']
    
    recent = Prep_MDobj(model_group, model_name, experiment_recent,
                        subexperiment_recent, frequency, input_variables,
                        regrid_data, select_lats, lons_to_360)
    past = Prep_MDobj(model_group, model_name, experiment_past,
                      subexperiment_past, frequency, input_variables,
                      regrid_data, select_lats, lons_to_360)
    
    return recent, past



def PMIP_dry_decomposition_refhgt(model_group, model_name, experiment_recent,
                                  experiment_past, subexperiment_recent, 
                                  subexperiment_past, frequency, 
                                  regrid_data=False, select_lats=True,
                                  lons_to_360=False):
    
    variables = ['tas','orog','ps']
    
    recent = Prep_MDobj(model_group, model_name, experiment_recent,
                        subexperiment_recent, frequency, variables, regrid_data, 
                        select_lats,lons_to_360)
    past = Prep_MDobj(model_group, model_name, experiment_past,
                      subexperiment_past, frequency, variables, regrid_data, 
                      select_lats,lons_to_360)
    
    dT_full = recent.tas.tas - past.tas.tas
    
    theta_recent = recent.tas.tas * (100000 / recent.ps.ps) ** (Rd / cp)
    theta_past = past.tas.tas * (100000 / past.ps.ps) ** (Rd / cp)
    dtheta = theta_recent - theta_past
    
    exner = past.tas.tas / theta_past
        
    dZ = recent.orog.orog - past.orog.orog
    
    lr = gamma_d
    
    dT_adiabatic = -1 * lr * dZ
    
    dT_diabatic = exner * dtheta
    
    #dT_residual = dT_full - (dT_adiabatic + dT_diabatic)
    
    #dT_nonadiabatic = dT_full - dT_adiabatic
    
    return dT_full, dT_diabatic, dT_adiabatic



def PMIP_moist_decomposition_refhgt(model_group, model_name, experiment_recent,
                                    experiment_past, subexperiment_recent, 
                                    subexperiment_past, frequency, 
                                    hurs = 'hurs', huss = 'huss',
                                    regrid_data=False, select_lats=True,
                                    lons_to_360=False):
    
    variables = ['tas','orog','ps',hurs,huss]
    
    recent = Prep_MDobj(model_group, model_name, experiment_recent,
                        subexperiment_recent, frequency, variables, regrid_data, 
                        select_lats,lons_to_360)
    past = Prep_MDobj(model_group, model_name, experiment_past,
                      subexperiment_past, frequency, variables, regrid_data, 
                      select_lats,lons_to_360)
    
    dT_full = recent.tas.tas - past.tas.tas
    
    theta_recent = recent.tas.tas * (100000 / recent.ps.ps) ** (Rd / cp)
    theta_past = past.tas.tas * (100000 / past.ps.ps) ** (Rd / cp)
    dtheta = theta_recent - theta_past
    
    exner = past.tas.tas / theta_past
        
    dZ = recent.orog.orog - past.orog.orog
    
    try:
        huss = past.huss.huss
        hurs = past.hurs.hurs

    except:
        #print('No huss and hurs variables available')
        huss = past.hus.hus.bfill(dim='plev').isel(plev=0)
        hurs = past.hur.hur.bfill(dim='plev').isel(plev=0)
            
    ws = ((100 * huss) / (hurs * (1 - huss)))
            
    lr = gamma_d * (1 + ((Lv * ws) / (Rd * past.tas.tas))) / (1 + ((epsilon * \
                ws * (Lv ** 2)) / (cp * Rd * (past.tas.tas ** 2))))
    
    #dT_adiabatic = -1 * gamma_s * dz
    
    #dT_diabatic = exner_recent * dtheta
    
    dT_diabatic = (dtheta * exner * lr / (gamma_d * (1 + ((Lv * ws) / \
                  (Rd * past.tas.tas)))))
    
    dT_adiabatic = -1 * lr * dZ
    
    #dT_residual = dT_full - (dT_adiabatic + dT_diabatic)
    
    #dT_nonadiabatic = dT_full - dT_adiabatic
    
    return dT_full, dT_diabatic, dT_adiabatic, lr
    


def MD_dry_decomposition_lowhgt(model_group, experiment_recent,
                                experiment_past, 
                                frequency, regrid_data=False, 
                                select_lats=True,lons_to_360=False,
                                time_average=True):
    
    variables = ['T','Z3','hyam','hybm','P0','PS']
    
    recent = Prep_MDobj(model_group, '', experiment_recent,
                        '', frequency, variables, regrid_data, select_lats,
                        lons_to_360, time_average)
    past = Prep_MDobj(model_group, '', experiment_past,
                      '', frequency, variables, regrid_data, select_lats,
                      lons_to_360, time_average)
    
    lats = recent.T.lat.values
    lons = recent.T.lon.values
    levs = recent.T.plev.values
    nlevs = len(levs)
    nlons = len(lons)
    nlats = len(lats)
    
    if time_average:
        p_recent = np.tile(recent.P0.values[np.newaxis,np.newaxis,np.newaxis], (nlevs,nlats,nlons)) * \
               np.tile(recent.hyam.values[...,np.newaxis,np.newaxis], (1,nlats,nlons)) + \
               np.tile(recent.PS.PS.values, (nlevs,1,1)) * \
               np.tile(recent.hybm.values[...,np.newaxis,np.newaxis], (1,nlats,nlons))
        p_past = np.tile(past.P0.values[np.newaxis,np.newaxis,np.newaxis], (nlevs,nlats,nlons)) * \
             np.tile(past.hyam.values[...,np.newaxis,np.newaxis], (1,nlats,nlons)) + \
             np.tile(past.PS.PS.values, (nlevs,1,1)) * \
             np.tile(past.hybm.values[...,np.newaxis,np.newaxis], (1,nlats,nlons))
        
        recent.P = xr.DataArray(p_recent, 
                                coords={'plev':levs,'lat':lats,'lon':lons},
                                dims=['plev','lat','lon'])
        past.P = xr.DataArray(p_past, 
                              coords={'plev':levs,'lat':lats,'lon':lons},
                              dims=['plev','lat','lon'])
    
    elif not time_average:
        times = recent.T.time.values
        ntimes = len(times)
            
        p_recent = np.tile(recent.P0.values[np.newaxis,np.newaxis,np.newaxis,np.newaxis], 
                           (nlevs,ntimes,nlats,nlons)) * \
               np.tile(recent.hyam.values[...,np.newaxis,np.newaxis,np.newaxis], (1,ntimes,nlats,nlons)) + \
               np.tile(recent.PS.PS.values, (nlevs,1,1,1)) * \
               np.tile(recent.hybm.values[...,np.newaxis,np.newaxis,np.newaxis], (1,ntimes,nlats,nlons))
        p_past = np.tile(past.P0.values[np.newaxis,np.newaxis,np.newaxis,np.newaxis], 
                         (nlevs,ntimes,nlats,nlons)) * \
             np.tile(past.hyam.values[...,np.newaxis,np.newaxis,np.newaxis], (1,ntimes,nlats,nlons)) + \
             np.tile(past.PS.PS.values, (nlevs,1,1,1)) * \
             np.tile(past.hybm.values[...,np.newaxis,np.newaxis,np.newaxis], (1,ntimes,nlats,nlons))
        
        recent.P = xr.DataArray(p_recent, 
                                coords={'plev':levs,'time':times,'lat':lats,'lon':lons},
                                dims=['plev','time','lat','lon'])
        past.P = xr.DataArray(p_past, 
                              coords={'plev':levs,'time':times,'lat':lats,'lon':lons},
                              dims=['plev','time','lat','lon'])
        
        past.T = past.T.T.assign_coords({"time":times})
        recent.T = recent.T.T.assign_coords({"time":times})
        past.Z3 = past.Z3.assign_coords({"time":times})
        recent.Z3 = recent.Z3.assign_coords({"time":times})
    
    dT_full = (recent.T.T - past.T.T).isel(plev=nlevs-1)
    
    recent.theta = recent.T.T * (recent.P0 / recent.P) ** (Rd / cp)
    past.theta = past.T.T * (past.P0 / past.P) ** (Rd / cp)
    dtheta = recent.theta - past.theta
    
    exner = past.T.T / past.theta
        
    dz = recent.Z3.Z3 - past.Z3.Z3
    
    dZ = dz.isel(plev=nlevs-1)
    
    lr = gamma_d
    
    #dT_adiabatic = -1 * lr * dZ
    
    dT_diabatic = (exner * dtheta).isel(plev=nlevs-1)
    
    #dT_advection = dT_full - (dT_adiabatic + dT_diabatic)
    
    #dT_nonadiabatic = dT_full - dT_adiabatic
    
    return dT_full, dT_diabatic, dZ, lr



def MD_moist_decomposition_lowhgt(model_group, experiment_recent,
                                  experiment_past,
                                  frequency, regrid_data=False, 
                                  select_lats=True, lons_to_360=False,
                                  time_average=True):
    
    variables = ['T','Z3','hyam','hybm','P0','Q','RELHUM','PS']
    
    recent = Prep_MDobj(model_group, '', experiment_recent,
                        '', frequency, variables, regrid_data, select_lats,
                        lons_to_360, time_average)
    past = Prep_MDobj(model_group, '', experiment_past,
                      '', frequency, variables, regrid_data, select_lats,
                      lons_to_360, time_average)
    
    lats = recent.T.lat.values
    lons = recent.T.lon.values
    levs = recent.T.plev.values
    nlevs = len(levs)
    nlons = len(lons)
    nlats = len(lats)
    
    if time_average:
        p_recent = np.tile(recent.P0.values[np.newaxis,np.newaxis,np.newaxis], (nlevs,nlats,nlons)) * \
               np.tile(recent.hyam.values[...,np.newaxis,np.newaxis], (1,nlats,nlons)) + \
               np.tile(recent.PS.PS.values, (nlevs,1,1)) * \
               np.tile(recent.hybm.values[...,np.newaxis,np.newaxis], (1,nlats,nlons))
        p_past = np.tile(past.P0.values[np.newaxis,np.newaxis,np.newaxis], (nlevs,nlats,nlons)) * \
             np.tile(past.hyam.values[...,np.newaxis,np.newaxis], (1,nlats,nlons)) + \
             np.tile(past.PS.PS.values, (nlevs,1,1)) * \
             np.tile(past.hybm.values[...,np.newaxis,np.newaxis], (1,nlats,nlons))
             
        recent.P = xr.DataArray(p_recent, 
                                coords={'plev':levs,'lat':lats,'lon':lons},
                                dims=['plev','lat','lon'])
        past.P = xr.DataArray(p_past, 
                              coords={'plev':levs,'lat':lats,'lon':lons},
                              dims=['plev','lat','lon'])
        
    elif not time_average:
        times = recent.T.time.values
        ntimes = len(times)
        
        p_recent = np.tile(recent.P0.values[np.newaxis,np.newaxis,np.newaxis,np.newaxis], 
                           (nlevs,ntimes,nlats,nlons)) * \
               np.tile(recent.hyam.values[...,np.newaxis,np.newaxis,np.newaxis], (1,ntimes,nlats,nlons)) + \
               np.tile(recent.PS.PS.values, (nlevs,1,1,1)) * \
               np.tile(recent.hybm.values[...,np.newaxis,np.newaxis,np.newaxis], (1,ntimes,nlats,nlons))
        p_past = np.tile(past.P0.values[np.newaxis,np.newaxis,np.newaxis,np.newaxis], 
                         (nlevs,ntimes,nlats,nlons)) * \
             np.tile(past.hyam.values[...,np.newaxis,np.newaxis,np.newaxis], (1,ntimes,nlats,nlons)) + \
             np.tile(past.PS.PS.values, (nlevs,1,1,1)) * \
             np.tile(past.hybm.values[...,np.newaxis,np.newaxis,np.newaxis], (1,ntimes,nlats,nlons))
        
        recent.P = xr.DataArray(p_recent, 
                                coords={'plev':levs,'time':times,'lat':lats,'lon':lons},
                                dims=['plev','time','lat','lon'])
        past.P = xr.DataArray(p_past, 
                              coords={'plev':levs,'time':times,'lat':lats,'lon':lons},
                              dims=['plev','time','lat','lon'])
        
        past.T = past.T.T.assign_coords({"time":times})
        recent.T = recent.T.T.assign_coords({"time":times})
        past.Z3 = past.Z3.assign_coords({"time":times})
        recent.Z3 = recent.Z3.assign_coords({"time":times})
        past.Q = past.Q.assign_coords({"time":times})
        recent.Q = recent.Q.assign_coords({"time":times})
        past.RELHUM = past.RELHUM.assign_coords({"time":times})
        recent.RELHUM = recent.RELHUM.assign_coords({"time":times})
    
    dT_full = (recent.T.T - past.T.T).isel(plev=nlevs-1)
    
    recent.theta = recent.T.T * (recent.P0 / recent.P) ** (Rd / cp)
    past.theta = past.T.T * (past.P0 / past.P) ** (Rd / cp)
    dtheta = recent.theta - past.theta
    
    exner = past.T.T / past.theta
        
    dz = recent.Z3.Z3 - past.Z3.Z3
    
    dZ = dz.isel(plev=nlevs-1)
    
    huss = past.Q.Q
    hurs = past.RELHUM.RELHUM
    
    ws = ((100 * huss) / (hurs * (1 - huss)))
    
    lr = (gamma_d * (1 + ((Lv * ws) / (Rd * past.T.T))) / (1 + ((epsilon * \
         ws * (Lv ** 2)) / (cp * Rd * (past.T.T ** 2))))).isel(plev=nlevs-1)
    
    dT_diabatic = (dtheta * exner * lr / (gamma_d * (1 + ((Lv * ws) / \
                  (Rd * past.T.T))))).isel(plev=nlevs-1)
    
    return dT_full, dT_diabatic, dZ, lr