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 typing import List

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_updated_{month}.nc'
    #filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/masks/{sim}/MSKCLD_2017{month}.nc'
    ds = xr.open_dataarray(filename)
    #ds = xr.open_dataset(filename)
    #print(ds['MSKCLD2'])
    #ds = ds['MSKCLD2'] == 1 # Cloudy points 
    #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

from scipy.ndimage import label

def compute_updraft_width(dataarray: xr.DataArray) -> xr.DataArray:
    """
    Compute all updraft widths (horizontally adjacent non-NaN grid cells)
    for each vertical level in the given 3D xarray DataArray.

    Parameters:
    dataarray (xr.DataArray): 3D DataArray with dimensions (level, south_north, west_east)

    Returns:
    xr.DataArray: 2D DataArray where each row corresponds to a level and contains all updraft widths.
    """
    levels = dataarray.level
    width_lists = []
    
    # Define the structure for horizontal connectivity (left-right adjacency)
    structure = np.array([[1, 1, 1], [0, 0, 0], [1, 1, 1]])
    
    for level in range(dataarray.shape[0]):
        slice_2d = dataarray.isel(level=level).values  # Extract 2D slice
        mask = ~np.isnan(slice_2d)  # Identify valid updraft points
        
        # Label connected updraft regions
        labeled_array, num_features = label(mask, structure=structure)
        
        # Collect updraft widths for this level
        widths = []
        for i in range(1, num_features + 1):
            indices = np.where(labeled_array == i)
            min_x, max_x = indices[1].min(), indices[1].max()
            width = max_x - min_x + 1  # Compute horizontal width of the cluster
            widths.append(float(width))  # Force float type here
        
        # Ensure we have at least one value (0 if no valid updrafts)
        width_lists.append(widths if widths else [float(0)])

    # Find the maximum number of widths across all levels
    max_widths = max(map(len, width_lists))

    # Pad with NaNs and explicitly cast to float
    padded_widths = np.array(
        [np.pad(w, (0, max_widths - len(w)), constant_values=np.nan) for w in width_lists], dtype=float
    )

    return xr.DataArray(padded_widths, coords=[levels, np.arange(max_widths)], dims=["level", "width"])

