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

################## Specific Humidity #####################

# Function to calculate actual vapor pressure
def get_vap_pressure(Td):
    """
    Calculates actual vapor pressure (e) from dew point temperature (Td in Kelvin).
    """
    #Td = Td + 273.16 # C to K

    a1 = 6.1121  # hPa
    a3 = 17.502
    a4 = 32.19   # K
    To = 273.16  # K
    t = ((Td - To) / (Td - a4)) * a3
    e  = a1 * np.exp(t)

    return e

def get_specific_humidity(ds_dew_point, ds_pressure):
    """
    Calculates specific humidity (q)
    
    Inputs:
    - ds_dew_point: xarray.DataArray for dew point (Kelvin).
    - ds_pressure: xarray.DataArray for pressure (hPa).

    Outputs:
    - q as an xarray.DataArray (kg/kg).
    """
    E = 0.622  # kg/kg, ratio of gas constants for dry air and water vapor

    # Calculate vapor pressure
    e = get_vap_pressure(ds_dew_point)

    # Calculate q
    q = (E * e) / (ds_pressure - ((1 - E) * e))
    return q

def read_in_monthly_data(month, hour_interval, climate_state):
    """
    Reads in all NetCDF files for a given month, calculates specific humidity at 850mb,
    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 specific humidity at 850mb (g/kg)
    """
    
    # 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}*'
        #file_path = f'/global/cfs/projectdirs/m4486/yuwei/Climate_Impact/long-term/data/{climate_state}/{hour_interval}/wrfout_d01_2017-04-01_00:00:00'
    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 dew point temperature
        pressure = getvar(ncfile, "pressure")  # 3D pressure in hPa
        td = getvar(ncfile, "td", units="K")  # 3D dew point in Kelvin
        
        # Interpolate to 850 hPa level
        td_850 = interplevel(td, pressure, 850.0)
        
        # Create a DataArray for pressure at 850mb (constant value)
        pressure_850 = xr.DataArray(
            np.full_like(td_850.values, 850.0),
            coords=td_850.coords,
            dims=td_850.dims
        )
        
        # Calculate specific humidity at 850mb
        q_850 = get_specific_humidity(td_850, pressure_850)
        q_850.attrs['description'] = 'Specific humidity at 850 hPa'
        q_850.attrs['units'] = 'g/kg'
        q_850.name = 'q_850'
        
        array_list.append(q_850)
        ncfile.close()

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

    return combined_ds

################### Relative Humidity #####################

def read_in_monthly_data_rh(month, hour_interval, climate_state):
    """
    Reads in all NetCDF files for a given month, calculates relative humidity at 850mb,
    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 relative humidity at 850mb (%)
    """
    
    # 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 relative humidity
        pressure = getvar(ncfile, "pressure")  # 3D pressure in hPa
        rh = getvar(ncfile, "rh")  # 3D relative humidity in %
        
        # Interpolate to 850 hPa level
        rh_850 = interplevel(rh, pressure, 850.0)
        
        rh_850.attrs['description'] = 'Relative humidity at 850 hPa'
        rh_850.attrs['units'] = '%'
        rh_850.name = 'rh_850'
        
        array_list.append(rh_850)
        ncfile.close()

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

    return combined_ds

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

################### Temperature #####################

def read_in_monthly_data_temp850(month, hour_interval, climate_state):
    """
    Reads in all NetCDF files for a given month, calculates temperature at 850mb,
    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 temperature at 850mb (C)
    """
    
    # 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 temperature
        pressure = getvar(ncfile, "pressure")  # 3D pressure in hPa
        temp = getvar(ncfile, "temp", units="C")  # 3D temperature in Celsius
        
        # Interpolate to 850 hPa level
        temp_850 = interplevel(temp, pressure, 850.0)
        
        temp_850.attrs['description'] = 'Temperature at 850 hPa'
        temp_850.attrs['units'] = 'C'
        temp_850.name = 'temp_850'
        
        array_list.append(temp_850)
        ncfile.close()

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

    return combined_ds

################### 2m Dewpoint Temperature #####################

def read_in_monthly_data_td2(month, hour_interval, climate_state):
    """
    Reads in all NetCDF files for a given month, extracts 2m dewpoint temperature,
    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 2m dewpoint temperature (C)
    """
    
    # 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))
    print(file_list)

    array_list = []
    for file in file_list:
        ##-- read file            
        ncfile = netCDF4.Dataset(file, 'r')
        print(file)
        
        # Get 2m dewpoint temperature
        td2 = getvar(ncfile, "td2", units="C")  # 2m dewpoint in Celsius
        
        td2.attrs['description'] = '2-meter dewpoint temperature'
        td2.attrs['units'] = 'C'
        td2.name = 'td2'
        
        array_list.append(td2)
        ncfile.close()

    print('done')
    combined_ds = xr.concat(array_list, dim='Time')
    combined_ds = combined_ds.to_dataset()
    combined_ds[td2.name].attrs['projection'] = str(combined_ds[td2.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:
        '''
        # Specific Humidity
        ds_q_850 = read_in_monthly_data(month, '3hr', sim)
        output_path = f'/pscratch/sd/d/dbrooks/acc2017_analysis/humidity_data/{sim}/specific_humidity_850mb_month{month}.nc'
        ds_q_850.to_netcdf(output_path)
        print(f'Saved specific humidity data to {output_path}')
        # Relative Humidity
        ds_rh_850 = read_in_monthly_data_rh(month, '3hr', sim)
        output_path2 = f'/pscratch/sd/d/dbrooks/acc2017_analysis/humidity_data/{sim}/relative_humidity_850mb_month{month}.nc'
        ds_rh_850.to_netcdf(output_path2)
        print(f'Saved relative humidity data to {output_path2}')
        '''
        # Temperature
        #ds_temp_850 = read_in_monthly_data_temp850(month, '3hr', sim)
        #output_path = f'/pscratch/sd/d/dbrooks/acc2017_analysis/temp_data/{sim}/temperature_850mb_month{month}.nc'
        #ds_temp_850.to_netcdf(output_path)
        #print(f'Saved temperature data to {output_path}')

        # Dewpoint Temperature
        ds_td2 = read_in_monthly_data_td2(month, '3hr', sim)
        output_path_td2 = f'/pscratch/sd/d/dbrooks/acc2017_analysis/humidity_data/{sim}/dewpoint_2m_month{month}.nc'
        ds_td2.to_netcdf(output_path_td2)
        print(f'Saved dewpoint temperature data to {output_path_td2}')

        
