############################### For looking at the clausius claperyon relationship ##############################
from netCDF4 import Dataset
#import h5py
import matplotlib.pyplot as plt
import matplotlib.colors as mcolors
import matplotlib.colors as Normalize
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
from dask.diagnostics import ProgressBar
import os
import re
import warnings
import gc
from datetime import datetime
import seaborn as sns
from matplotlib.colors import LinearSegmentedColormap, TwoSlopeNorm
import math

monlist = ['06'] # months in the simulation
sim_list = ['current','future','future_urban']
select_subregion = False
select_time_window = False

# Keep chunks modest to avoid loading full domain arrays into memory.
CHUNKS = {"Time": 24, "south_north": 120, "west_east": 120}
dask.config.set({"array.slicing.split_large_chunks": True})

# Performance controls: bounded sampling keeps runtime predictable and RAM safe.
ENABLE_DASK_PROGRESS = True
TIME_SLICE_FOR_SAMPLING = 8

if ENABLE_DASK_PROGRESS:
    ProgressBar().register()


def compute_exact_binned_extremes(temp_da, precip_da, sat_deficit_da, temp_bins, min_data_points, time_chunk=8, label=''):
    """
    Reproduce original binning exactly, but stream through time slices to avoid OOM.
    """
    n_bins = len(temp_bins) - 1
    counts = np.zeros(n_bins, dtype=np.int64)
    bin_paths = [f'/tmp/cc_bin_{os.getpid()}_{label}_{i}.bin' for i in range(n_bins)]

    # Ensure old temp files from a previous interrupted run do not leak into current stats.
    for path in bin_paths:
        if os.path.exists(path):
            os.remove(path)

    n_time = int(temp_da.sizes.get('Time', 0))
    chunk_iter = tqdm(
        range(0, n_time, time_chunk),
        desc=f'Chunks {label}',
        leave=False,
        dynamic_ncols=True,
    )

    try:
        for t0 in chunk_iter:
            t1 = min(t0 + time_chunk, n_time)
            temp_chunk = temp_da.isel(Time=slice(t0, t1)).compute().values
            precip_chunk = precip_da.isel(Time=slice(t0, t1)).compute().values
            sd_chunk = sat_deficit_da.isel(Time=slice(t0, t1)).compute().values

            valid_base = (
                np.isfinite(temp_chunk)
                & np.isfinite(precip_chunk)
                & np.isfinite(sd_chunk)
                & (sd_chunk < 0.5)
            )

            if not np.any(valid_base):
                del temp_chunk, precip_chunk, sd_chunk, valid_base
                gc.collect()
                continue

            temp_c = temp_chunk - 273.15

            for i in range(n_bins):
                in_bin = (temp_c >= temp_bins[i]) & (temp_c < temp_bins[i + 1])
                mask = valid_base & in_bin
                if not np.any(mask):
                    continue

                vals = precip_chunk[mask].astype(np.float32, copy=False)
                with open(bin_paths[i], 'ab') as f:
                    vals.tofile(f)
                counts[i] += vals.size

            del temp_chunk, precip_chunk, sd_chunk, valid_base, temp_c
            gc.collect()

        extreme_precip = []
        bin_centers = []

        reduce_iter = tqdm(
            range(n_bins),
            desc=f'Reduce {label}',
            leave=False,
            dynamic_ncols=True,
        )
        for i in reduce_iter:
            if counts[i] < min_data_points or not os.path.exists(bin_paths[i]):
                continue

            bin_precip = np.fromfile(bin_paths[i], dtype=np.float32)
            if bin_precip.size < min_data_points:
                continue

            threshold = np.percentile(bin_precip, 99)
            extreme_bin_precip = bin_precip[bin_precip >= threshold]
            extreme_precip.append(float(extreme_bin_precip.mean()))
            bin_centers.append((temp_bins[i] + temp_bins[i + 1]) / 2)

            del bin_precip, extreme_bin_precip

        return extreme_precip, bin_centers
    finally:
        for path in bin_paths:
            if os.path.exists(path):
                os.remove(path)

if select_subregion == False:
    region = 'Full Domain'