# updraft values 
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'/global/cfs/projectdirs/m4486/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_d01_2017-{month}*'
        #file_path = f'/global/cfs/projectdirs/m4486/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_d01_2017-04-01_00:00:00'
    elif hour_interval == '1hr':
        file_path = f'/global/cfs/projectdirs/m4486/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))

    # Read in percentile updraft values for given sim and month
    thresh_file = f'/pscratch/sd/d/dbrooks/acc2017_analysis/vertical_motion/current/updraft_percentiles_month{month}.nc' # use only the thresholds from the current sim
    #thresh_file = f'/pscratch/sd/d/dbrooks/acc2017_analysis/vertical_motion/current/downdraft_percentiles_month{month}.nc'
    thresholds = xr.open_dataset(thresh_file)
  
    # Initialize dictionaries for storing updraft statistics
    updraft_lists = {key: [] for key in thresholds.keys()}
    sum_updraft_lists = {key: [] for key in thresholds.keys()}

    # Define vertical levels (these should match your dataset)
    levels = np.linspace(0.5, 15, 59)

    # Initialize storage for all data points per level for IQR computation
    all_data_points = {key: [[] for _ in range(len(levels))] for key in thresholds.keys()}

    for i, file in enumerate(file_list):
        print(f"Index: {i}, Processing: {file}")

        # Open NetCDF file and get vertical velocity data
        ncfile = netCDF4.Dataset(file, 'r')
        data = getvar(ncfile, var_name)
        data = remove_lateral_boundaries(data)

        # Apply convective core mask
        mask = mask2d[i, :, :]
        expanded_mask = mask.expand_dims(dim={"bottom_top": data.sizes["bottom_top"]})
        data = xr.where(expanded_mask, data, float("nan"))

        # Interpolate to standard levels
        height = getvar(ncfile, "height_agl", units="km")
        height = remove_lateral_boundaries(height)
        data = interplevel(data, vert=height, desiredlev=levels)
        #data = data * -1 # for downdrafts

        # Loop through each threshold to compute frequency, sum, and collect data for IQR
        for key in thresholds.keys():
            threshold = thresholds[key]
            #print(threshold.values)

            # Filter data: keep only values above threshold
            data_above_thresh = data.where(data > threshold)

            # Compute frequency and sum (keeping existing logic unchanged)
            data_freq = (data_above_thresh > threshold).sum(dim=['south_north', 'west_east'])
            data_freq_ds = data_freq.to_dataset(name=f'{key}_freq')
            updraft_lists[key].append(data_freq_ds)
            #print(updraft_lists[key])

            data_sum = data_above_thresh.sum(dim=['south_north', 'west_east'])
            data_sum_ds = data_sum.to_dataset(name=f'{key}_sum')
            sum_updraft_lists[key].append(data_sum_ds)

            # Collect all data points per level for IQR calculation
            for level_idx, level in enumerate(levels):
                level_data = data_above_thresh.sel(level=level).values.flatten()
                level_data = level_data[~np.isnan(level_data)]  # Remove NaNs

                if len(level_data) > 0:
                    all_data_points[key][level_idx].extend(level_data.tolist())

        ncfile.close()

    print('done')

    # Compute Q25 and Q75 per level for each threshold
    q25_updraft_final = {}
    q75_updraft_final = {}
    for key in thresholds.keys():
        q25_per_level = np.zeros(len(levels))
        q75_per_level = np.zeros(len(levels))
        for level_idx in range(len(levels)):
            if len(all_data_points[key][level_idx]) > 0:
                q25_per_level[level_idx] = np.percentile(all_data_points[key][level_idx], 25)
                q75_per_level[level_idx] = np.percentile(all_data_points[key][level_idx], 75)
            else:
                q25_per_level[level_idx] = np.nan
                q75_per_level[level_idx] = np.nan
        q25_updraft_final[key] = q25_per_level
        q75_updraft_final[key] = q75_per_level

    # Convert to xarray datasets with both Q25 and Q75
    quartile_updraft_ds = {
        key: xr.Dataset({
            f'{key}_q25': (["level"], q25_updraft_final[key]),
            f'{key}_q75': (["level"], q75_updraft_final[key])
        }, coords={"level": levels}) 
        for key in thresholds.keys()
    }

    # Merge datasets
    #freq_ds_up = xr.merge([xr.concat(updraft_lists[key], dim='Time') for key in thresholds.keys()])
    #mean_ds_up = xr.merge([xr.concat(sum_updraft_lists[key], dim='Time') for key in thresholds.keys()])
    quartile_ds_up = xr.merge([quartile_updraft_ds[key] for key in thresholds.keys()])
    print(quartile_ds_up)
    print(month,climate_state) 
    
    #mean_ds_up = mean_ds_up.drop_attrs(deep=True)
    #freq_ds_up = freq_ds_up.drop_attrs(deep=True)

    #print(mean_ds_up['w_99_sum'])

    #return freq_ds_up, mean_ds_up, iqr_ds_up
    return quartile_ds_up

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=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='QHAIL')

        '''
        ##### Add individual hydrometeor classes to dataset ######
        q_vars=['QRAIN','QICE','QSNOW','QGRAUP']
        for var in q_vars:
            q_data = getvar(ncfile,var) * 1000

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

            q_data = interplevel(q_data,vert=height, desiredlev=np.linspace(0.5,15,59))
            q_data = q_data.mean(dim=['south_north','west_east'])

            data[var] = q_data
        '''

        #data = cloud_top_height(data)
        #data = vertical_velocity_mask(data)
        array_list.append(data)
        ncfile.close()
        #print(data)
        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

