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)

########## Downloading metling level ############
def cloud_top_height(ds):
    """
    Determines the maximum cloud height.
    
    Parameters:
        ds (xarray.Dataset): The input dataset with QCLOUD and vertical levels.

    Returns:
        xarray.DataArray: 2D mask where True indicates grid points with max cloud height in 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 max 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_top_height_km",
        attrs={"description": "Max height (in km) where QTOTAL is >1e-5 kg/kg"}
    )
    
    #min_height_da.to_dataset(name=cloud_base_height_km)

    return max_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}*'
        #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))

    array_list=[]
    for file in file_list:
        ##-- read file  
        print(file)          
        ncfile = netCDF4.Dataset(file,'r')
        #print(ncfile) 
        data = getvar(ncfile,var_name) + getvar(ncfile,'QRAIN') + getvar(ncfile,'QICE') + getvar(ncfile,'QSNOW') + getvar(ncfile,'QGRAUP')
        data = data.where(data > 1e-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 = cloud_top_height(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()

    #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_data(month, '3hr', sim, 'QCLOUD')
        if sim == 'current':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/cloud_stuff/current/cloud_top_height_1e-5_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/acc2017_analysis/cloud_stuff/future/cloud_top_height_1e-5_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/acc2017_analysis/cloud_stuff/future_urban/cloud_top_height_1e-5_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)
