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

################### Geopotential Height #####################

def read_in_monthly_data_z500(month, hour_interval, climate_state):
    """
    Reads in all NetCDF files for a given month, calculates geopotential height at 500mb,
    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'

    Returns:
        xarray dataset of full month of geopotential height at 500mb (m)
    """
    
    # 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}*'
    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))

    array_list = []
    for file in file_list:
        ##-- read file            
        ncfile = netCDF4.Dataset(file, 'r')
        print(file)
        
        # Get pressure and geopotential height
        pressure = getvar(ncfile, "pressure")  # 3D pressure in hPa
        z = getvar(ncfile, "z")  # 3D geopotential height in m
        
        # Interpolate to 500 hPa level
        z_500 = interplevel(z, pressure, 500.0)
        
        z_500.attrs['description'] = 'Geopotential height at 500 hPa'
        z_500.attrs['units'] = 'm'
        z_500.name = 'z_500'
        print(z_500)
        
        array_list.append(z_500)
        ncfile.close()

    print('done')
    combined_ds = xr.concat(array_list, dim='Time')
    combined_ds = combined_ds.to_dataset()
    combined_ds[z_500.name].attrs['projection'] = str(combined_ds[z_500.name].attrs['projection'])
    print(combined_ds)

    return combined_ds

# Run the processing for each month and simulation
for month in monlist:
    for sim in sim_list:
        # Geopotential Height
        ds_z_500 = read_in_monthly_data_z500(month, '3hr', sim)
        output_path = f'/pscratch/sd/d/dbrooks/acc2017_analysis/pressure_data/{sim}/500mb/geopotential_height_500mb_month{month}.nc'
        ds_z_500.to_netcdf(output_path)
        print(f'Saved geopotential height data to {output_path}')