from netCDF4 import Dataset
import matplotlib.pyplot as plt
import matplotlib.colors as mcolors
import matplotlib.colors as Normalize
import cartopy.crs as ccrs
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
from matplotlib.colors import LinearSegmentedColormap, TwoSlopeNorm
from scipy.stats import linregress, pearsonr

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

def load_in_mask(sim, month):
    #filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/masks/{sim}/convective_core_mask_withtimes_month{month}.nc'
    filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/masks/{sim}/convective_wind_mask_{month}'
    #filename = f'/pscratch/sd/d/dbrooks/radar_data/{sim}/composite_ref_month{month}.nc'
    ds = xr.open_dataarray(filename)
    #print(ds)
    #ds = ds['__xarray_dataarray_variable__']
    #ds = ds['MSKCLD2'] == 1
    #ds = ds['mdbz'] > 5 
    return ds

def remove_lateral_boundaries(current_ds):
    # Access the latitude and longitude arrays (XLAT, XLONG)
    lats = current_ds['XLAT']
    lons = current_ds['XLONG']

    # Get the shape of the latitude and longitude arrays
    n_lat, n_lon = lats.shape
    # Exclude 15 grid cells from each side (latitude and longitude)
    lat_slice = slice(15, n_lat - 15)
    lon_slice = slice(15, n_lon - 15)

    # Subset the data using the grid cell indices
    current_ds = current_ds.isel(south_north=lat_slice, west_east=lon_slice)

    return current_ds

def mean_rh_below_cold_pool(rh: xr.DataArray, ds_cph: xr.Dataset) -> xr.DataArray:
    """
    Compute the mean relative humidity (RH) below the cold pool height at each horizontal grid cell.

    Parameters:
    ----------
    rh : xr.DataArray
        3D relative humidity with dimensions (level, south_north, west_east).
        The 'level' coordinate must be in meters.
    ds_cph : xr.Dataset
        Dataset containing 'cold_pool_height' in kilometers (km) with dimensions (south_north, west_east).

    Returns:
    -------
    xr.DataArray
        2D DataArray (south_north, west_east) of mean RH below the cold pool height.
    """

    cph_meters = ds_cph['cold_pool_height']

    # Broadcast cold pool height to match RH shape
    cph_broadcast = cph_meters.expand_dims({'level': rh.level})

    # Create a 3D level array for comparison (shape: [level, 1, 1])
    level_values = rh['level'].values[:, np.newaxis, np.newaxis]

    # Create the mask: True where level <= cold pool height
    mask = level_values <= cph_broadcast

    # Apply the mask and compute the mean along the vertical level dimension
    rh_masked = rh.where(mask)
    mean_rh = rh_masked.mean(dim='level', skipna=True)

    return mean_rh



def read_in_monthly_data(month, hour_interval, climate_state, var_name, mask2d):
    """
    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_0*: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))

    # load in proper cp height file
    cp_height = xr.open_dataset(f'/pscratch/sd/d/dbrooks/acc2017_analysis/vertical_motion/{climate_state}/coldpool_height_month{month}_accurate.nc')

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

        # Interpolate to standard levels
        height = getvar(ncfile, "height_agl", units="m")
        data = interplevel(data, vert=height, desiredlev=np.linspace(30,6000,60))

        data = remove_lateral_boundaries(data)
        #print(data)

        # Get mask
        mask = mask2d.sel(Time=data.Time.values)
        cp_height2 = cp_height.sel(Time=data.Time.values)
        #print(cp_height['cold_pool_height'].max())

        # Expand mask to include the 'bottom_top' dimension
        expanded_mask = mask.expand_dims(dim={"level": data.sizes["level"]})

        # Apply the mask
        data = xr.where(expanded_mask, data, float("nan"))
        #print(data)

        data = mean_rh_below_cold_pool(data, cp_height2)
        data=data.to_dataset(name='rh_up_to_cpheight')

        array_list.append(data)
        ncfile.close()

    combined_da = xr.concat(array_list, dim='Time')
    return combined_da

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

for month in monlist:
    for sim in sim_list:
        mask = load_in_mask(sim,month)
        #print(mask)
        #ds2 = read_in_monthly_data3(month, '3hr', sim, mask)
        ds2 = read_in_monthly_data(month, '3hr', sim, 'rh', mask)
        ds2.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/humidity_data/{sim}/rh_upto_cp_height_month{month}_accurate.nc')
        print(f'done {sim} {month}')