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

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

def load_in_mask(sim, month):
    filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/masks/{sim}/convective_core_mask_withtimes_month{month}.nc'
    #filename = f'/home/dbrooks/wind_analysis/radar_data/{sim}/composite_ref_month{month}.nc'
    #print(ds)
    ds = xr.open_dataarray(filename)
    #ds = ds['__xarray_dataarray_variable__']
    #print(ds)
    #ds = ds['mdbz'] > 5 
    return ds

def extract_temp_at_cloud_top(cloud_top_height_km, climate_state, month):
    """
    Extract temperature at cloud top height for each time step from WRF outputs.

    Parameters
    ----------
    cloud_top_height_km : xr.DataArray
        3D array (Time, south_north, west_east) of cloud top heights in km.

    Returns
    -------
    xr.DataArray
        Temperatures (degC) at cloud top heights with same shape as cloud_top_height_km.
    """
    file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/3hr/wrfout_d01_2017-{month}*'
    #file_path = f'/pscratch/sd/y/yuwei/Climate_Impact/long-term/data/{climate_state}/3hr/wrfout_d01_2017-04-29_21:00:00'

    
    wrf_files = sorted(glob.glob(file_path))

    times, ny, nx = cloud_top_height_km['cloud_top_height_km'].shape
    temp_data = np.full((times, ny, nx), np.nan, dtype=np.float32)

    mask = load_in_mask(climate_state,month)

    for t_idx, wrf_file in enumerate(wrf_files):
        with Dataset(wrf_file) as nc:
            print(wrf_file)
            temp_3d = getvar(nc, "temp", units="degC")  # (nz, ny, nx)
            time = temp_3d.Time.values
            #print(time, type(time))
            z_3d = getvar(nc, "z", units="km")          # (nz, ny, nx)

            # Cloud top heights for this time slice
            cth_slice = cloud_top_height_km['cloud_top_height_km'].sel(Time=time).data  # (ny, nx)

            # Find nearest vertical index at each (y,x)
            idx_k = np.abs(z_3d.data - cth_slice[None, :, :]).argmin(axis=0)  # (ny, nx)

            # Gather temps at those indices using take_along_axis
            temp_slice = np.take_along_axis(
                temp_3d.data, idx_k[None, :, :], axis=0
            )[0]  # remove vertical axis after gather

            temp_data[t_idx] = temp_slice

    # Build DataArray
    temp_da = xr.DataArray(
        temp_data,
        coords=cloud_top_height_km.coords,
        dims=cloud_top_height_km.dims,
        name="temp_at_cloud_top_degC",
        attrs={"units": "degC", "description": "Temperature at cloud top height"}
    )

    temp_da = xr.where(mask, temp_da, float("nan"))

    return temp_da

q_thresh = ['1e-5','1e-6']
for thresh in q_thresh:
    for month in monlist:
        for sim in sim_list:
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/cloud_stuff/{sim}/cloud_top_height_{thresh}_month{month}.nc'
            print(filename)
            cloud_top_ds = xr.open_dataset(filename)
            cc_mask = load_in_mask(sim,month)
            cloud_top_ds = xr.where(cc_mask, cloud_top_ds, float("nan"))

            print(sim, month)
            temp_da = extract_temp_at_cloud_top(cloud_top_ds, sim, month)
            temp_da.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/cloud_stuff/{sim}/cloud_height_{thresh}_temp_month{month}.nc')