# -*- coding: utf-8 -*-
"""
Created on Tue Mar 18 21:36:35 2025

@author: Travis Hahn
"""

# ## Load in Data
import pandas as pd
import pickle
import json

path_to_files = "D:\\Research\\BNL\\Data\\wrfout\\wrfout\\2022-06-20\\Tracking_Data"

wrf_mergers = pd.read_csv(f"{path_to_files}/merge_data.csv")
wrf_mergers = wrf_mergers.drop(wrf_mergers.columns[[0]],axis=1)
wrf_splitters = pd.read_csv(f"{path_to_files}/split_data.csv")
wrf_splitters = wrf_splitters.drop(wrf_splitters.columns[[0]],axis=1)


WRF_tracks = pd.read_hdf(f"{path_to_files}\\Tracking.h5","table")

import xarray as xr

WRF_segmentation_2d = xr.open_dataset("D:\\Research\\BNL\\Data\\wrfout\\wrfout\\2022-06-20\\Tracking_Data\\Mask_Segmentation_TWC.nc")

from copy import deepcopy
feature_segmentation = (WRF_segmentation_2d.segmentation_mask - 0).rename("Feature_Segmentation")
cell_segmentation = deepcopy(feature_segmentation).rename("Cell_Segmentation")

frame_groups = WRF_tracks.groupby("frame")

# Loop over tracks, replacing feature_id values with cell_id values in the cell_segmenation DataArray
for frame in frame_groups:
    cell_segmentation_frame = cell_segmentation[frame[0]].values
    # Loop over each feature in that frame
    for feature in frame[1].itertuples():
        # Replace the feature_id with the cell_id
        cell_segmentation_frame[
            cell_segmentation_frame == feature.feature
        ] = feature.cell

print(wrf_mergers)
#%%

# Find the mergers or splitters
from copy import deepcopy

split_or_merge = deepcopy(wrf_splitters) # wrf_mergers
mergers = False
splitters = not(mergers)

for i in split_or_merge.itertuples():
    print(i)
    
    index = i[0]

    frames = i[1].replace("(", "").replace(")", "").split(", ")
    frames = [int(f) for f in frames if f != ""]
    first_frame, last_frame = int(frames[0]), int(frames[1])

    if mergers:
        parents = i[2].replace("(", "").replace(")", "").split(" ")
        parents = [int(p) for p in parents if p != ""]
        first_parent, last_parent = int(parents[0]), int(parents[1])

        merged_cell = i[3]
    
        # Both cells must live longer than 5 frames
        lifetime = 5    
        if len(WRF_tracks[WRF_tracks["cell"] == first_parent]) > lifetime:
            if len(WRF_tracks[WRF_tracks["cell"] == last_parent]) > lifetime:
                print(split_or_merge.iloc[index])
    
    if splitters:
        parent = i[2]

        split_cells = i[3].replace("(", "").replace(")", "").split(", ")
        split_cells = [int(p) for p in split_cells if p != ""]
        first_split, last_split = int(split_cells[0]), int(split_cells[1])

        # Both cells must live longer than 5 frames
        lifetime = 5        
        if len(WRF_tracks[WRF_tracks["cell"] == first_split]) > lifetime:
            if len(WRF_tracks[WRF_tracks["cell"] == last_split]) > lifetime:
                print(split_or_merge.iloc[index])
                
#%%
from shapely import geometry
from osgeo import gdal
import rasterio.features
import geopandas as gpd
import numpy as np
import datetime
import matplotlib.ticker as mticker
import xarray as xr
import matplotlib.pyplot as plt
from copy import deepcopy
import matplotlib as mpl

