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

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

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

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 _cpi_column(b_profile, z_profile):
    """Compute CPI for a single column."""
    if np.all(np.isnan(b_profile)):
        return np.nan

    # Find Zh index (first where b > -0.005)
    mask = b_profile > -0.005
    if not np.any(mask):
        return np.nan

    zh_idx = np.argmax(mask)

    b_seg = b_profile[:zh_idx]
    z_seg = z_profile[:zh_idx]

    if len(b_seg) < 2 or np.any(np.isnan(b_seg)):
        return np.nan

    integral = np.trapz(b_seg, z_seg)
    return np.sqrt(-2 * integral) if integral < 0 else 0.0


def compute_cpi_optimized(buoyancy: xr.DataArray) -> xr.DataArray:
    """
    Optimized CPI computation using xarray vectorization.
    """
    z = buoyancy['level'].values

    # Broadcast z across the 3D domain to match buoyancy shape
    z_broadcasted = xr.DataArray(
        np.broadcast_to(z[:, None, None], buoyancy.shape),
        dims=buoyancy.dims,
        coords=buoyancy.coords
    )

    cpi = xr.apply_ufunc(
        _cpi_column,
        buoyancy,
        z_broadcasted,
        input_core_dims=[["level"], ["level"]],
        output_core_dims=[[]],
        vectorize=True,
        dask="parallelized" if hasattr(buoyancy.data, "compute") else None,
        output_dtypes=[float],
    )

    return cpi.rename("CPI")

def compute_cold_pool_height(buoyancy: xr.DataArray, threshold: float = -0.005) -> xr.DataArray:
    """
    Compute cold pool height (Zh) from buoyancy profile.
    
    Parameters:
        buoyancy (xr.DataArray): 3D buoyancy array with dims ('level', 'south_north', 'west_east')
        threshold (float): Buoyancy threshold to define end of cold pool (default -0.005 m/s^2)
    
    Returns:
        xr.DataArray: 2D array of cold pool heights (in same units as 'level' coordinate)
    """
    z = buoyancy['level'].values  # Assume 'level' is in meters (or height)
    z_broadcast = xr.DataArray(
        np.broadcast_to(z[:, None, None], buoyancy.shape),
        dims=buoyancy.dims,
        coords=buoyancy.coords
    )

    # Find first index where mask is True
    def _first_valid_height(b_col, z_col):
        valid = np.where(b_col > threshold)[0]
        return z_col[valid[0]] if valid.size > 0 else np.nan

    cold_pool_height = xr.apply_ufunc(
        _first_valid_height,
        buoyancy,
        z_broadcast,
        input_core_dims=[["level"], ["level"]],
        output_core_dims=[[]],
        vectorize=True,
        dask="parallelized" if hasattr(buoyancy.data, "compute") else None,
        output_dtypes=[float],
    )
    cold_pool_height = cold_pool_height.to_dataset(name='cold_pool_height')

    return cold_pool_height

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-04-22_18: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))
  
    downdraft_list=[]
    updraft_list=[]
    i=0
    for file in file_list:
        ##-- read file            
        ncfile = netCDF4.Dataset(file,'r')
        #print(ncfile) 
        print(file)
        data = getvar(ncfile,var_name)

        mask = mask2d[i,:,:]

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

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


        ##### Downdrafts #####
        data1 = data.where(data < -1) # mask to only get updraft values below -1 m/s
        height = getvar(ncfile, "height_agl", units="km")  # Height in km (ABOVE GROUND LEVEL)
        #data['z'] = height

        data1 = interplevel(data1,vert=height, desiredlev=np.linspace(0.5,15,59))
        data1 = data1.to_dataset(name=var_name)

        # Step 1: Flatten the DataArray
        flat_array1 = data1[var_name].values.ravel()  # or .flatten()
        #print(len(flat_array))

        # Step 2: Remove NaN values
        downdraft_array = flat_array1[~np.isnan(flat_array1)]
        #print(len(downdraft_array))


        ####### Updrafts ###########
        data2 = data.where(data > 0) # mask to only get updraft values above 0 m/s
        height = getvar(ncfile, "height_agl", units="km")  # Height in km (ABOVE GROUND LEVEL)
        #data['z'] = height

        data2 = interplevel(data2,vert=height, desiredlev=np.linspace(0.5,15,59))
        data2 = data2.to_dataset(name=var_name)

        # Step 1: Flatten the DataArray
        flat_array2 = data2[var_name].values.ravel()  # or .flatten()
        #print(len(flat_array))

        # Step 2: Remove NaN values
        updraft_array = flat_array2[~np.isnan(flat_array2)]
        #print(len(updraft_array))

        '''
        data2 = data['wa'].mean(dim=['south_north','west_east'])
        data2 = data2.to_dataset(name='w<-1')

        data['w<-5'] = data['wa'].where(data['wa'] < -5)
        data2['w<-5'] = data['w<-5'].mean(dim=['south_north','west_east'])

        data['w<-8'] = data['wa'].where(data['wa'] < -8)
        data2['w<-8'] = data['w<-8'].mean(dim=['south_north','west_east'])
        '''
        #print(data)
        downdraft_list.append(downdraft_array)
        updraft_list.append(updraft_array)

        ncfile.close()
        #print(i)
        i=i+1


    print('done')
    #combined_ds = xr.concat(array_list, dim='Time')
    combined_down = np.concatenate(downdraft_list)
    combined_up = np.concatenate(updraft_list)

    p99_downdraft = np.percentile(combined_down, 1)
    print("Downdraft 99 percentile: ", p99_downdraft)

    p99_updraft = np.percentile(combined_up, 99)
    print("Updraft 99 percentile: ", p99_updraft)

    print(month,climate_state)

    #combined_ds = combined_ds.to_dataset()

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

    return combined_down