# Finding percentiles
def read_in_monthly_data3(month, hour_interval, climate_state, var_name, mask2d, num_bins=3900):
    """
    Reads in all NetCDF files for a given month, extracts updraft and downdraft
    vertical velocity data, and computes the 99th percentile per vertical level.

    Parameters:
        month (str): Month you want data for, e.g., '04'
        climate_state (str): 'current', 'future', or 'future_urban'
        var_name (str): Variable to load, e.g., "wa"
        mask2d (xarray.DataArray): 2D mask to apply to the data
        num_bins (int): Number of bins for histogram per level

    Returns:
        Two xarray Datasets containing 99th percentile values of downdrafts and updrafts.
    """

    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-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}*'
    file_list = sorted(glob.glob(file_path))

    # Define vertical levels (these should match your dataset)
    levels = np.linspace(0.5, 15, 59)

    # Histogram bins for all levels
    bin_edges = np.linspace(1, 40, num_bins)  # Adjust range based on dataset

    # Initialize histogram count arrays for each level
    #downdraft_counts = np.zeros((len(levels), len(bin_edges) - 1))
    updraft_counts = np.zeros((len(levels), len(bin_edges) - 1))

    for i, file in enumerate(file_list):
        print(f"Index: {i}, Processing: {file}")

        # Open the file lazily
        ncfile = netCDF4.Dataset(file, 'r')

        # Read vertical velocity data
        data = getvar(ncfile, var_name)
    
        # Apply the mask
        mask = mask2d[i, :, :]
        expanded_mask = mask.expand_dims(dim={"bottom_top": data.sizes["bottom_top"]})
        data = xr.where(expanded_mask, data, float("nan"))

        # Get height levels
        height = getvar(ncfile, "height_agl", units="km")

        ##### Downdrafts #####
        downdrafts = data.where(data < -1) * -1  # Mask for downdrafts
        downdrafts = interplevel(downdrafts, vert=height, desiredlev=levels)
        #print(downdrafts)

        ##### Updrafts #####
        #updrafts = data.where(data > 2)  # Mask for updrafts
        #updrafts = interplevel(updrafts, vert=height, desiredlev=levels)

        # Flatten data per level and update histograms
        for level_idx, level in enumerate(levels):
            #level_downdrafts = downdrafts.sel(level=level).values.flatten()
            level_updrafts = downdrafts.sel(level=level).values.flatten()

            # Update histograms per level
            #downdraft_counts[level_idx] += np.histogram(level_downdrafts[~np.isnan(level_downdrafts)], bins=bin_edges)[0]
            updraft_counts[level_idx] += np.histogram(level_updrafts[~np.isnan(level_updrafts)], bins=bin_edges)[0]

        ncfile.close()

    print("Estimating final 99th percentiles per level...")

    # Compute the nth percentile per level
    def compute_percentile(bin_edges, bin_counts, percentile):
        cdf = np.cumsum(bin_counts, axis=1) / np.sum(bin_counts, axis=1, keepdims=True)
        return np.array([bin_edges[np.searchsorted(cdf[i], percentile)] for i in range(len(bin_counts))])

    #final_downdraft_99 = compute_percentile(bin_edges, downdraft_counts, 0.99)
    final_updraft_99 = compute_percentile(bin_edges, updraft_counts, percentile=0.99)
    final_updraft_90 = compute_percentile(bin_edges, updraft_counts, percentile=0.90)
    final_updraft_70 = compute_percentile(bin_edges, updraft_counts, percentile=0.70)
    final_updraft_50 = compute_percentile(bin_edges, updraft_counts, percentile=0.50)

    # Store as xarray Dataset
    #final_downdraft_ds = xr.Dataset({"wa": (["level"], final_downdraft_99)}, coords={"level": levels})
    final_updraft_ds = xr.Dataset({"w_99": (["level"], final_updraft_99),
                                   "w_90": (["level"], final_updraft_90),
                                   "w_70": (["level"], final_updraft_70),
                                   "w_50": (["level"], final_updraft_50)}, coords={"level": levels})
    

    print(f"Completed processing for {month}, {climate_state}")

    return final_updraft_ds