def get_gdf(segmentation, t = int, cell: list[int, int] | None = None):

    if cell is None:
        cell_seg = deepcopy(segmentation.values[t])
    else:
        cell_seg_full = deepcopy(segmentation.values[t])
        cell_seg = -np.ones_like(cell_seg_full)
        if np.sum(cell_seg_full == cell) == 0:
            print(f"None of the requested cells found at this altitude.\nThe available cells are {np.unique(cell_seg_full)}")
            return None

        cell_seg = np.where(cell_seg_full == cell, cell, cell_seg)
            
    lat_arr = segmentation.latitude.values
    lon_arr = segmentation.longitude.values

    myShapes = rasterio.features.shapes(cell_seg, connectivity=8)
    x_y_coords_in_latlon = []
    values = []
    for shape in myShapes:
        x_y_coords_in_latlon.append(shape[0]["coordinates"][0])
        values.append(int(shape[1]))

    polygons = []
    cell_list = []
    for ind, row in zip(values, x_y_coords_in_latlon):
        lat_list = []
        lon_list = []
        row_list = []
        skip = False
        for x, y in row:
            x_set = x - 1
            y_set = y - 1

            lat = lat_arr[int(y_set), int(x_set)]
            lon = lon_arr[int(y_set), int(x_set)]

            lat_list.append(lat)
            lon_list.append(lon)

            latlon_coord = (lat, lon)

            row_list.append(latlon_coord)

        if not(skip):
            cell_list.append(ind)
            polygons.append(geometry.Polygon(list(zip(lon_list, lat_list))))
    
    crs = "EPSG:4326"

    cell_df = pd.DataFrame({"cell" : cell_list})
    my_gdf2 = gpd.GeoDataFrame(cell_df, crs=crs, geometry=polygons)
    return my_gdf2


