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


def hex_to_rgb(value):
    """
    Converts hex to RGB tuple.
    """
    value = value.strip("#")
    lv = len(value)
    return tuple(int(value[i:i + lv // 3], 16) for i in range(0, lv, lv // 3))


def create_custom_diverging_colormap(levels):
    """
    Notebook-style diverging colormap: cool blues to warm reds.
    """
    custom_rgb = [
        '#37569e', '#5578af', '#719cc1', '#b8e3e7', '#ffffff',
        '#ffd784', '#fda75e', '#eb7949', '#d24d37', '#b12122'
    ]

    cmap = LinearSegmentedColormap.from_list("CoolWarmCustom", custom_rgb, N=levels)
    cmap.set_under(np.array(hex_to_rgb('#0a348c')) / 255)
    cmap.set_over(np.array(hex_to_rgb('#840000')) / 255)
    return cmap

# Remove boundaries
def remove_lateral_boundaries(current_ds,future_ds,future_urban_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)
    future_ds = future_ds.isel(south_north=lat_slice, west_east=lon_slice)
    future_urban_ds = future_urban_ds.isel(south_north=lat_slice, west_east=lon_slice)

    return current_ds,future_ds,future_urban_ds

monlist=['04']

# rh850mb 
c_rh850 = xr.open_dataset(f'/pscratch/sd/d/dbrooks/acc2017_analysis/humidity_data/current/relative_humidity_850mb_month{monlist[0]}.nc')
f_rh850 = xr.open_dataset(f'/pscratch/sd/d/dbrooks/acc2017_analysis/humidity_data/future/relative_humidity_850mb_month{monlist[0]}.nc')
fu_rh850 = xr.open_dataset(f'/pscratch/sd/d/dbrooks/acc2017_analysis/humidity_data/future_urban/relative_humidity_850mb_month{monlist[0]}.nc')

#print(c_rh850)

# cloud mask
c_cloud = xr.open_dataset(f'/pscratch/sd/d/dbrooks/acc2017_analysis/masks/current/MSKCLD_2017{monlist[0]}.nc')
f_cloud = xr.open_dataset(f'/pscratch/sd/d/dbrooks/acc2017_analysis/masks/future/MSKCLD_2017{monlist[0]}.nc')
fu_cloud = xr.open_dataset(f'/pscratch/sd/d/dbrooks/acc2017_analysis/masks/future_urban/MSKCLD_2017{monlist[0]}.nc')

c_rh850,f_rh850,fu_rh850 = remove_lateral_boundaries(c_rh850,f_rh850,fu_rh850)
c_cloud,f_cloud,fu_cloud = remove_lateral_boundaries(c_cloud,f_cloud,fu_cloud)

def get_clear_sky_points(c_ds, f_ds, fu_ds):
    c_clear_points = c_cloud['MSKCLD1'] < 1
    f_clear_points = f_cloud['MSKCLD1'] < 1
    fu_clear_points = fu_cloud['MSKCLD1'] < 1

    c_ds = xr.where(c_clear_points, c_ds, float("nan"))
    f_ds = xr.where(f_clear_points, f_ds, float("nan"))
    fu_ds = xr.where(fu_clear_points, fu_ds, float("nan"))

    return c_ds, f_ds, fu_ds

c_rh850,f_rh850,fu_rh850 = get_clear_sky_points(c_rh850,f_rh850,fu_rh850)

print('done1')

from scipy.stats import gaussian_kde


def plot_param_pdf(c_ds, f_ds, var, threshold, max_bin, iters, fu_ds=None, bin_width=0.25):
    """
    Plots the smoothed PDF of a variable from three simulations with mean in legend,
    using a 3-bin rolling average and scaling frequency to 0–100%.
    """
    
    def preprocess(dataarray):
        """Flatten, clean, and clip data."""
        arr = dataarray.values.astype(np.float32).flatten()
        arr = arr[~np.isnan(arr)]
        #if max_bin is not None:
        #    arr = np.clip(arr, None, max_bin)
        return arr


    def bootstrap_from_kde(data, bins, iters=1000, size=100_000):
        kde = gaussian_kde(data)
        resampled = kde.resample(iters * size).reshape(iters, size)

        counts = np.zeros((iters, len(bins)), dtype=np.float32)

        for i in range(iters):
            sample = resampled[i]
            # Don't filter; let values > bins[-1] fall into last bin via clip
            indices = np.digitize(sample, bins, right=False) - 1
            indices = np.clip(indices, 0, len(bins) - 1)

            counts[i, :] = np.bincount(indices, minlength=len(bins))

        freqs = (counts / counts.sum(axis=1, keepdims=True)) * 100
        avg_freq = np.nanmean(freqs, axis=0)
        avg_freq = avg_freq[1:-1] # removes the last index

        print('done bootstrapping')
        #return pd.Series(avg_freq).rolling(window, center=True).mean().to_numpy()
        return avg_freq


    
    # Preprocess each dataset
    c_filtered = preprocess(c_ds[var].where(c_ds[var] >= threshold))
    f_filtered = preprocess(f_ds[var].where(f_ds[var] >= threshold))
    fu_filtered = preprocess(fu_ds[var].where(fu_ds[var] >= threshold))

    # Shared bin range for all datasets
    #min_val = min(c_filtered.min(), f_filtered.min(), fu_filtered.min())
    min_val = threshold
    effective_max = max_bin if max_bin is not None else max(c_filtered.max(), f_filtered.max()) #,fu_filtered.max())
    bins = np.arange(min_val, effective_max + bin_width, bin_width)

    # With this:
    c_kde = bootstrap_from_kde(c_filtered, bins, iters=iters, size=100_000)
    f_kde = bootstrap_from_kde(f_filtered, bins, iters=iters, size=100_000)
    fu_kde = bootstrap_from_kde(fu_filtered, bins, iters=iters, size=100_000)

    bins = bins[1:-1] # removes the last index

    # Means
    c_mean = np.mean(c_filtered)
    f_mean = np.mean(f_filtered)
    fu_mean = np.mean(fu_filtered)

    print(c_kde.sum())
    print(f_kde.sum())
    #print(c_kde)

    # Plot
    plt.figure(figsize=(4, 3))
    plt.plot(bins, c_kde, label=f"Current (Mean: {c_mean:.2f})", color="black", linewidth=2)
    plt.plot(bins, f_kde, label=f"Future (Mean: {f_mean:.2f})", color="#1E88E5", linewidth=2)
    plt.plot(bins, fu_kde, label=f"Future+Urban (Mean: {fu_mean:.2f})", color="#D81B60", linewidth=2)

    # Labels
    label_dict = {
        'cape_2d': ('CAPE', '(J kg$^{-1}$)'),
        'mcin': ('CIN', '(J kg$^{-1}$)'),
        'helicity': ('0-3km SRH', '(m$^{2}$ s$^{-2}$)'),
        '0_1km_srh': ('0-1km SRH', '(m$^{2}$ s$^{-2}$)'),
        'ehi': ('EHI', ''),
        'stp': ('STP', ''),
        'q': ('2m q', '(g kg$^{-1}$)'),
        'rh2m': ('2m RH', '(%)'),
        'rh_850': ('850mb RH', '(%)'),
        'bulk_shear_0_6km': ('0-6km Bulk Shear', '(m s$^{-1}$)'),
        'lcl': ('LCL', '(m)')
    }
    title, units = label_dict.get(var, (var, ''))

    # Month detection (you likely set monlist earlier)
    try:
        month = {'04': 'April', '05': 'May', '06': 'June'}.get(monlist[0], '')
    except:
        month = ''

    plt.title(f"{title} ({month})")
    #plt.title(f"{title} ({month})")
    plt.xlabel(f"{title} {units}")
    plt.ylabel("Normalized Frequency (%)")
    #plt.ylim(0,5)
    #plt.yscale('log')
    #plt.yscale('symlog', linthresh=0.01)
    #custom_ticks = [0.001,0.01,0.1,1]
    plt.xlim(10,100)
    #custom_ticks = [0.1,1]
    plt.minorticks_off()
    #plt.yticks(custom_ticks)
    plt.legend(fontsize=8, loc=0)
    #plt.legend(fontsize=8, loc='lower left')
    plt.grid(alpha=0.5, linestyle=':')
    plt.tight_layout()
    # SAVE THE PLOT FIRST 
    plt.savefig(f'rh850_pdf_{month}2.png', dpi=300, bbox_inches='tight')
    #plt.show()


def _get_month_name(month_code):
    return {'04': 'April', '05': 'May', '06': 'June'}.get(month_code, month_code)


def plot_mean_monthly_rh850_spatial(current_ds, future_ds, future_urban_ds, var='rh_850', save_fig=True):
    """
    Plot monthly mean 850mb RH and differences between simulations.
    Layout follows the notebook style: Current, Future-Current, Future+Urban-Future.
    """
    month_name = _get_month_name(monlist[0])

    # Core fields (already masked to clear-sky points upstream in this script)
    mean_current = current_ds[var].mean(dim='Time')
    mean_future = future_ds[var].mean(dim='Time')
    #mean_future_urban = future_urban_ds[var].mean(dim='Time')

    diff_future_current = mean_future - mean_current
    #diff_future_urban_future = mean_future_urban - mean_future

    lats = current_ds['XLAT']
    lons = current_ds['XLONG']

    # RH magnitude colormap
    base_cmap = matplotlib.colormaps['plasma']
    rh_levels = np.arange(10, 101, 10)
    rh_cmap = mcolors.ListedColormap(base_cmap(np.linspace(0.1, 0.8, len(rh_levels))))
    rh_norm = mcolors.BoundaryNorm(rh_levels, rh_cmap.N)
    rh_cmap.set_under('white')
    rh_cmap.set_over(base_cmap(0.99))

    # Difference colormap (discrete diverging) using notebook custom palette
    diff_levels = [-6, -4, -2, -1, 1, 2, 4, 6]
    import colormaps 
    diff_cmap1 = colormaps.rdbu_11_r
    diff_cmap = diff_cmap1[1:10]

    # Extract individual colors from the base colormap
    colors = [diff_cmap1(i / (len(diff_levels) - 1)) for i in range(len(diff_levels))]

    diff_cmap.set_over(colors[-1])   # Upper bound color
    diff_cmap.set_under(colors[0])  # Lower bound color

    # Create a normalization for the contour levels
    diff_norm = mcolors.BoundaryNorm(diff_levels, diff_cmap.N)

    fig, axs = plt.subplots(1, 2, figsize=(12, 4), subplot_kw={'projection': ccrs.PlateCarree()})

    # Panel 1: Current monthly mean
    ax = axs[0]
    pb = ax.pcolormesh(lons, lats, mean_current, cmap=rh_cmap, norm=rh_norm, transform=ccrs.PlateCarree(), zorder=1)
    ax.add_feature(cfeature.STATES, edgecolor='black', linewidths=0.35, alpha=0.8)
    ax.add_feature(cfeature.COASTLINE, edgecolor='black', linewidths=0.35, alpha=0.8)
    gl = ax.gridlines(draw_labels=True, linewidth=1, color='gray', alpha=0.5, linestyle='--', zorder=2)
    gl.top_labels = False
    gl.right_labels = False
    gl.left_labels = True
    gl.bottom_labels = True
    ax.set_title('Mean 850mb RH (%)', loc='left', fontsize=12)
    ax.set_title('(Current)', loc='right', fontsize=12)
    plt.colorbar(pb, ax=ax, orientation='vertical', fraction=0.05, pad=0.01, shrink=0.8, extend='both')

    current_mean = mean_current.mean(dim=['south_north', 'west_east'])
    current_max = mean_current.max(dim=['south_north', 'west_east'])
    txt = f"Mean: {current_mean:.2f}\\nMax: {current_max:.2f}"
    ax.text(
        -112.0,
        42.3,
        txt,
        fontsize=8,
        color='black',
        weight='bold',
        transform=ccrs.PlateCarree(),
        bbox=dict(facecolor='white', alpha=0.9, boxstyle='round,pad=0.5')
    )

    # Panel 2: Future - Current
    ax = axs[1]
    pb = ax.pcolormesh(lons, lats, diff_future_current, cmap=diff_cmap, norm=diff_norm, transform=ccrs.PlateCarree(), zorder=1)
    ax.add_feature(cfeature.STATES, edgecolor='black', linewidths=0.35, alpha=0.8)
    ax.add_feature(cfeature.COASTLINE, edgecolor='black', linewidths=0.35, alpha=0.8)
    gl = ax.gridlines(draw_labels=True, linewidth=1, color='gray', alpha=0.5, linestyle='--', zorder=2)
    gl.top_labels = False
    gl.right_labels = False
    gl.left_labels = False
    gl.bottom_labels = True
    ax.set_title(r'$\Delta$850mb RH (%)', loc='left', fontsize=12)
    ax.set_title('(Warming Effect)', loc='right', fontsize=12)
    plt.colorbar(pb, ax=ax, orientation='vertical', fraction=0.05, pad=0.01, shrink=0.8, extend='both')

    fc_mean = diff_future_current.mean(dim=['south_north', 'west_east'])
    fc_min = diff_future_current.min(dim=['south_north', 'west_east'])
    fc_max = diff_future_current.max(dim=['south_north', 'west_east'])
    txt = f"Mean: {fc_mean:.2f} (%)\nMin: {fc_min:.2f} (%)\nMax: {fc_max:.2f} (%)"
    ax.text(
        -112.0,
        43.2,
        txt,
        fontsize=8,
        color='black',
        weight='bold',
        transform=ccrs.PlateCarree(),
        bbox=dict(facecolor='white', alpha=0.9, boxstyle='round,pad=0.5')
    )
    '''
    # Panel 3: Future+Urban - Future
    ax = axs[2]
    pb = ax.pcolormesh(lons, lats, diff_future_urban_future, cmap=diff_cmap, norm=diff_norm, transform=ccrs.PlateCarree(), zorder=1)
    ax.add_feature(cfeature.STATES, edgecolor='black', linewidths=0.35, alpha=0.8)
    ax.add_feature(cfeature.COASTLINE, edgecolor='black', linewidths=0.35, alpha=0.8)
    gl = ax.gridlines(draw_labels=True, linewidth=1, color='gray', alpha=0.5, linestyle='--', zorder=2)
    gl.top_labels = False
    gl.right_labels = False
    gl.left_labels = False
    gl.bottom_labels = True
    ax.set_title(r'$\Delta$850mb RH (%)', loc='left', fontsize=12)
    ax.set_title('(Urbanization Effect)', loc='right', fontsize=12)
    plt.colorbar(pb, ax=ax, orientation='vertical', fraction=0.05, pad=0.01, shrink=0.8, extend='both')

    uf_mean = diff_future_urban_future.mean(dim=['south_north', 'west_east'])
    uf_min = diff_future_urban_future.min(dim=['south_north', 'west_east'])
    uf_max = diff_future_urban_future.max(dim=['south_north', 'west_east'])
    txt = f"Mean: {uf_mean:.2f} (%)\\nMin: {uf_min:.2f} (%)\\nMax: {uf_max:.2f} (%)"
    ax.text(
        -112.0,
        44.2,
        txt,
        fontsize=8,
        color='black',
        weight='bold',
        transform=ccrs.PlateCarree(),
        bbox=dict(facecolor='white', alpha=0.9, boxstyle='round,pad=0.5')
    )
    '''

    fig.suptitle(f'Mean 850mb RH and Differences Between Simulations in {month_name}', fontsize=14)
    plt.tight_layout(rect=[0, 0.01, 1, 0.99])

    if save_fig:
        plt.savefig(f'rh850_spatial_monthly_{month_name}.png', dpi=300, bbox_inches='tight')
    else:
        plt.show()

#plot_param_pdf(c_rh850,f_rh850, fu_ds=fu_rh850, var='rh_850', bin_width=1, max_bin=99, threshold=0.5, iters=100)

plot_mean_monthly_rh850_spatial(c_rh850, f_rh850, fu_rh850, var='rh_850', save_fig=True)