# updraft width
def read_in_monthly_data4(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_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))
  
    # Initialize lists and statistics
    updraft_list = []

    # Standard deviation components (Welford's Method)
    levels = np.linspace(0.5, 15, 59)  # Match vertical levels
    mean_widths = np.zeros(len(levels))
    M2_widths = np.zeros(len(levels))  # Sum of squared differences
    count_widths = np.zeros(len(levels))

    for i, file in enumerate(file_list):
        print(f"Processing: {file}")

        # Read NetCDF file
        ncfile = netCDF4.Dataset(file, 'r')
        data = getvar(ncfile, var_name)

        # Get height and interpolate to standard levels
        height = getvar(ncfile, "height_agl", units="km")
        data = interplevel(data, vert=height, desiredlev=levels)

        # Select updrafts above 5 m/s
        updrafts = data.where(data > 5)

        # Compute updraft widths for each level
        widths = compute_updraft_width(updrafts)
        widths = widths.where(widths > 1) # only keep widths with more than 1 grid cell

        # Store the mean width for each level
        mean_widths_ds = widths.mean(dim="width").to_dataset(name="mean_updraft_width")
        updraft_list.append(mean_widths_ds)

        # Update running statistics for standard deviation
        for level_idx in range(len(levels)):
            valid_widths = widths.isel(level=level_idx).values
            valid_widths = valid_widths[~np.isnan(valid_widths)]

            if len(valid_widths) == 0:
                continue  # Skip if no valid updrafts

            for w in valid_widths:
                count_widths[level_idx] += 1
                delta = w - mean_widths[level_idx]
                mean_widths[level_idx] += delta / count_widths[level_idx]
                M2_widths[level_idx] += delta * (w - mean_widths[level_idx])

        ncfile.close()

    print('Done processing all files.')

    # Calculate final standard deviation
    std_widths = np.sqrt(M2_widths / (count_widths - 1))

    # Create xarray dataset for the final standard deviation
    std_widths_ds = xr.Dataset({"std_updraft_width": (["level"], std_widths)}, coords={"level": levels})

    # Combine all mean width data into a single dataset
    combined_ds = xr.concat(updraft_list, dim='Time')
    combined_ds['std_updraft_width'] = std_widths_ds['std_updraft_width']


    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_ds