def plot_tracked_cells(
        segmentation : xr.Dataset,
        tracks : gpd.GeoDataFrame,
        merge_frames: list[int, int],
        dataset_name: str = "WRF",
        cell: list[int, int] | None = None,
        bounds: list[tuple, tuple] | None = None,
        save = None,
        **args,
        ):

    plt.figure(figsize=[9, 9])
    ax = plt.subplot(111)

    # url = "/share/D3/data/hweiner/ne_10m_admin_1_states_provinces.zip"
    # world = gpd.read_file(url)
    # world.plot(ax=ax, edgecolor='black', facecolor='none')

    total_time = segmentation.shape[0]

    # there has got to be a better way to do this
    try:
        lifetimes = np.unique(tracks["time"])
        full_datetime_lifetimes = [datetime.datetime.strptime(lt, "%Y-%m-%d %H:%M:%S") for lt in lifetimes]
        datetime_lifetimes_str = [datetime.datetime.strftime(lt, "%H:%M:%S") for lt in full_datetime_lifetimes]
        datetime_lifetimes = [datetime.datetime.strptime(lt, "%H:%M:%S") for lt in datetime_lifetimes_str]

    except:
        lifetimes = np.unique(tracks["time"]).astype('datetime64[s]').astype(str)
        full_datetime_lifetimes = [datetime.datetime.strptime(lt, "%Y-%m-%dT%H:%M:%S") for lt in lifetimes]
        datetime_lifetimes_str = [datetime.datetime.strftime(lt, "%H:%M:%S") for lt in full_datetime_lifetimes]
        datetime_lifetimes = [datetime.datetime.strptime(lt, "%H:%M:%S") for lt in datetime_lifetimes_str]

    minute_diff_from_init_lifetimes = [(n - datetime_lifetimes[0]).total_seconds() / 60 for n in datetime_lifetimes]

    for cind, c in enumerate(cell):
        # call the first geodataframe, then walk through time according to dt, make a new gdf and join it to the accumulating gdf
        if "t1_c1_cut" in args and cind == 1:
            c1_end = args["t1_c1_cut"]
            t1 = np.max(tracks[tracks["cell"] == c]["frame"]) - c1_end
        else:
            t1 = np.max(tracks[tracks["cell"] == c]["frame"])

        t0 = np.min(tracks[tracks["cell"] == c]["frame"])

        alpha_index = np.linspace(1, 0.1, t1 - t0 + 1)


        gdf_to_plot = get_gdf(segmentation=segmentation, t=t0, cell = c)
        print(alpha_index)
        gdf_to_plot["lifetime"] = tracks.query("frame==@t0").time.values.min() # lifetimes[t0]
        gdf_to_plot["minutes_from_init"] = 0#minute_diff_from_init_lifetimes[t0]
        gdf_to_plot["alpha"] = alpha_index[0]

        for t_step, gdf_t in enumerate(np.arange(t0 + 1, t1 + 1, 1)):
            next_gdf = get_gdf(segmentation=segmentation, t=gdf_t, cell = c)
            if next_gdf is not None:
                next_gdf["lifetime"] = tracks.query("frame==@gdf_t").time.values.min()
                next_gdf["minutes_from_init"] = 0#minute_diff_from_init_lifetimes[gdf_t]
                next_gdf["alpha"] = alpha_index[t_step+1]
                gdf_to_plot = pd.concat((gdf_to_plot, next_gdf))
            else:
                print(f'the gdf at time {gdf_t} is None')
                continue


        cell_lifetimes = np.array(gdf_to_plot["lifetime"])
        cell_minute_diff_from_init = np.array(gdf_to_plot["minutes_from_init"])
        legend_tick_spacer = 1#int(np.floor((t1 - t0) / 2))

        # legend_tick_spacer = 1
        if cind == 1:
            cell_lifetimes_for_legend = [0 for i in cell_lifetimes]
            cell_lifetimes_for_legend = [0 for i in cell_lifetimes_for_legend]
            gdf_to_plot.plot(ax=ax, edgecolor='black', alpha= gdf_to_plot["alpha"],
                    column="minutes_from_init", legend=False, cmap="Reds_r",
                    legend_kwds={"format": mticker.FixedFormatter(cell_lifetimes_for_legend[::legend_tick_spacer]), 
                                    "orientation": "horizontal", 
                                    "shrink": 1,
                                    "aspect": 80,
                                    "pad": 0.09, 
                                    "extend": "both", 
                                    "ticks" : cell_minute_diff_from_init[::legend_tick_spacer]})

            merge_lats  = []
            merge_lons = []

            tracks_cell_1 = tracks[tracks["cell"] == cell[-1]]
            tracks_cell_1_merge_frame = tracks_cell_1[tracks["frame"] == merge_frames[-1]]

            tracks_cell_2 = tracks[tracks["cell"] == cell[0]]
            tracks_cell_2_merge_frame = tracks_cell_2[tracks["frame"] == merge_frames[0]]

            print(cell)
            print(merge_frames[0])

            merge_lats.append(tracks_cell_2_merge_frame["latitude"].values[0])
            merge_lats.append(tracks_cell_1_merge_frame["latitude"].values[0])

            merge_lons.append(tracks_cell_2_merge_frame["longitude"].values[0])
            merge_lons.append(tracks_cell_1_merge_frame["longitude"].values[0])

            if "label" in args:
                l = args["label"]
            else:
                l = "merge event"
            plt.plot(merge_lons, merge_lats, linestyle = '--', color='xkcd:electric green', alpha=0.7, label=l)

        else:
            cell_lifetimes_for_legend = [0 for i in cell_lifetimes]
            cell_lifetimes_for_legend = [0 for i in cell_lifetimes_for_legend]
            gdf_to_plot.plot(ax=ax, edgecolor='black', alpha= gdf_to_plot["alpha"],
                        column="minutes_from_init", legend=False, cmap="Blues_r",
                        legend_kwds={"format": mticker.FixedFormatter(cell_lifetimes_for_legend[::2]), 
                                    "orientation": "horizontal", 
                                    "shrink": 1,
                                    "aspect": 80,
                                    "pad": -0.09, 
                                    "extend": "both", 
                                    "ticks" : cell_minute_diff_from_init[::2]})

            if bounds is None:
            # derive the minimum and maximum plot bounds
                minx = np.min(gdf_to_plot['geometry'].bounds['minx'])
                miny = np.min(gdf_to_plot['geometry'].bounds['miny'])
                maxx = np.max(gdf_to_plot['geometry'].bounds['maxx'])
                maxy = np.max(gdf_to_plot['geometry'].bounds['maxy'])

                ax.set_xlim(minx * 1.001, maxx * 0.999)
                ax.set_ylim(miny * 1.001, maxy * 0.999)
            
            else:
                ax.set_xlim(*bounds[0])
                ax.set_ylim(*bounds[1])

            # append degrees to the x and y labels
            xtl = ax.get_xticklabels()
            ytl = ax.get_yticklabels()
            new_xtl = []
            new_ytl = []
            for ctl_x, ctl_y in zip(xtl, ytl):
                tx = ctl_x.get_text()

                try: # mpl uses weird negative signs, so just use a try/except
                    float(tx) < 0
                    new_xtl.append(fr'{tx}$\degree$E')
                except:
                    new_xtl.append(fr'{tx[1:]}$\degree$W')

                ty = ctl_y.get_text()
                try:
                    float(ty) < 0
                    new_ytl.append(fr'{ty}$\degree$N')
                except:
                    new_ytl.append(fr'{ty[1:]}$\degree$S')


            # write x and y labels
            #ax.set_xticks(np.linspace(-96.8,-95.7))
            #ax.set_xticklabels(new_xtl, rotation=45, fontsize=10)
            #ax.set_yticklabels(new_ytl, rotation=45, fontsize=10)
            ax.set_xlabel("Longitude")
            ax.set_ylabel("Latitude")
        
        if "t1_c1_cut" in args and cind == 1:
            cell_tracks = tracks[tracks["cell"] == c]
            cell_tracks = cell_tracks[:-args["t1_c1_cut"]]

        elif "t1_c0_cut" in args and cind == 0:
            cell_tracks = tracks[tracks["cell"] == c][:-args["t1_c0_cut"]]

        else:
            cell_tracks = tracks[tracks["cell"] == c]

        lats = cell_tracks['latitude'].values
        lons = cell_tracks['longitude'].values
        total_frames = len(cell_tracks["frame"])

        if cind == 1:
            colors_in_frame = plt.cm.Reds_r(np.linspace(0,1,total_frames))
        else:
            colors_in_frame = plt.cm.Blues_r(np.linspace(0,1,total_frames))


        for f in range(total_frames):
            plt.scatter(lons[f], lats[f], marker='o', color = colors_in_frame[f])
        
            # plt.plot(lons[f], lats[f], linestyle='--', alpha=1, color = colors_in_frame[f])
        # colored_line(lons, lats, ax=ax, c = np.arange(total_frames), cmap = "Greens_r" if cind == 1 else "Blues_r")
        else:
            plt.plot(lons, lats, linestyle='--', alpha = 0.6, color = "red" if cind == 1 else "blue")

    # ax.set_title(f"{dataset_name} Merging Example", fontsize=12)

    plt.legend(loc='upper left')
    plt.grid(alpha=0.4)
    plt.tight_layout()

    if save != None:
        plt.savefig(save, dpi=600)

    plt.show()
    return gdf_to_plot

