from netCDF4 import Dataset
import matplotlib.pyplot as plt
import matplotlib.colors as mcolors
import matplotlib.colors as Normalize
import cartopy.crs as crs
import cartopy.feature as cfeature
from cartopy.mpl.ticker import LongitudeFormatter, LatitudeFormatter
import matplotlib.ticker as mticker
import matplotlib
import xarray as xr
import netCDF4
import numpy as np
import pandas as pd
import glob
import dask
import os

import wrf
from wrf import (getvar, vinterp, interplevel, to_np, latlon_coords, get_cartopy,
                 cartopy_xlim, cartopy_ylim)

def read_in_monthly_data2(month, hour_interval, climate_state, var_name):
    """
    Reads in all NetCDF files for a given month and combines them into one dataset.

    Parameters:
        month (str): month you want data for, e.g., '04'
        hour_interval (str): '1hr' or '3hr'
        climate_state (str): 'current', 'future', or 'future_urban'
        var_name (str): what variable you want to load in, e.g. "wspd_wdir10"

    Returns:
        xarray dataset of full month of data
    """

    # Determine the file path based on the input parameters
    if hour_interval == "3hr":
        file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_d01_2017-{month}*'
        #file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_d01_2017-06-01_00:00:00'
    elif hour_interval == "1hr":
        file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_hourly_d01_2017-{month}*'

    # Use glob to find all matching files for the month
    file_list = sorted(glob.glob(file_path))

    array_list=[]
    for file in file_list:
        ##-- read file            
        ncfile = netCDF4.Dataset(file,'r')
        #print(ncfile) 
        print(file)
        data = getvar(ncfile,var_name, fill_nocloud=True, missing=-999, units='K')
        #data1 = data.sel(mcape_mcin_lcl_lfc='mcape')
        #data2 = data.sel(mcape_mcin_lcl_lfc='mcin')
        #cloudfrac = getvar(ncfile,'cloudfrac')

        # Interpolate to standard levels
        #height = getvar(ncfile, "height_agl", units="km")
        #data = interplevel(data, vert=height, desiredlev=[3,8])
        #print(data)
        #lapse_rate = (data[1,:,:] - data[0,:,:]) / 5

        #data = lapse_rate.to_dataset(name='3_8km_lapse_rate')

        data = data.to_dataset(name=var_name)
        #print(data)
        #data3['mcin'] = data2
        #print(data3)
        array_list.append(data)
        ncfile.close()

    print('done')
    combined_ds = xr.concat(array_list, dim='Time')

    #combined_ds = combined_ds.to_dataset(name='vert_velo_mask')

    combined_ds[var_name].attrs['projection'] = str(combined_ds[var_name].attrs['projection'])
    #combined_ds['mcin'].attrs['projection'] = str(combined_ds[var_name].attrs['projection'])

    return combined_ds

def read_in_monthly_data3(month, hour_interval, climate_state):
    """
    Reads in all NetCDF files for a given month and combines them into one dataset.

    Parameters:
        month (str): month you want data for, e.g., '04'
        hour_interval (str): '1hr' or '3hr'
        climate_state (str): 'current', 'future', or 'future_urban'
        var_name (str): what variable you want to load in, e.g. "wspd_wdir10"

    Returns:
        xarray dataset of full month of data
    """

    # Determine the file path based on the input parameters
    if hour_interval == "3hr":
        file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_d01_2017-{month}*'
        #file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_d01_2017-06-01_00:00:00'
    elif hour_interval == "1hr":
        file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_hourly_d01_2017-{month}*'

    # Use glob to find all matching files for the month
    file_list = sorted(glob.glob(file_path))

    array_list=[]
    for file in file_list:
        ##-- read file            
        ncfile = netCDF4.Dataset(file,'r')
        #print(ncfile) 
        print(file)
        #data = getvar(ncfile,var_name)
        data = getvar(ncfile,'uvmet')
        data2 = getvar(ncfile,'uvmet10')

        # Interpolate to standard levels
        height = getvar(ncfile, "height_agl", units="km")
        data = interplevel(data, vert=height, desiredlev=[6])
        #data2 = interplevel(data2, vert=height, desiredlev=[0.01,1,3,6,10,12])
        #print(data)

        # 0-6km shear
        u_1 = data[0,:,:] - data2[0,:,:] # u component
        v_1 = data[1,:,:] - data2[1,:,:] # v component

        bulk_shear_1 = np.sqrt(u_1**2 + v_1**2)

        #print(data)
        
        '''
        # 6-10km shear
        u_6_10 = data[0,2,:,:] - data[0,1,:,:] # u component
        v_6_10 = data[1,2,:,:] - data[1,1,:,:] # v component

        bulk_shear_6_10 = np.sqrt(u_6_10**2 + v_6_10**2)

        # 3-12km mean speed
        data3 = getvar(ncfile,'wspd',units='m s-1')
        data3 = interplevel(data3, vert=height, desiredlev=np.arange(3,12.1,1))
        mean_speed = data3.mean(dim='level')
        '''
        #data = data.to_dataset(name=var_name)
        dataset = bulk_shear_1.to_dataset(name='bulk_shear_0_6km')
        #dataset['bulk_shear_6_10km'] = bulk_shear_6_10
        #dataset['3_12km_mean'] = mean_speed
        #data['uvmet10'] = data2
        #data3['mcin'] = data2
        #print(data3)
        array_list.append(dataset)
        ncfile.close()

    print('done')
    combined_ds = xr.concat(array_list, dim='Time')

    #combined_ds = combined_ds.to_dataset(name='vert_velo_mask')
    #combined_ds[var].attrs['projection'] = str(combined_ds[var].attrs['projection'])
    #combined_ds['mcin'].attrs['projection'] = str(combined_ds[var_name].attrs['projection'])

    return combined_ds

monlist = ['04','05','06'] # months in the simulation
sim_list = ['current','future','future_urban']

for month in monlist:
    for sim in sim_list:
        ds = read_in_monthly_data2(month, '3hr', sim, 'ctt')
        #ds.to_netcdf(f'/pscratch/sd/d/dbrooks/thermo_data/{sim}/bulk_shear_month{month}.nc')
        ds.to_netcdf(f'/pscratch/sd/d/dbrooks/cloud_stuff/{sim}/cloudtoptemp_month{month}.nc')
        print(sim)
    print(month)