# Binning updraft speeds
def read_in_monthly_data5(month, hour_interval, climate_state, var_name, mask2d, num_bins=50):
    """
    Reads in all NetCDF files for a given month, extracts updraft and downdraft
    vertical velocity data.

    Parameters:
        month (str): Month you want data for, e.g., '04'
        climate_state (str): 'current', 'future', or 'future_urban'
        var_name (str): Variable to load, e.g., "wa"
        mask2d (xarray.DataArray): 2D mask to apply to the data
        num_bins (int): Number of bins for histogram per level

    Returns:
        Two xarray Datasets containing histogram of downdrafts and updrafts.
    """

    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-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}*'
    file_list = sorted(glob.glob(file_path))

    # Define vertical levels (these should match your dataset)
    levels = np.linspace(0.5, 15, 59)

    # Histogram bins for all levels
    bin_edges = np.linspace(1, 60, num_bins)  # Adjust range based on dataset
    print(bin_edges)

    # Initialize histogram count arrays for each level
    #downdraft_counts = np.zeros((len(levels), len(bin_edges) - 1))
    updraft_counts = np.zeros((len(levels), len(bin_edges) - 1))
    #print(np.shape(updraft_counts))
    
    for i, file in enumerate(file_list):
        print(f"Index: {i}, Processing: {file}")

        # Open the file lazily
        ncfile = netCDF4.Dataset(file, 'r')

        # Read vertical velocity data
        data = getvar(ncfile, var_name)
        #print(data)
        data = remove_lateral_boundaries(data)
    
        # Apply the mask
        mask = mask2d[i, :, :]
        expanded_mask = mask.expand_dims(dim={"bottom_top": data.sizes["bottom_top"]})
        data = xr.where(expanded_mask, data, float("nan"))

        # Get height levels
        height = getvar(ncfile, "height_agl", units="km")
        height=remove_lateral_boundaries(height)

        ##### Downdrafts #####
        downdrafts = data.where(data < -1) * -1  # Mask for downdrafts
        downdrafts = interplevel(downdrafts, vert=height, desiredlev=levels)

        print(downdrafts)

        ##### Updrafts #####
        #updrafts = data.where(data > 2)  # Mask for updrafts
        #updrafts = interplevel(updrafts, vert=height, desiredlev=levels)

        # Flatten data per level and update histograms
        for level_idx, level in enumerate(levels):
            #level_downdrafts = downdrafts.sel(level=level).values.flatten()
            level_updrafts = downdrafts.sel(level=level).values.flatten()

            # Update histograms per level
            #downdraft_counts[level_idx] += np.histogram(level_downdrafts[~np.isnan(level_downdrafts)], bins=bin_edges)[0]
            updraft_counts[level_idx] += np.histogram(level_updrafts[~np.isnan(level_updrafts)], bins=bin_edges)[0]
            #print(updraft_counts)

        ncfile.close()

    updraft_da = xr.DataArray(
        updraft_counts,
        dims=["level", "bin_edge"],
        coords={
            "level": levels,
            "bin_edge": bin_edges[:-1]  # Use left bin edges since it's a count
        },
        name="downdraft_count"
    )

    updraft_da = updraft_da.to_dataset(name='w_count')
    #print(updraft_da)

    print(f"Completed processing for {month}, {climate_state}")

    #return final_updraft_ds
    return updraft_da

# Binning downdraft speeds by selected dates
def read_in_monthly_data6(month, hour_interval, climate_state, var_name, mask2d, selected_dates: List[pd.Timestamp], num_bins=55):
    """
    Reads NetCDF files for a month and computes downdraft histograms for 12Z-12Z days from selected_dates.
    """
    # Define file path
    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}*'

    file_list = sorted(glob.glob(file_path))
    levels = np.linspace(0.5, 15, 59)
    bin_edges = np.linspace(1, 55, num_bins)

    # Histogram accumulator
    downdraft_counts = np.zeros((len(levels), len(bin_edges) - 1))

    for i, file in enumerate(file_list):
        print(f"Index: {i}, Processing: {file}")
        ncfile = netCDF4.Dataset(file, 'r')

        # Read vertical velocity
        data = getvar(ncfile, var_name)
        #print(data)

        # Get file timestamp from metadata or file name (whichever is reliable)
        try:
            time_str = data.Time.values  # works for WRF files
            print(time_str)
            file_time = pd.to_datetime(time_str)
        except:
            print("Could not read timestamp from file, skipping...")
            ncfile.close()
            continue

        # Determine if file_time falls in any 12Z–12Z window
        in_range = False
        for day_start in selected_dates:
            if day_start <= file_time < (day_start + pd.Timedelta(hours=24)):
                in_range = True
                break
        if not in_range:
            print(f"Skipping {file_time} — not in selected 12Z windows")
            ncfile.close()
            continue
        

        data = remove_lateral_boundaries(data)

        # Apply mask
        mask = mask2d[i, :, :]
        expanded_mask = mask.expand_dims(dim={"bottom_top": data.sizes["bottom_top"]})
        data = xr.where(expanded_mask, data, float("nan"))

        # Get height
        height = getvar(ncfile, "height_agl", units="km")
        height = remove_lateral_boundaries(height)

        # Downdrafts only
        downdrafts = data.where(data < -1) * -1
        downdrafts = interplevel(downdrafts, vert=height, desiredlev=levels)

        # Per-level histograms
        for level_idx, level in enumerate(levels):
            level_vals = downdrafts.sel(level=level).values.flatten()
            downdraft_counts[level_idx] += np.histogram(level_vals[~np.isnan(level_vals)], bins=bin_edges)[0]

        ncfile.close()

    # Build result
    downdraft_da = xr.DataArray(
        downdraft_counts,
        dims=["level", "bin_edge"],
        coords={"level": levels, "bin_edge": bin_edges[:-1]},
        name="downdraft_count"
    )

    return downdraft_da.to_dataset(name="w_count")