def read_in_monthly_data2(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-22_18: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))

    # To also get hail data
    hail_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/Ze/wrfout_hourly_d01_2017-{month}*'
    hail_list = sorted(glob.glob(hail_path))

    i=0
    array_list=[]
    for file, hail_file in zip(file_list, hail_list):
        ##-- read file            
        ncfile = netCDF4.Dataset(file,'r')

        ncfile2 = netCDF4.Dataset(hail_file,'r')
        hail = xr.open_dataset(xr.backends.NetCDF4DataStore(ncfile2))
        hail=hail['QHAIL'][0,:,:,:]
        #print(hail.mean())
        print(file)
        print(hail_file)
        data = getvar(ncfile,var_name) + getvar(ncfile,'QRAIN') + getvar(ncfile,'QICE') + getvar(ncfile,'QSNOW') + getvar(ncfile,'QGRAUP') + hail
        data = data * 1000 # kg/kg to g/kg
        #data = data.where(data > 10e-5) # mask to only get values that denote cloud base value
        #data = data.where(data>=0)
        #print(data)

        mask = mask2d[i,:,:]

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

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

        #data = vinterp(wrfin=ncfile, field=data, vert_coord='ght_agl', interp_levels=np.arange(1,6.1,0.1))
        height = getvar(ncfile, "height_agl", units="km")  # Height in km (ABOVE GROUND LEVEL)
        #data = data.to_dataset(name=var_name)
        #data['z'] = height

        data = interplevel(data,vert=height, desiredlev=np.linspace(0.5,15,59))
        data = data.mean(dim=['south_north','west_east'])
        data = data.to_dataset(name='QTOTAL')
        #print(data)
        #data = cloud_top_height(data)
        #data = vertical_velocity_mask(data)
        array_list.append(data)
        ncfile.close()
        print(i)
        i=i+1

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

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

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

    return combined_ds

