#!/usr/bin/env python3
import matplotlib.pyplot as plt
import numpy as np
import xarray as xr
import tobac
import os

def load_palm(fname):
    D = xr.open_mfdataset(fname, parallel=True, engine='h5netcdf', decode_times=False).drop_duplicates(dim='time').dropna('time')
    D['time'] = np.datetime64('2000-01-01T00:00:00') + D.time.astype('timedelta64[s]').data
    D = D.isel(time=(D.time > np.array('2000-01-01T12:00:00.000000000').astype('datetime64[ns]')) & (D.time < np.array('2000-01-01T16:00:00.000000000').astype('datetime64[ns]')))

    dx = dy = np.gradient(D.x).max()
    dt = (D.time.data[1:] - D.time.data[:-1]).mean().astype('timedelta64[s]').astype(np.int32).item()

    if 'lwp*_xy' in D:
        lwp_series = xr.where(D['lwp*_xy'] > 1e-3, D['lwp*_xy'], np.nan).squeeze()
        dz = 1
    else:
        dz = xr.DataArray(np.gradient(D.zw_3d), dims='zu_3d')
        def lwp(ql, eps=1e-3):
            _ = (ql * dz).sum('zu_3d')
            #return _
            return xr.where(_ > eps, _, np.nan)

        lwp_series = lwp(D.ql)

    lwp_series['time'] = lwp_series.time.data.astype('datetime64[s]').astype('datetime64[ns]')
    lwp_series.name = 'lwp'

    return dx, dz, dt, lwp_series.compute()

def load_uclales(fname):
    D = xr.open_dataset(fname, decode_times=False)
    D['time'] = (np.datetime64('2000-01-01T00:00:00') + [ np.timedelta64(int(_),'s') for _ in D.time ]).astype('datetime64[ns]')

    dz = xr.DataArray(np.gradient(D.zm), dims='zt')
    dx = dy = np.gradient(D.xt).max()
    dt = (D.time[1] - D.time[0]).data.astype('timedelta64[s]').astype(np.int32).item()


    def lwp(l, eps=1e-12):
        _ = (l * dz).sum('zt')
        #return _
        return xr.where(_ > eps, _, np.nan)

    time_slice = slice(32, 64, 1)

    lwp_series = lwp(D.l.isel(time=time_slice)).rename(dict(xt='x', yt='y'))
    return dx, dz, dt, lwp_series


def run(fname, load_func, feature_thresh=[.1,], traj_args={}):

    dx, dz, dt, lwp_series = load_func(fname)

    print(dx, dz, dt, lwp_series)
    f = tobac.feature_detection_multithreshold(lwp_series, dx, feature_thresh)
    print(f"features: {f}")
    traj = tobac.linking_trackpy(f, lwp_series, dt=dt, dxy=dx, **traj_args)
    print(f"trajectories: {traj}")


    def plot_features(itime):
        plt.figure()
        plt.clf()
        lwp_series.isel(time=itime).plot()
        _f = f[f.frame==itime]
        plt.plot(_f.x, _f.y, 'x', color='orange')

        traj_in_this_frame = traj[traj.frame==itime]
        for cellid in traj_in_this_frame.cell.unique():
            _t = traj[traj.cell==cellid]
            plt.plot(_t.x, _t.y, '-', markevery=[-1], marker='o',markerfacecolor='r',markeredgecolor='r', markersize=.5)


    plot_features(len(lwp_series.time)-2)


    def mean_path_distances():
        path_distance = []
        for cellid in traj.cell.unique():
            _t = traj[traj.cell==cellid]
            path_dx, path_dy = [_t[_].iloc[-1] - _t[_].iloc[0] for _ in ('x','y')]
            path_distance.append((path_dx, path_dy))
        return np.array(path_distance).mean(axis=0)

    mpd = mean_path_distances()
    print(f"{os.path.basename(fname)}: mean path distance {mpd}")
    return dict(
            dx=dx,
            dz=dz,
            dt=dt,
            lwp=lwp_series,
            features=f,
            traj=traj,
            mpd=mpd,
            )


mpd = dict(
    twostr = run('/project/meteo/work/Richard.Maier/Clean_Installs/PALM/build/JOBS/big_domain_twostr/OUTPUT/big_domain_*_xy.0*.nc', load_palm, feature_thresh=[.05,], traj_args=dict(v_max=10)),
    ts     = run('/project/meteo/work/Richard.Maier/Clean_Installs/PALM/build/JOBS/big_domain_ts/OUTPUT/big_domain_*_xy.0*.nc', load_palm, feature_thresh=[.05,], traj_args=dict(v_max=10)),
#run('/project/meteo/work/Richard.Maier/Clean_Installs/PALM/build/JOBS/big_domain_twostr/OUTPUT/big_domain_*_3d.02*.nc', load_palm, feature_thresh=[.2,], traj_args=dict(v_max=10)),
#run('/project/meteo/work/Richard.Maier/Clean_Installs/PALM/build/JOBS/big_domain_ts/OUTPUT/big_domain_ts_3d.02*.nc', load_palm, feature_thresh=[.2,], traj_args=dict(v_max=10)),
#run('/archive/meteo/work/Fabian.Jakub/ucla_cases/acor/acor_3d180_.5_5_418500/acor_3d180.merged.nc', load_uclales, feature_thresh=[.1,], traj_args=dict(v_max=10)),
#run('/archive/meteo/work/Fabian.Jakub/ucla_cases/acor/acor_3d90_.5_5_418500/acor_3d90.merged.nc', load_uclales, feature_thresh=[.1,], traj_args=dict(v_max=10)),
)