# Binning hydrometeors
def read_in_monthly_data7(
    month,
    hour_interval,
    climate_state,
    mask2d,
    hydrometeor_vars= ['QCLOUD','QRAIN','QICE','QSNOW','QGRAUP','QHAIL'],
    selected_dates=None,
    num_bins=500,
    vertical_levels=np.linspace(0.5, 15, 59)
):
    """
    Reads in NetCDF files for a month and computes vertical histograms for specified hydrometeors.
    
    Parameters:
        month (str): Month string like '04'
        hour_interval (str): '1hr' or '3hr'
        climate_state (str): 'current', 'future', 'future_urban'
        hydrometeor_vars (list): Variables to include (e.g., ['QHAIL', 'QRAIN'])
        mask2d (xr.DataArray): (Time, south_north, west_east)
        selected_dates (list): Optional list of pd.Timestamp (12Z start of day)
        num_bins (int): Histogram bins
        vertical_levels (array): Vertical interpolation levels (in km)
    
    Returns:
        xr.Dataset: Combined histograms by hydrometeor
    """
    # File lists
    if hour_interval == '3hr':
        file_path = f'/global/cfs/projectdirs/m4486/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-01_00:00:00'
    elif hour_interval == '1hr':
        file_path = f'/global/cfs/projectdirs/m4486/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_hourly_d01_2017-{month}*'
    
    file_list = sorted(glob.glob(file_path))
    
    hail_path = f'/global/cfs/projectdirs/m4486/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-04-01_00:00:00'
    hail_list = sorted(glob.glob(hail_path))

    # Warm cloud points
    #wc = xr.open_dataarray(f'/pscratch/sd/d/dbrooks/acc2017_analysis/cloud_stuff/{climate_state}/warm_cloud_points_convcores_month{month}.nc')
    #print(wc)
    #print(file_list)
    #print(hail_list)

    bin_edges = np.linspace(0.01, 5.0, num_bins)  # g/kg
    histograms = {
        var: np.zeros((len(vertical_levels), num_bins - 1)) for var in hydrometeor_vars
    }

    for i, (file, hail_file) in enumerate(zip(file_list, hail_list)):

        try:
            ncfile = netCDF4.Dataset(file, 'r')
            ts = getvar(ncfile, 'T2')
            time_str = ts.Time.values
            if time_str is None:
                continue
            file_time = pd.to_datetime(time_str)
            print(file_time)
        except:
            print('time is empty')
            continue

        try:
            height = getvar(ncfile, "height_agl", units="km")
            height = remove_lateral_boundaries(height)
        except:
            continue

        hail_nc = netCDF4.Dataset(hail_file, 'r')
        hail = xr.open_dataset(xr.backends.NetCDF4DataStore(hail_nc))
        #print(hail)
        hail = hail['QHAIL'][0,:,:,:] * 1000 # kg/kg to g/kg

        hail['XLAT'] = ts['XLAT']
        hail['XLONG'] = ts['XLONG']
        hail = remove_lateral_boundaries(hail)

        var_data = {"QHAIL": hail}
        for var in hydrometeor_vars:
            if var != "QHAIL":
                try:
                    raw = getvar(ncfile, var) * 1000
                    #print(raw)

                    raw = remove_lateral_boundaries(raw)
                    var_data[var] = raw
                except:
                    continue
        
        mask = mask2d[i, :, :]
        expanded_mask = mask.expand_dims(dim={"bottom_top": next(iter(var_data.values())).sizes["bottom_top"]})
        for var in var_data:
            var_data[var] = xr.where(expanded_mask, var_data[var], np.nan)

        for var in hydrometeor_vars:
            if var in var_data:
                interp = interplevel(var_data[var], vert=height, desiredlev=vertical_levels)
                for level_idx, level in enumerate(vertical_levels):
                    vals = interp.sel(level=level).values.flatten()
                    vals = vals[~np.isnan(vals)]
                    histograms[var][level_idx] += np.histogram(vals, bins=bin_edges)[0]

        ncfile.close()
        hail_nc.close()

    ds_out = xr.Dataset()
    for var in hydrometeor_vars:
        ds_out[f"{var}_hist"] = xr.DataArray(
            histograms[var],
            dims=["level", "bin_edge"],
            coords={"level": vertical_levels, "bin_edge": bin_edges[:-1]},
        )

    return ds_out

