from netCDF4 import Dataset
import h5py
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 re
import warnings

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

def read_in_monthly_data(month, hour_interval, climate_state, var_name, pressure_level=200, units='m s-1', pressure_var=True):
    """
    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))

    if pressure_var == False:
        array_list=[]
        for file in file_list:
            ##-- read file            
            ncfile = netCDF4.Dataset(file,'r') 
            data = getvar(ncfile,var_name, units=units)
            data = data.to_dataset(name=var_name)
            #print(data.attrs)
            array_list.append(data)
            ncfile.close()

        combined_ds = xr.concat(array_list, dim='Time')
        combined_ds[var_name].attrs['projection'] = str(combined_ds[var_name].attrs['projection'])
        combined_ds.attrs['projection'] = str(combined_ds[var_name].attrs['projection'])

    else:
        p_list=[]
        u_list=[]
        v_list=[]
        wind_list=[]
        for file in file_list:
            ##-- read file            
            ncfile = netCDF4.Dataset(file,'r') 

            # Extract the pressure, geopotential height, and wind variables
            p = getvar(ncfile, "pressure")
            z = getvar(ncfile, "z", units="dm")
            ua = getvar(ncfile, "ua", units="m s-1")
            va = getvar(ncfile, "va", units="m s-1")
            #wspd = getvar(ncfile, "wspd_wdir", units="m s-1")[0,:]

            # Interpolate geopotential height, u, and v winds to 500 hPa
            ht_500 = interplevel(z, p, pressure_level)
            u_500 = interplevel(ua, p, pressure_level)
            v_500 = interplevel(va, p, pressure_level)
            #wspd_500 = interplevel(wspd, p, pressure_level)

            p_list.append(ht_500)
            u_list.append(u_500)
            v_list.append(v_500)
            #wind_list.append(wspd_500)
            print(file)
            ncfile.close()

        print('done1')
        combined_p = xr.concat(p_list, dim='Time')
        combined_u = xr.concat(u_list, dim='Time')
        combined_v = xr.concat(v_list, dim='Time')
        #combined_wind = xr.concat(wind_list, dim='Time')

        combined_p = combined_p.to_dataset()
        combined_u = combined_u.to_dataset()
        combined_v = combined_v.to_dataset()
        #combined_wind = combined_wind.to_dataset()

        #print(combined_p)
        combined_p['height_interp'].attrs['projection'] = str(combined_p['height_interp'].attrs['projection'])
        combined_u['ua_interp'].attrs['projection'] = str(combined_u['ua_interp'].attrs['projection'])
        combined_v['va_interp'].attrs['projection'] = str(combined_v['va_interp'].attrs['projection'])
        #combined_wind['wspd_wdir_interp'].attrs['projection'] = str(combined_wind['wspd_wdir_interp'].attrs['projection'])

    return combined_p, combined_u, combined_v #combined_wind

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

current_list=[]
future_list=[]
future_urban_list=[]

for month in monlist:
    for sim in sim_list:
        combined_p, combined_u, combined_v = read_in_monthly_data(month, '3hr', sim, var_name='pressure', pressure_var=True)
        #print(combined_u)
        #print(combined_v)
        #print(combined_wind)
        if sim == 'current':
            filename1 = f'/pscratch/sd/d/dbrooks/pressure_data/Current/200mb/200mb_month{month}.nc'
            filename2 = f'/pscratch/sd/d/dbrooks/pressure_data/Current/200mb/u_month{month}.nc'
            filename3 = f'/pscratch/sd/d/dbrooks/pressure_data/Current/200mb/v_month{month}.nc'
            filename4 = f'/pscratch/sd/d/dbrooks/pressure_data/Current/200mb/wind_month{month}.nc'
            combined_p.to_netcdf(filename1)
            combined_u.to_netcdf(filename2)
            combined_v.to_netcdf(filename3)

            #ds = xr.open_dataset(filename)
            #current_list.append(ds)
        elif sim == 'future':
            filename1 = f'/pscratch/sd/d/dbrooks/pressure_data/Future/200mb/200mb_month{month}.nc'
            filename2 = f'/pscratch/sd/d/dbrooks/pressure_data/Future/200mb/u_month{month}.nc'
            filename3 = f'/pscratch/sd/d/dbrooks/pressure_data/Future/200mb/v_month{month}.nc'
            filename4 = f'/pscratch/sd/d/dbrooks/pressure_data/Future/200mb/wind_month{month}.nc'
            combined_p.to_netcdf(filename1)
            combined_u.to_netcdf(filename2)
            combined_v.to_netcdf(filename3)

            #ds = xr.open_dataset(filename)
            #future_list.append(ds)
        elif sim == 'future_urban':
            filename1 = f'/pscratch/sd/d/dbrooks/pressure_data/Future_urban/200mb/200mb_month{month}.nc'
            filename2 = f'/pscratch/sd/d/dbrooks/pressure_data/Future_urban/200mb/u_month{month}.nc'
            filename3 = f'/pscratch/sd/d/dbrooks/pressure_data/Future_urban/200mb/v_month{month}.nc'
            filename4 = f'/pscratch/sd/d/dbrooks/pressure_data/Future_urban/200mb/wind_month{month}.nc'
            combined_p.to_netcdf(filename1)
            combined_u.to_netcdf(filename2)
            combined_v.to_netcdf(filename3)

            #ds = xr.open_dataset(filename)
            #future_urban_list.append(ds)
        else:
            print('incorrect input simulation name')
    print(sim)
    print(month)