# Buoyancy calculation
def read_in_monthly_data3(month, hour_interval, climate_state, 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-05-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))
 
    # To also get hail data
    hail_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/Ze/wrfout_hourly_d01_2017-{month}*'
    #hail_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/Ze/wrfout_hourly_d01_2017-05-01_00:00:00'
    hail_list = sorted(glob.glob(hail_path))

    i=0
    array_list=[]
    for file, hail_file in zip(file_list, hail_list):
        ##-- read file            
        ncfile = netCDF4.Dataset(file,'r')

        ncfile2 = netCDF4.Dataset(hail_file,'r')
        hail = xr.open_dataset(xr.backends.NetCDF4DataStore(ncfile2))
        hail=hail['QHAIL'][0,:,:,:]
        #print(hail.mean())
        print(file)
        print(hail_file)
        q_total = getvar(ncfile,'QCLOUD') + getvar(ncfile,'QRAIN') + getvar(ncfile,'QICE') + getvar(ncfile,'QSNOW') + getvar(ncfile,'QGRAUP') + hail
        q_vapor = getvar(ncfile,'QVAPOR')
        theta = getvar(ncfile,'theta')

        height = getvar(ncfile, "height_agl", units="m")  # Height in km (ABOVE GROUND LEVEL)

        q_total = interplevel(q_total,vert=height, desiredlev=np.linspace(30,6000,60))
        q_vapor = interplevel(q_vapor,vert=height, desiredlev=np.linspace(30,6000,60))
        theta = interplevel(theta,vert=height, desiredlev=np.linspace(30,6000,60))

        q_total = remove_lateral_boundaries(q_total)
        q_vapor = remove_lateral_boundaries(q_vapor)
        theta = remove_lateral_boundaries(theta)

        domain_mean_theta = theta.mean(dim=['south_north','west_east'])
        domain_mean_q_vapor = q_vapor.mean(dim=['south_north','west_east'])

        #print(domain_mean_theta)

        # Get mask
        mask = mask2d.sel(Time=domain_mean_theta.Time.values)

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

        # Apply the mask to 'wa'
        q_total = xr.where(expanded_mask, q_total, float("nan"))
        q_vapor = xr.where(expanded_mask, q_vapor, float("nan"))
        theta = xr.where(expanded_mask, theta, float("nan"))

        g=9.81 # gravity constant (m/s^2)

        total_buoyancy = g * (((theta - domain_mean_theta)/domain_mean_theta) + 0.61*(q_vapor-domain_mean_q_vapor) - q_total) # with condensate loading
        #thermal_buoyancy = g * (((theta - domain_mean_theta)/domain_mean_theta)) # thermal buouyancy only
        #wv_buoyancy = g * (0.61*(q_vapor-domain_mean_q_vapor))
        #cl_buoyancy = -g * q_total
        #print(total_buoyancy)
        #cpi = compute_cpi_optimized(total_buoyancy)
        cp_height = compute_cold_pool_height(total_buoyancy)
        print(cp_height['cold_pool_height'].mean())

        #total_buoyancy = total_buoyancy.mean(dim=['south_north','west_east'])
        #thermal_buoyancy = thermal_buoyancy.mean(dim=['south_north','west_east'])
        #wv_buoyancy = wv_buoyancy.mean(dim=['south_north','west_east'])
        #cl_buoyancy = cl_buoyancy.mean(dim=['south_north','west_east'])

        #buoyancy = total_buoyancy.to_dataset(name='B_total')
        #buoyancy['B_thermal'] = thermal_buoyancy
        #buoyancy['B_vapor'] = wv_buoyancy
        #buoyancy['B_cl'] = cl_buoyancy
        #print(buoyancy)

        #array_list.append(buoyancy)
        array_list.append(cp_height)
        ncfile.close()

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

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

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

    return combined_ds


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_chunked(month, '3hr', sim, 'QRAIN', mask)
        ds2.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/hydrometeor_data/{sim}/evap_rate_qrain_month{month}_accurate.nc')
        ds3 = read_in_monthly_data_chunked(month, '3hr', sim, 'QCLOUD', mask)
        ds3.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/hydrometeor_data/{sim}/evap_rate_qcloud_month{month}_accurate.nc')