# Frequency of cloudy points (Qtotal > 1e-6 kg/kg)
def compute_cloudy_point_frequency(
    month,
    hour_interval,
    climate_state,
    mask2d,
    hydrometeor_vars=['QCLOUD','QRAIN','QICE','QSNOW','QGRAUP','QHAIL'],
    threshold_kg_kg=1e-6,
    vertical_levels=np.linspace(0.5, 15, 59)
):
    """
    Computes the frequency of "cloudy points" where Qtotal > threshold for each vertical level.
    
    Parameters:
        month (str): Month string like '04'
        hour_interval (str): '1hr' or '3hr'
        climate_state (str): 'current', 'future', 'future_urban'
        mask2d (xr.DataArray): (Time, south_north, west_east)
        hydrometeor_vars (list): Hydrometeor variables to sum
        threshold_kg_kg (float): Threshold in kg/kg (default 1e-6)
        vertical_levels (array): Vertical interpolation levels (in km)
    
    Returns:
        xr.Dataset: Frequency of cloudy points per vertical level
    """
    # File paths
    if hour_interval == '3hr':
        file_path = f'/global/cfs/projectdirs/m4486/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_d01_2017-{month}*'
    elif hour_interval == '1hr':
        file_path = f'/global/cfs/projectdirs/m4486/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_hourly_d01_2017-{month}*'
    
    file_list = sorted(glob.glob(file_path))
    
    hail_path = f'/global/cfs/projectdirs/m4486/yuwei/Climate_Impact/long-term/data/{climate_state}/Ze/wrfout_hourly_d01_2017-{month}*'
    hail_list = sorted(glob.glob(hail_path))

    # Convert threshold to g/kg for consistency with data processing
    threshold_g_kg = threshold_kg_kg * 1000

    # Initialize frequency counter for each level
    cloudy_frequency = np.zeros(len(vertical_levels))

    for i, (file, hail_file) in enumerate(zip(file_list, hail_list)):
        print(f"Processing file {i+1}/{len(file_list)}: {file}")

        try:
            ncfile = netCDF4.Dataset(file, 'r')
            ts = getvar(ncfile, 'T2')
            time_str = ts.Time.values
            if time_str is None:
                continue
            file_time = pd.to_datetime(time_str)
        except:
            print('Could not read time, skipping...')
            continue

        try:
            height = getvar(ncfile, "height_agl", units="km")
            height = remove_lateral_boundaries(height)
        except:
            print('Could not read height, skipping...')
            ncfile.close()
            continue

        # Read QHAIL from separate file
        try:
            hail_nc = netCDF4.Dataset(hail_file, 'r')
            hail = xr.open_dataset(xr.backends.NetCDF4DataStore(hail_nc))
            hail = hail['QHAIL'][0,:,:,:] * 1000  # kg/kg to g/kg
            hail['XLAT'] = ts['XLAT']
            hail['XLONG'] = ts['XLONG']
            hail = remove_lateral_boundaries(hail)
        except:
            print('Could not read hail, skipping...')
            ncfile.close()
            continue

        # Read other hydrometeor variables
        var_data = {"QHAIL": hail}
        for var in hydrometeor_vars:
            if var != "QHAIL":
                try:
                    raw = getvar(ncfile, var) * 1000  # kg/kg to g/kg
                    raw = remove_lateral_boundaries(raw)
                    var_data[var] = raw
                except:
                    print(f'Could not read {var}, skipping variable...')
                    continue
        
        # Apply mask
        mask = mask2d[i, :, :]
        expanded_mask = mask.expand_dims(dim={"bottom_top": next(iter(var_data.values())).sizes["bottom_top"]})
        
        # Sum all hydrometeors to get Qtotal
        qtotal = None
        for var in var_data:
            masked_var = xr.where(expanded_mask, var_data[var], 0.0)  # Use 0 for summation
            if qtotal is None:
                qtotal = masked_var
            else:
                qtotal = qtotal + masked_var
        
        # Apply mask again to set non-masked regions to NaN
        qtotal = xr.where(expanded_mask, qtotal, np.nan)
        
        # Interpolate to vertical levels
        qtotal_interp = interplevel(qtotal, vert=height, desiredlev=vertical_levels)
        
        # Count cloudy points per level
        for level_idx, level in enumerate(vertical_levels):
            level_data = qtotal_interp.sel(level=level).values.flatten()
            level_data = level_data[~np.isnan(level_data)]
            cloudy_frequency[level_idx] += np.sum(level_data > threshold_g_kg)

        ncfile.close()
        hail_nc.close()

    print(f"Completed processing for {month}, {climate_state}")

    # Create output dataset
    ds_out = xr.Dataset({
        "cloudy_point_frequency": xr.DataArray(
            cloudy_frequency,
            dims=["level"],
            coords={"level": vertical_levels},
            attrs={"description": f"Frequency of grid points where Qtotal > {threshold_kg_kg} kg/kg",
                   "units": "count"}
        )
    })

    return ds_out

