from netCDF4 import Dataset
import matplotlib.pyplot as plt
import matplotlib.colors as mcolors
import matplotlib.colors as Normalize
import cartopy.crs as ccrs
import cartopy.feature as cfeature
from cartopy.mpl.ticker import LongitudeFormatter, LatitudeFormatter
import matplotlib.ticker as mticker
import matplotlib
import xarray as xr
import netCDF4
import numpy as np
import pandas as pd
import glob
import dask
import os
from matplotlib.colors import LinearSegmentedColormap, TwoSlopeNorm
from scipy.stats import linregress, pearsonr

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


def _integrate_profile(b_profile, z_profile, threshold):
    # Find first index where buoyancy > threshold
    indices = np.where(b_profile > threshold)[0]
    #print(indices)
    if indices.size > 0:
        k_threshold = indices[0]
        b_sub = b_profile[:k_threshold]
        z_sub = z_profile[:k_threshold]
        return np.trapz(b_sub, z_sub)
    else:
        return np.nan

def integrate_buoyancy_fast(buoyancy: xr.DataArray, threshold=-0.005):
    """
    Fast version: integrate buoyancy up to first level > threshold.
    Uses xarray.apply_ufunc for vectorization.
    """
    level_heights = buoyancy['level'].values  # (level,)

    # Apply ufunc over south_north, west_east
    result = xr.apply_ufunc(
        _integrate_profile,
        buoyancy,
        xr.DataArray(level_heights, dims=["level"]),
        kwargs={"threshold": threshold},
        input_core_dims=[["level"], ["level"]],
        output_core_dims=[[]],
        vectorize=True,
        dask="parallelized",
        output_dtypes=[np.float32]
    )

    result = result.where(result < 0)
    result.attrs['description'] = 'Integrated buoyancy from lowest level to cold pool height'
    result.attrs['units'] = 'm^2/s^2'

    return result 

from scipy.ndimage import uniform_filter

# Convert to NumPy, apply uniform filter on each level
def apply_spatial_filter(data_array, size):
    # Apply to each level separately
    filtered = np.empty_like(data_array)
    for k in range(data_array.shape[0]):
        filtered[k] = uniform_filter(data_array[k], size=size, mode='nearest')
    return filtered

# Buoyancy calculation
def read_in_monthly_data(month, hour_interval, climate_state, mask2d=None):
    """
    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-29_21: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-04-29_21:00:00'
    hail_list = sorted(glob.glob(hail_path))

    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 m (ABOVE GROUND LEVEL)

        '''
        # Threshold: 1000 m AGL
        threshold = 2000.0

        # Count levels < threshold at each horizontal grid point
        mask = height < threshold  # shape: (bottom_top, south_north, west_east)
        counts = np.sum(mask, axis=0)  # shape: (south_north, west_east)

        # Compute mean across domain
        mean_levels_1km = np.mean(counts)

        print(f"Mean number of model levels in first 2 km AGL: {mean_levels_1km:.2f}")
        '''

        q_total = interplevel(q_total,vert=height, desiredlev=np.linspace(30,2000,12))
        q_vapor = interplevel(q_vapor,vert=height, desiredlev=np.linspace(30,2000,12))
        theta = interplevel(theta,vert=height, desiredlev=np.linspace(30,2000,12))

        #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'])

        window=50 # Assume you want a 100kmx100km (50x50 since resolution is 2km) box filter

        # Apply to xarray DataArray
        smoothed_values = apply_spatial_filter(theta.values, size=(window, window))
        theta_smoothed = xr.DataArray(
            smoothed_values,
            dims=theta.dims,
            coords=theta.coords,
            name="theta_smoothed"
        )

        smoothed_values = apply_spatial_filter(q_vapor.values, size=(window, window))
        qvapor_smoothed = xr.DataArray(
            smoothed_values,
            dims=q_vapor.dims,
            coords=q_vapor.coords,
            name="qvapor_smoothed"
        )
        

        #print(theta_smoothed)

        '''
        # 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)

        ####### Need to change mean theta to a mean over smaller area #######
        total_buoyancy = g * (((theta - theta_smoothed)/theta_smoothed) + 0.61*(q_vapor-qvapor_smoothed) - 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)
        #integrated_buoyancy = compute_cpi_optimized(total_buoyancy)
        integrated_buoyancy = integrate_buoyancy_fast(total_buoyancy)
        #print(integrated_buoyancy.min(skipna=True))
        #print(integrated_buoyancy.mean(skipna=True))
        #cp_height = compute_cold_pool_height(total_buoyancy)
        #print(cp_height['cold_pool_height'].max())
        #print(integrated_buoyancy)

        #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(total_buoyancy)
        array_list.append(integrated_buoyancy)
        ncfile.close()
 

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

    combined_ds = combined_ds.to_dataset(name='integrated_b')

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

    return combined_ds

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

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