#%%

SMALL_SIZE = 22
MEDIUM_SIZE = 22
BIGGER_SIZE = 23

plt.rc('font', size=SMALL_SIZE)          # controls default text sizes
plt.rc('axes', titlesize=BIGGER_SIZE)     # fontsize of the axes title
plt.rc('axes', labelsize=MEDIUM_SIZE)    # fontsize of the x and y labels
plt.rc('xtick', labelsize=SMALL_SIZE)    # fontsize of the tick labels
plt.rc('ytick', labelsize=SMALL_SIZE)    # fontsize of the tick labels
plt.rc('legend', fontsize=SMALL_SIZE)    # legend fontsize
plt.rc('figure', titlesize=BIGGER_SIZE)  # fontsize of the figure title

segmentation = cell_segmentation
tracks = WRF_tracks
bounds = [(-96.15, -95.65), (29.6, 30)]
a = plot_tracked_cells(segmentation=cell_segmentation, tracks=tracks, dataset_name = "WRF", bounds = bounds,
                        # cell = [ 2032, 376], merge_frames = [4, 5])
                        # cell = [1678, 3973], merge_frames = [5, 6]) # bounds = [(-64.9, -64.55), (-9.76, -9.45)]
                        cell = [110, 131], merge_frames = [120, 121], t1_c1_cut = 1, save = "../plots/merge_example_good.png") # bounds = [(-53.7, -52.7), (-8.2, -7.89)]