#print(selected_12z_days)

for month in monlist:
    for sim in sim_list:
        # Load in mask
        mask = load_in_mask(sim,month) 
        mask = remove_lateral_boundaries(mask)
    

        ds1 = read_in_monthly_data(month, '3hr', sim, 'wa', mask)
        #ds1 = read_in_monthly_data5(month, '3hr', sim, 'wa', mask)
        #ds1 = read_in_monthly_data7(month, '3hr', sim, mask)
        #ds2 = read_in_monthly_data7(month, '3hr', sim, mask, selected_dates=selected_12z_days_cgtf)

        #ds1 = compute_cloudy_point_frequency(
        #    month,
        #    '3hr',
        #    sim,
        #    mask)
        
        #ds1.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/hydrometeor_data/{sim}/q_var_histogram_allclouds_month{month}.nc')
        #ds1.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/hydrometeor_data/{sim}/q_var_histogram_non_warmclouds_convcores_month{month}.nc')
        #ds2.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/hydrometeor_data/{sim}/q_var_histogram_ccountdays_month{month}.nc')
        #print(ds1)
        
        #ds2.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/vertical_motion/{sim}/downdraft_sum_bylevel_UPDATED_month{month}.nc')
        #ds3.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/vertical_motion/{sim}/downdraft_std_bylevel_UPDATED_month{month}.nc')
        ds1.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/vertical_motion/{sim}/updraft_iqr_bylevel_UPDATED_month{month}.nc')
        #ds1.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/cloud_stuff/{sim}/vertical_cloud_frequency_month{month}.nc')
        print('done',sim, month)
        #updraft_99_sum.to_netcdf(f'/pscratch/sd/d/dbrooks/vertical_motion/{sim}/downdraft_99_sum_month{month}.nc')
        

