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 gc
from datetime import datetime
import seaborn as sns
from matplotlib.colors import LinearSegmentedColormap, TwoSlopeNorm
import matplotlib.lines as mlines
import matplotlib.patheffects as pe


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

from joblib import Parallel, delayed
from scipy.stats import mannwhitneyu
from tqdm import tqdm

def calculate_frequency_per_bin(precip_data, bins):
    """Return a list of frequency counts per bin."""
    for lo, hi in bins:
        return [(precip_data >= lo) & (precip_data < hi)].sum()

def bootstrap_chunk(current, future, bins, n_iterations):
    """Run a small batch of bootstrap iterations serially inside one worker."""
    rel_changes = []
    p_vals = []

    for _ in range(n_iterations):
        #print('1a')
        sample_size_cur = min(10000000, len(current))
        sample_size_fut = min(10000000, len(future))

        cur_sample = current[np.random.randint(0, len(current), sample_size_cur)]
        fut_sample = future[np.random.randint(0, len(future), sample_size_fut)]

        #print('1b')

        # Frequencies per bin
        cur_freq = np.array(calculate_frequency_per_bin(cur_sample, bins), dtype=np.float32)
        fut_freq = np.array(calculate_frequency_per_bin(fut_sample, bins), dtype=np.float32)

        with np.errstate(divide='ignore', invalid='ignore'):
            rel = ((fut_freq / cur_freq) - 1) * 100
            rel[cur_freq == 0] = np.nan

        try:
            _, p = mannwhitneyu(cur_sample, fut_sample, alternative='two-sided')
        except ValueError:
            p = np.nan

        rel_changes.append(rel)
        p_vals.append(p)

    return rel_changes, p_vals

def bootstrap_relative_change_chunked(current, future, bins, n_iterations=1000, batch_size=10, n_jobs=-1):
    n_chunks = n_iterations // batch_size
    remaining = n_iterations % batch_size

    print(f"Running {n_iterations} iterations in {n_chunks} chunks of {batch_size} with {n_jobs} workers...")

    # Create task list
    tasks = [batch_size] * n_chunks
    if remaining > 0:
        tasks.append(remaining)

    # Run in parallel over batches
    results = Parallel(n_jobs=n_jobs)(
        delayed(bootstrap_chunk)(current, future, bins, n_iter)
        for n_iter in tqdm(tasks, desc="Bootstrapping (chunked)")
    )

    # Flatten results
    rel_changes = np.concatenate([np.array(r[0]) for r in results], axis=0)
    p_values = np.concatenate([np.array(r[1]) for r in results], axis=0)

    rel_changes = np.stack(rel_changes)  # shape: (n_iterations, n_bins)

    ci_lower = np.nanpercentile(rel_changes, 2.5, axis=0)
    ci_upper = np.nanpercentile(rel_changes, 97.5, axis=0)
    p_sig_fraction = np.mean(p_values < 0.05)

    return ci_lower.tolist(), ci_upper.tolist(), p_sig_fraction


def calculate_frequency_per_bin(precip_data, bins):
    bin_frequencies = []
    for lower, upper in bins:
        # Count occurrences within each bin across all time steps and grid points
        bin_mask = (precip_data >= lower) & (precip_data < upper)
        bin_frequency = bin_mask.sum().item()  # Total occurrences in this bin
        bin_frequencies.append(bin_frequency)
    return bin_frequencies


# Relative change functions
def calculate_relative_change(future, current):
    return ((future / current) - 1) * 100

def calculate_relative_change_urbanization(future, current, future_urban):
    f = ((future / current) - 1) * 100
    fu = ((future_urban / current) - 1) * 100
    return fu - f

# Remove boundaries
def remove_lateral_boundaries(current_ds,future_ds,future_urban_ds):
    # Access the latitude and longitude arrays (XLAT, XLONG)
    lats = current_ds['XLAT']
    lons = current_ds['XLONG']

    # Get the shape of the latitude and longitude arrays
    n_lat, n_lon = lats.shape
    # Exclude 15 grid cells from each side (latitude and longitude)
    lat_slice = slice(15, n_lat - 15)
    lon_slice = slice(15, n_lon - 15)

    # Subset the data using the grid cell indices
    current_ds = current_ds.isel(south_north=lat_slice, west_east=lon_slice)
    future_ds = future_ds.isel(south_north=lat_slice, west_east=lon_slice)
    future_urban_ds = future_urban_ds.isel(south_north=lat_slice, west_east=lon_slice)

    return current_ds,future_ds,future_urban_ds

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

# Precip
for month in monlist:
    for sim in sim_list:
        if sim == 'current':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/precip_data/{sim}/hourly_precip_data_month{month}.nc'
            ds = xr.open_dataset(filename)
            #precip_current_list.append(ds)
            precip_current_ds = ds
        elif sim == 'future':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/precip_data/{sim}/hourly_precip_data_month{month}.nc'
            ds = xr.open_dataset(filename)
            #precip_future_list.append(ds)
            precip_future_ds = ds
        elif sim == 'future_urban':
            filename = f'/pscratch/sd/d/dbrooks/acc2017_analysis/precip_data/{sim}/hourly_precip_data_month{month}.nc'
            ds = xr.open_dataset(filename)
            #precip_future_urban_list.append(ds)
            precip_future_urban_ds = ds
        else:
            print('incorrect input simulation name')
        print(sim)
    print(month)

    # Adjust first time value in hourly precip data
    # Extract time values as a pandas Index
    time_values = precip_current_ds.Time.values

    # Adjust only the first time value by subtracting 1 hour
    time_values[0] = pd.Timestamp(time_values[0]) - pd.Timedelta(hours=1)

    # Reassign the modified time back to the dataset
    precip_current_ds = precip_current_ds.assign_coords(Time=time_values)
    precip_future_ds = precip_future_ds.assign_coords(Time=time_values)
    precip_future_urban_ds = precip_future_urban_ds.assign_coords(Time=time_values)


    precip_current_ds, precip_future_ds, precip_future_urban_ds = remove_lateral_boundaries(precip_current_ds, precip_future_ds, precip_future_urban_ds)

    ########### Bootstrapping Test ##############

    # Sample precipitation bins and labels
    PRECIP_BINS = [(0.25, 2.5), (2.5, 10), (10, 50), (50, np.inf)]
    BIN_LABELS = ['0.25-2.5', '2.5-10', '10-50', '≥50']


    all_precip_current = precip_current_ds['RAINNC']
    all_precip_future = precip_future_ds['RAINNC']
    all_precip_future_urban = precip_future_urban_ds['RAINNC']

    f_array = all_precip_future.values.flatten()
    c_array = all_precip_current.values.flatten()
    fu_array = all_precip_future_urban.values.flatten()


    # Run bootstrapping
    acc_ci_lower, acc_ci_upper, acc_sig = bootstrap_relative_change_chunked(
    c_array, f_array, PRECIP_BINS, n_iterations=1000, batch_size=10, n_jobs=2)

    print('ACC:', acc_ci_lower, acc_ci_upper)

    #urb_ci_lower, urb_ci_upper, urb_sig = bootstrap_relative_change_chunked(
    #f_array, fu_array, PRECIP_BINS, n_iterations=1000, batch_size=10, n_jobs=2)

    #print('Urb:', urb_ci_lower, urb_ci_upper)
    