else:
    bounding_box=[-97.926091,27.270908,-83.709230,38.648338] # min_lon,min_lat,max_lon,max_lat (southeast/gulf coast)
    region = 'Southeast' # can change based on bounding box

if select_time_window == True:
    start_time = '2017-06-19T18:00:00'
    end_time = '2017-06-24T06:00:00'
else:
    start_time = 0
    end_time = 0

def hex_to_rgb(value):
    '''
    Converts hex to rgb colours
    value: string of 6 characters representing a hex colour.
    Returns: list length 3 of RGB values'''
    value = value.strip("#") # removes hash symbol if present
    lv = len(value)
    return tuple(int(value[i:i + lv // 3], 16) for i in range(0, lv, lv // 3))

# Downscaling function
def downscsale_precip(ds):
    # Assuming `ds` is your xarray dataset
    ds_downscaled = ds.coarsen(
        south_north=6,  # Downsampling factor of 6 for south_north (to get resolution of 12km)
        west_east=6,     # Downsampling factor of 6 for west_east
        boundary="trim"
    ).mean()  

    return ds_downscaled

def read_in_monthly_data(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}*'
    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') 
        data = getvar(ncfile,var_name)
        data = data.to_dataset(name=var_name)
        #print(data.attrs)
        array_list.append(data)
        ncfile.close()

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

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

    return combined_ds


def open_chunked_dataset(filename, primary_var):
    ds = xr.open_dataset(filename, chunks=CHUNKS)
    keep_vars = [primary_var]
    for coord_var in ("XLAT", "XLONG"):
        if coord_var in ds.data_vars:
            keep_vars.append(coord_var)
    return ds[keep_vars]

# Downscaling function
def downscsale_precip(ds1):
    # Assuming `ds` is your xarray dataset
    ds_downscaled1 = ds1.coarsen(
        south_north=6,  # Downsampling factor of 6 for south_north (to get resolution of 12km)
        west_east=6,     # Downsampling factor of 6 for west_east
        boundary="trim"
    ).mean()
    
    return ds_downscaled1

precip_current_list=[]
precip_future_list=[]
precip_future_urban_list=[]

# Precip
for month in monlist:
    for sim in sim_list:
        if sim == 'current':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/precip_data/current/hourly_precip_data_month{month}.nc'
            ds = open_chunked_dataset(filename, 'RAINNC')
            #ds = downscsale_precip(ds)
            precip_current_list.append(ds)
            #precip_current_ds = ds
        elif sim == 'future':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/precip_data/future/hourly_precip_data_month{month}.nc'
            ds = open_chunked_dataset(filename, 'RAINNC')
            #ds = downscsale_precip(ds)
            precip_future_list.append(ds)
            #precip_future_ds = ds
        elif sim == 'future_urban':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/precip_data/future_urban/hourly_precip_data_month{month}.nc'
            ds = open_chunked_dataset(filename, 'RAINNC')
            #ds = downscsale_precip(ds)
            precip_future_urban_list.append(ds)
            #precip_future_urban_ds = ds
        else:
            print('incorrect input simulation name')
        print(sim)
    print(month)
print('precip')


temp_current_list=[]
temp_future_list=[]
temp_future_urban_list=[]


# Temperature
for month in monlist:
    for sim in sim_list:
        #ds = read_in_monthly_data(month, '3hr', sim, 'T2')
        if sim == 'current':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/temp_data/current/temp_data_month{month}.nc'
            #print(ds)
            #ds.to_netcdf(filename)
            ds = open_chunked_dataset(filename, 'T2')
            #ds = downscsale_precip(ds)
            temp_current_list.append(ds)
        elif sim == 'future':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/temp_data/future/temp_data_month{month}.nc'
            #ds.to_netcdf(filename)
            ds = open_chunked_dataset(filename, 'T2')
            #ds = downscsale_precip(ds)
            temp_future_list.append(ds)
        elif sim == 'future_urban':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/temp_data/future_urban/temp_data_month{month}.nc'
            #ds.to_netcdf(filename)
            ds = open_chunked_dataset(filename, 'T2')
            #ds = downscsale_precip(ds)
            temp_future_urban_list.append(ds)
        else:
            print('incorrect input simulation name')
        print(sim)
    print(month)
print('temp')


dptemp_current_list=[]
dptemp_future_list=[]
dptemp_future_urban_list=[]

# Dew Point Temperature
for month in monlist:
    for sim in sim_list:
        #ds = read_in_monthly_data(month, '3hr', sim, 'T2')
        if sim == 'current':
            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/humidity_data/current/dewpoint2m_data_month{month}.nc'
            #print(ds)
            #ds.to_netcdf(filename)
            ds = open_chunked_dataset(filename, 'td2')
            #ds = downscsale_precip(ds)
            dptemp_current_list.append(ds)
        elif sim == 'future':
            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/humidity_data/future/dewpoint2m_data_month{month}.nc'
            #ds.to_netcdf(filename)
            ds = open_chunked_dataset(filename, 'td2')
            #ds = downscsale_precip(ds)
            dptemp_future_list.append(ds)
        elif sim == 'future_urban':
            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/humidity_data/future_urban/dewpoint2m_data_month{month}.nc'
            #ds.to_netcdf(filename)
            ds = open_chunked_dataset(filename, 'td2')
            #ds = downscsale_precip(ds)
            dptemp_future_urban_list.append(ds)
        else:
            print('incorrect input simulation name')
        print(sim)
    print(month)
print('temp')

# Surface pressure
pressure_current_list=[]
pressure_future_list=[]
pressure_future_urban_list=[]

for month in monlist:
    for sim in sim_list:
        if sim == 'current':
            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/pressure_data/Current/slp_data_month{month}.nc'
            ds = open_chunked_dataset(filename, 'slp')
            pressure_current_list.append(ds)
            #precip_current_ds = ds
        elif sim == 'future':
            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/pressure_data/Future/slp_data_month{month}.nc'
            ds = open_chunked_dataset(filename, 'slp')
            pressure_future_list.append(ds)
            #precip_future_ds = ds
        elif sim == 'future_urban':
            filename = f'/global/cfs/projectdirs/m2637/dbrooks/2017_seasonal_analysis/pressure_data/Future_urban/slp_data_month{month}.nc'
            ds = open_chunked_dataset(filename, 'slp')
            pressure_future_urban_list.append(ds)
            #precip_future_urban_ds = ds
        else:
            print('incorrect input simulation name')
        print(sim)
    print(month)


precip_current_ds = xr.concat(precip_current_list, dim='Time')
del precip_current_list
print('precip current done')
gc.collect() 
precip_future_ds = xr.concat(precip_future_list, dim='Time')
del precip_future_list
print('precip future done')
gc.collect() 
precip_future_urban_ds = xr.concat(precip_future_urban_list, dim='Time')
del precip_future_urban_list
print('precip future urban done')
gc.collect()  # Force garbage collection to free up memory

temp_current_ds = xr.concat(temp_current_list, dim='Time')
del temp_current_list
print('temp current done')
gc.collect() 
temp_future_ds = xr.concat(temp_future_list, dim='Time')
del temp_future_list
print('temp future done')
gc.collect() 
temp_future_urban_ds = xr.concat(temp_future_urban_list, dim='Time')
del temp_future_urban_list
print('temp future urban done')
gc.collect()  # Force garbage collection to free up memory


dptemp_current_ds = xr.concat(dptemp_current_list, dim='Time')
dptemp_future_ds = xr.concat(dptemp_future_list, dim='Time')
dptemp_future_urban_ds = xr.concat(dptemp_future_urban_list, dim='Time')

pressure_current_ds = xr.concat(pressure_current_list, dim='Time')
pressure_future_ds = xr.concat(pressure_future_list, dim='Time')
pressure_future_urban_ds = xr.concat(pressure_future_urban_list, dim='Time')



del dptemp_current_list,dptemp_future_list,dptemp_future_urban_list
gc.collect()  # Force garbage collection to free up memory

del pressure_current_list,pressure_future_list,pressure_future_urban_list
gc.collect()  # Force garbage collection to free up memory


    # 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

print('done')
precip_current_ds, precip_future_ds, precip_future_urban_ds = remove_lateral_boundaries(precip_current_ds, precip_future_ds, precip_future_urban_ds)
temp_current_ds, temp_future_ds, temp_future_urban_ds = remove_lateral_boundaries(temp_current_ds, temp_future_ds, temp_future_urban_ds)
dptemp_current_ds, dptemp_future_ds, dptemp_future_urban_ds = remove_lateral_boundaries(dptemp_current_ds, dptemp_future_ds, dptemp_future_urban_ds)
pressure_current_ds, pressure_future_ds, pressure_future_urban_ds = remove_lateral_boundaries(pressure_current_ds, pressure_future_ds, pressure_future_urban_ds)
print('done')

# Select subregion if desired
if select_subregion == True:
    def select_region(bounding_box, current_ds, future_ds, future_urban_ds):
        min_lon,min_lat,max_lon,max_lat = bounding_box[0], bounding_box[1], bounding_box[2], bounding_box[3]

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

        # Create a boolean mask for the region of interest
        region_mask = (lats >= min_lat) & (lats <= max_lat) & (lons >= min_lon) & (lons <= max_lon)

        # Subset the data using the bounding box
        current_ds = current_ds.where(region_mask, drop=True)
        future_ds = future_ds.where(region_mask, drop=True)
        future_urban_ds = future_urban_ds.where(region_mask, drop=True)

        return current_ds, future_ds, future_urban_ds

    #precip_current_ds, precip_future_ds, precip_future_urban_ds = select_region(bounding_box, precip_current_ds, precip_future_ds, precip_future_urban_ds)

# resample precip to 3 hour timesteps
def resample_precip_to_3hr(precip_ds):
    """
    Resamples the precipitation dataset to match the 3-hour timestep of the wind dataset by summing
    the precipitation over each 3-hour interval.

    Parameters:
    precip_ds (xarray.Dataset): The original hourly precipitation dataset.

    Returns:
    xarray.Dataset: The resampled precipitation dataset at 3-hour intervals.
    """
    # Resample the dataset to 3-hour intervals and sum the precipitation over those intervals
    precip_resampled = precip_ds.resample(Time='3h').mean()

    return precip_resampled

# Extract time values as a pandas Index
time_values = precip_current_ds.Time.values

# Adjust only the first time value by subtracting 1 hour
time_values[0] = pd.Timestamp(time_values[0]) - pd.Timedelta(hours=1)

# Reassign the modified time back to the dataset
precip_current_ds = precip_current_ds.assign_coords(Time=time_values)
precip_future_ds = precip_future_ds.assign_coords(Time=time_values)
precip_future_urban_ds = precip_future_urban_ds.assign_coords(Time=time_values)

precip_current_ds = resample_precip_to_3hr(precip_current_ds)
precip_future_ds = resample_precip_to_3hr(precip_future_ds)
precip_future_urban_ds = resample_precip_to_3hr(precip_future_urban_ds) 

if select_time_window == True:
    def subset_time_window(ds, start_time, end_time):
        """
        Subset the dataset based on a specified time window.

        Parameters:
        ds (xarray.Dataset): The dataset to subset.
        start_time (str or datetime): The start time of the window (inclusive).
        end_time (str or datetime): The end time of the window (inclusive).

        Returns:
        xarray.Dataset: The subset of the dataset within the specified time window.
        """
        # Subset the dataset by time
        subset_ds = ds.sel(Time=slice(start_time, end_time))
        
        return subset_ds

    #precip_current_ds = subset_time_window(precip_current_ds, start_time, end_time)
    #precip_future_ds = subset_time_window(precip_future_ds, start_time, end_time)
    #precip_future_urban_ds = subset_time_window(precip_future_urban_ds, start_time, end_time) 

simulation_data = {
    "Current": (temp_current_ds, precip_current_ds, dptemp_current_ds, pressure_current_ds),
    "Future": (temp_future_ds, precip_future_ds, dptemp_future_ds, pressure_future_ds),
    "Future-Urban": (temp_future_urban_ds, precip_future_urban_ds, dptemp_future_urban_ds, pressure_future_urban_ds)
}

del temp_current_ds, precip_current_ds, dptemp_current_ds, pressure_current_ds
del temp_future_ds, precip_future_ds, dptemp_future_ds, pressure_future_ds
del temp_future_urban_ds, precip_future_urban_ds, dptemp_future_urban_ds, pressure_future_urban_ds
gc.collect()  # Force garbage collection to free up memory

######################################## Functions to get exponential fits ##########################################
from scipy.optimize import curve_fit
import numpy as np
import matplotlib.pyplot as plt
from tqdm import tqdm

def epi_scaling_equation(T, EPI20, r):
    """
    The exponential form of the EPI-temperature relationship:
    EPI = EPI20 * (1 + r)**(T - 20)
    """
    return EPI20 * ((1 + r) ** (T - 20))

def fit_scaling_curve_bootstrap(temp_values, precip_values, label, n_bootstrap=1000):
    """
    Fits an exponential curve to the extreme precipitation data using bootstrapping.
    Returns the best-fit r value and its 95% confidence interval.

    Parameters:
    - temp_values: Array of temperature bin centers (°C)
    - precip_values: Array of extreme precipitation (mm/hr)
    - label: Name of the simulation (for printing)
    - n_bootstrap: Number of bootstrap resamples (default: 1000)

    Returns:
    - r: Best-fit scaling rate per °C
    - r_confidence_interval: 95% confidence interval of r from bootstrap
    """
    # Remove any NaN values
    temp_values = np.array(temp_values)
    precip_values = np.array(precip_values)

    valid_indices = ~np.isnan(temp_values) & ~np.isnan(precip_values)
    temp_values = temp_values[valid_indices]
    precip_values = precip_values[valid_indices]

    # Fit the curve to the original data to get the "best estimate"
    popt, pcov = curve_fit(epi_scaling_equation, temp_values, precip_values, 
                           p0=[np.mean(precip_values[temp_values == 15]), 0.07], 
                           bounds=([0, 0], [np.inf, 1]))

    # Extract the initial best-fit r
    EPI20_best, r_best = popt

    # ----------------------------------------------------
    # Bootstrap Resampling to Calculate Uncertainty
    # ----------------------------------------------------
    r_bootstrap = []  # Store r values from each resample
    EPI20_bootstrap = []

    # Run the bootstrap with tqdm progress bar
    for _ in tqdm(range(n_bootstrap), desc=f'Bootstrapping {label}'):
        # Randomly resample the data with replacement
        resample_idx = np.random.choice(len(temp_values), size=len(temp_values), replace=True)
        temp_resample = temp_values[resample_idx]
        precip_resample = precip_values[resample_idx]

        # Fit the curve to the resampled data
        try:
            popt_resample, _ = curve_fit(epi_scaling_equation, temp_resample, precip_resample, 
                                         p0=[EPI20_best, r_best], bounds=([0, 0], [np.inf, 1]))
            EPI20_bootstrap.append(popt_resample[0])
            r_bootstrap.append(popt_resample[1])
        except:
            # If fit fails, just append NaN (this happens occasionally)
            EPI20_bootstrap.append(np.nan)
            r_bootstrap.append(np.nan)

    # Convert results to arrays and remove failed fits
    r_bootstrap = np.array(r_bootstrap)
    r_bootstrap = r_bootstrap[~np.isnan(r_bootstrap)]

    # ----------------------------------------------------
    # Calculate the Best-Fit and 95% Confidence Interval
    # ----------------------------------------------------
    r_conf_low = np.percentile(r_bootstrap, 2.5)
    r_conf_high = np.percentile(r_bootstrap, 97.5)
    r_median = np.percentile(r_bootstrap, 50)

    # ----------------------------------------------------
    # Plot the Fitted Curve
    # ----------------------------------------------------
    T_fit = np.linspace(temp_values.min(), temp_values.max(), 100)
    EPI_fit = epi_scaling_equation(T_fit, EPI20_best, r_best)

    # ----------------------------------------------------
    # Print the Results
    # ----------------------------------------------------
    print(f"Simulation: {label}")
    print(f"Best-fit r: {r_best:.4f} ({r_best*100:.2f}%) per °C")
    print(f"95% Confidence Interval: [{r_conf_low:.4f}, {r_conf_high:.4f}]")
    print("--------------------------------------------------")
    
    return r_best, r_conf_low, r_conf_high, T_fit, EPI_fit, EPI20_best

# Function to calculate saturation vapor pressure
def get_sat_vap_pressure(T):
    """
    Calculates saturation vapor pressure (es) from temperature (T in Kelvin).
    """
    a1 = 6.1121  # hPa
    a3 = 17.502
    a4 = 32.19   # K
    To = 273.16  # K
    t = ((T - To) / (T - a4)) * a3
    return a1 * np.exp(t)

# Function to calculate actual vapor pressure
def get_vap_pressure(Td):
    """
    Calculates actual vapor pressure (e) from dew point temperature (Td in Kelvin).
    """
    Td = Td + 273.16 # C to K

    a1 = 6.1121  # hPa
    a3 = 17.502
    a4 = 32.19   # K
    To = 273.16  # K
    t = ((Td - To) / (Td - a4)) * a3
    e  = a1 * np.exp(t)

    return e


def get_saturation_deficit(ds_temperature, ds_dew_point, ds_pressure):
    """
    Calculates saturation deficit for the given xarray Datasets. Uses methods from Wang and Sun (2022)
    
    Inputs:
    - ds_temperature: xarray.DataArray for temperature (Kelvin).
    - ds_dew_point: xarray.DataArray for dew point (Kelvin).
    - ds_pressure: xarray.DataArray for pressure (hPa).

    Outputs:
    - Saturation deficit as an xarray.DataArray (g/kg).
    """
    E = 0.622  # kg/kg, ratio of gas constants for dry air and water vapor

    # Calculate vapor pressures
    es = get_sat_vap_pressure(ds_temperature)
    e = get_vap_pressure(ds_dew_point)

    #print(e)

    de = es - e

    # Calculate saturation deficit
    dq = (E * de) / (ds_pressure - (1 - E) * de)

    #print(dq*1000)
    return dq*1000 #g/kg


###################### Bivariate Scaling #############################

def plot_extreme_precip_scaling_with_saturation(simulation_data, month, bin_size=0.5, min_data_points=200):
    """
    Plots extreme precipitation scaling with temperature for a specified month, with lines for each simulation.
    Only includes points where the saturation deficit is below 0.5 g/kg.

    Parameters:
    - simulation_data: dict of xarray DataSets for simulations, each containing temperature and precipitation.
                       Expects the format { "label1": (temp_ds1, precip_ds1), "label2": (temp_ds2, precip_ds2), ... }
    - month: str, the month in 'MM' format to filter data, e.g., '04' for April.
    - bin_size: float, size of temperature bins in °C (default is 0.5).
    - min_data_points: int, minimum number of data points per bin to include it (default is 1000).
    """
    plt.figure(figsize=(5,4))

    colors = {
    "Current": 'black',
    "Future": '#1E88E5',
    "Future-Urban": '#D81B60'
    }

    exterme_precip_values = {}
    bin_centers_bysim = {}

    # Loop through each simulation and plot its scaling curve
    sim_iter = tqdm(
        simulation_data.items(),
        total=len(simulation_data),
        desc='Simulation progress',
        dynamic_ncols=True,
    )
    for label, (temp_ds, precip_ds, dp_ds, pressure_ds) in sim_iter:
        sim_iter.set_postfix_str(label)
        
        # Filter the data for the specified month
        month_idx = int(month)
        temp_month = temp_ds.sel(Time=temp_ds['Time.month'] == month_idx)
        precip_month = precip_ds.sel(Time=precip_ds['Time.month'] == month_idx)
        dp_month = dp_ds.sel(Time=dp_ds['Time.month'] == month_idx)
        pressure_month = pressure_ds.sel(Time=pressure_ds['Time.month'] == month_idx)
        

        saturation_deficit = get_saturation_deficit(
            temp_month['T2'],
            dp_month['td2'],
            pressure_month['slp']
            )

        sat_deficit_month = saturation_deficit.sel(Time=saturation_deficit['Time.month'] == month_idx)

        # Ensure all datasets have the same shape
        if temp_month['T2'].shape != precip_month['RAINNC'].shape or temp_month['T2'].shape != sat_deficit_month.shape:
            raise ValueError(f"Temperature, precipitation, and saturation deficit data shapes do not match for {label}")

        # Define temperature bins
        temp_min = 10
        temp_max = 22
        temp_bins = np.arange(temp_min, temp_max + bin_size, bin_size)

        extreme_precip, bin_centers = compute_exact_binned_extremes(
            temp_month['T2'],
            precip_month['RAINNC'],
            sat_deficit_month,
            temp_bins=temp_bins,
            min_data_points=min_data_points,
            time_chunk=TIME_SLICE_FOR_SAMPLING,
            label=label,
        )

        if not bin_centers:
            bin_centers_bysim[label] = []
            exterme_precip_values[label] = []
            continue

        bin_centers_bysim[label] = bin_centers
        exterme_precip_values[label]= extreme_precip

        # Get exponential fit 
        #r_best, r_conf_low, r_conf_high, T_fit, EPI_fit, EPI20_best = fit_scaling_curve_bootstrap(temp_values=bin_centers_bysim[label], precip_values=exterme_precip_values[label], label=label)

        # Apply a 3-bin moving average to smooth the results
        extreme_precip_smoothed = pd.Series(extreme_precip).rolling(window=3, center=True).mean().to_numpy()

        # Plot the results for this simulation
        plt.plot(bin_centers, extreme_precip_smoothed, linestyle='-', 
                 label=f'{label}', #(r={r_best*100:.2f}%)' 
                 color=colors[label], linewidth=2.5)

        # Plot the exponential fit
        #plt.plot(T_fit, EPI_fit, linestyle='--', color=colors[label])

        # Plot the Confidence Interval 
        #plt.fill_between(T_fit, 
        #                epi_scaling_equation(T_fit, EPI20_best, r_conf_low),
        #                epi_scaling_equation(T_fit, EPI20_best, r_conf_high),
        #                color=colors[label], alpha=0.15)

    if not bin_centers:
        raise RuntimeError('No temperature bins met min_data_points. Lower min_data_points or widen bin range.')

    temp_range = (bin_centers[0], bin_centers[-1])
    temp_start, temp_end = temp_range
    temp_values = np.linspace(temp_start, temp_end, 100)
    yints = [0.25,0.5,1,2,4,8,16]
    for y_intercept in yints:  # Generate lines with y-intercepts at 10, 20, ..., 50 mm/hr
        cc_line = [y_intercept * np.exp(0.07 * (T - temp_start)) for T in temp_values]
        plt.plot(temp_values, cc_line, color='gray', linestyle='--', linewidth=2, alpha=0.5)

    plt.plot(temp_values[0], cc_line[0], color='gray', linestyle='--', linewidth=2, label=f'CC Scaling (7%)', alpha=0.5)

    if monlist[0] == '04':
        month = 'April'
    elif monlist[0] == '05':
        month = 'May'
    elif monlist[0] == '06':
        month = 'June'

    # Plot formatting
    plt.xlabel('Temperature (°C)')
    plt.ylabel('Precipitation Rate (mm hr$^{-1}$)')
    plt.title(f'>99th Percentile Precip Rate Scaling with SD<0.5 g kg$^{{-1}}$\n{month}')
    plt.yscale('log')
    plt.ylim(1,60)
    plt.legend(loc='upper left', fontsize=9)
    plt.xlim(10,22)
    plt.grid(True, axis='x', alpha=0.5)

    # SAVE THE PLOT FIRST 
    plt.savefig(f'cc_scaling_{month}.png', dpi=300, bbox_inches='tight')
    #plt.show()

    return exterme_precip_values, bin_centers_bysim   


epi_values, bin_centers = plot_extreme_precip_scaling_with_saturation(simulation_data, month=monlist[0])

for ds_group in simulation_data.values():
    for ds in ds_group:
        ds.close()