from netCDF4 import Dataset
import matplotlib.pyplot as plt
import matplotlib.colors as mcolors
import matplotlib.colors as Normalize
import cartopy.crs as crs
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

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

def vertical_velocity_mask(ds):
    """
    Finds horizontal grid points where wa > 2 m/s for a vertical depth > 4 km.

    Parameters:
        ds (xarray.Dataset): Dataset containing vertical velocity 'wa' and height 'z'.

    Returns:
        xarray.DataArray: 2D mask where True indicates grid points with wa > 2 m/s
                          for more than 4 km of vertical depth.
    """
    # Extract data variables
    wa = ds["wa"].squeeze()  # Remove the Time dimension, shape: (bottom_top, south_north, west_east)
    z = ds["z"].squeeze()    # Remove the Time dimension, shape: (bottom_top, south_north, west_east)

    # Create a mask for wa > 2 m/s
    wa_mask = wa > 2  # True where wa > 2, shape: (bottom_top, south_north, west_east)

    # Compute the height differences between consecutive levels (bottom_top axis)
    dz = np.diff(z, axis=0)  # Shape: (bottom_top-1, south_north, west_east)

    # Compute the vertical depth of valid regions where wa > 2
    # For each column, sum the height differences where wa_mask is True
    valid_depth = np.sum(dz * wa_mask[:-1, :, :], axis=0)  # Shape: (south_north, west_east)

    # Create a mask for grid points where the valid depth is greater than 4 km
    thick_mask = valid_depth >= 4.0  # Boolean mask, shape: (south_north, west_east)

    # Convert the mask to an xarray.DataArray
    mask_da = xr.DataArray(
        thick_mask,
        dims=("south_north", "west_east"),
        coords={
            "XLONG": ds["XLONG"],
            "XLAT": ds["XLAT"],
        },
        name="vertical_velocity_mask",
        attrs={"description": "Mask for wa > 2 m/s over > 4 km depth"}
    )
    
    return mask_da

def cloud_top_height(ds):
    """
    Determines the horizontal grid points where the cloud thickness is greater than 5 km.
    
    Parameters:
        ds (xarray.Dataset): The input dataset with QCLOUD and vertical levels.

    Returns:
        xarray.DataArray: 2D mask where True indicates grid points with cloud thickness > 5 km.
    """
    # Extract QCLOUD and vertical height (z in km)
    qcloud = ds["QCLOUD"]
    
    # Calculate height in kilometers using wrf-python's getvar
    z = ds['z']  # Shape: (bottom_top, south_north, west_east)
    
    # Mask heights where QCLOUD is NaN
    valid_heights = np.where(~np.isnan(qcloud), z, np.nan)  # Retain only valid heights
    
    #print(np.shape(valid_heights))

    # Find the minimum valid height along the vertical (bottom_top axis)
    # Collapse the bottom_top dimension (axis=0) to get a 2D array
    max_height = np.nanmax(valid_heights, axis=0)  # Shape: (south_north, west_east)

    #print(np.shape(min_height))

    # Convert to xarray.DataArray for consistency
    max_height_da = xr.DataArray(
        max_height,
        dims=("south_north", "west_east"),
        coords={
            "XLONG": ds["XLONG"],
            "XLAT": ds["XLAT"]
        },
        name="cloud_base_height_km",
        attrs={"description": "Max height (in km) where QCLOUD is valid"}
    )
    
    #min_height_da.to_dataset(name=cloud_base_height_km)

    return max_height_da

def interpolate_lowest_valid_height(ds):
    """
    Interpolates the lowest height (in km) at each horizontal grid point
    where QCLOUD is not NaN. (to use for cloud base height (CBH))
    
    Parameters:
        ds (xarray.Dataset): The input dataset with QCLOUD and vertical levels.

    Returns:
        xarray.DataArray: 2D array of lowest height (in km) at each grid point.
    """
    # Extract QCLOUD and vertical height (z in km)
    qcloud = ds["QCLOUD"]
    
    # Calculate height in kilometers using wrf-python's getvar
    z = ds['z']  # Shape: (bottom_top, south_north, west_east)
    
    # Mask heights where QCLOUD is NaN
    valid_heights = np.where(~np.isnan(qcloud), z, np.nan)  # Retain only valid heights
    
    #print(np.shape(valid_heights))

    # Find the minimum valid height along the vertical (bottom_top axis)
    # Collapse the bottom_top dimension (axis=0) to get a 2D array
    min_height = np.nanmin(valid_heights, axis=0)  # Shape: (south_north, west_east)

    #print(np.shape(min_height))

    # Convert to xarray.DataArray for consistency
    min_height_da = xr.DataArray(
        min_height,
        dims=("south_north", "west_east"),
        coords={
            "XLONG": ds["XLONG"],
            "XLAT": ds["XLAT"]
        },
        name="cloud_base_height_km",
        attrs={"description": "Lowest height (in km) where QCLOUD is valid"}
    )
    
    #min_height_da.to_dataset(name=cloud_base_height_km)

    return min_height_da

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')
        #print(ncfile) 
        print(file)
        data = getvar(ncfile,var_name)
        #data = data.where(data > 10e-5) # mask to only get values that denote cloud base value
        #data = data.where(data>=0)

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

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

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

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

    return combined_ds

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

    array_list=[]
    for file in file_list:
        ##-- read file            
        ncfile = netCDF4.Dataset(file,'r')
        #print(ncfile) 
        print(file)
        data = getvar(ncfile,var_name)
        data = data.sel(mcape_mcin_lcl_lfc='mcape')
        cloudfrac = getvar(ncfile,'cloudfrac')

        # Mask all values where any level in low_mid_high is > 0
        mask = (cloudfrac > 0).any(dim="low_mid_high")
        cloud_mask = cloudfrac.where(~mask, np.nan).min(dim='low_mid_high') == 0


        # Apply cloud mask to only derive clear sky points
        data = xr.where(cloud_mask, data, float("nan"))

        data = data.to_dataset(name=var_name)
        array_list.append(data)
        ncfile.close()

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

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

    #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:
        ds = read_in_monthly_data2(month, '3hr', sim, 'cape_2d')
        ds.to_netcdf(f'/pscratch/sd/d/dbrooks/thermo_data/{sim}/mu_cape_clearsky_month{month}.nc')
        print(sim)
    print(month)

'''
for month in monlist:
    for sim in sim_list:
        ds = read_in_monthly_data(month, '3hr', sim, 'wa')
        if sim == 'current':
            filename = f'/pscratch/sd/d/dbrooks/vertical_motion/Current/w_component_month{month}.nc'
            #print(ds)
            ds.to_netcdf(filename)
            #ds = xr.open_dataset(filename)
            #current_list.append(ds)
        elif sim == 'future':
            filename = f'/pscratch/sd/d/dbrooks/vertical_motion/Future/w_component_month{month}.nc'
            ds.to_netcdf(filename)
            #ds = xr.open_dataset(filename)
            #future_list.append(ds)
        elif sim == 'future_urban':
            filename = f'/pscratch/sd/d/dbrooks/vertical_motion/Future_urban/w_component_month{month}.nc'
            ds.to_netcdf(filename)
            #ds = xr.open_dataset(filename)
            #future_urban_list.append(ds)
        else:
            print('incorrect input simulation name')
        print(sim)
    print(month)
'''