import numpy as np
import os
import matplotlib.pyplot as plt
from kuibit.simdir import SimDir
import kuibit.visualize_matplotlib as viz

home = os.path.expanduser("~/")

datadir = os.path.join(home, "Cactus", "wave2dx")
print(datadir)
sim = SimDir(datadir)
print(sim)

frame_number = 5

xyz_data = sim.gridfunctions.xyz["waveeqn_v"]
iterations = xyz_data.iterations
iteration = iterations[frame_number]
x0 = xyz_data[iteration].x0
x1 = xyz_data[iteration].x1
print(x0, x1)
print(len(iterations))

times = xyz_data.times
resolution = [3, 100, 100]

# Create x and y coord objects for use in pcolor
def get_coords(xxx, x0, x1):
    xvals = np.linspace(x0[2], x1[2], np.shape(xxx)[1])
    yvals = np.linspace(x0[1], x1[1], np.shape(xxx)[0])
    xcoord, ycoord = np.meshgrid(xvals, yvals)
    return xcoord, ycoord

### Using interpolation

def get_slice(iteration):
    xyz_data_unif = xyz_data[iteration].to_UniformGridData(shape=resolution,
                                        x0=x0,
                                        x1=x1,
                                        resample=True)
    # Coordintes are z, y, x
    return xyz_data_unif.data[1,:,:]

slice = get_slice(iteration)

xcoord, ycoord = get_coords(slice, x0, x1)

print(slice.min(), slice.max())
plt.pcolormesh(xcoord, ycoord, slice)
plt.savefig("kplot.png")

### Using raw data

def plot_frame(frame,vmin,vmax,z=0,where=plt):
    """
    We will plot a z=constant slice of a 3D data set
    """
    iteration = iterations[frame]
    data = xyz_data[iteration]
    grids = reversed(list(data.iter_from_finest()))
    for grid in grids:
        # the grid is a tuple of 3 things. It's the last one we care about
        grid_data = grid[2]
        x0 = grid_data.x0
        x1 = grid_data.x1

        # Note that OpenPMD does zyx
        if z < x0[0] or z > x1[0]:
            continue

        # Pick a different color scheme for each refinement level
        cmlist = [plt.cm.Reds, plt.cm.Oranges, plt.cm.Greens, plt.cm.Blues, plt.cm.Purples]
        cm = cmlist[grid_data.ref_level % len(cmlist)]

        # Extract the numpy grid
        np_grid = grid_data.data

        # Find the z=constant slice
        if z == x0[0]:
            np_z = 0
        elif z == x1[0]:
            np_z = np_grid.shape[0]-1
        else:
            dz = (x1[0] - x0[0])/(np_grid.shape[0]-1)
            np_z = int(round((z - x0[0])/dz))

        # Extract the z=constant slice
        xy_np_data = np_grid[np_z,:,:]

        xcoord, ycoord = get_coords(xy_np_data, x0, x1)
        where.pcolor(xcoord, ycoord, xy_np_data, cmap=cm, vmin=vmin, vmax=vmax)

plot_frame(frame_number, -1, 1)
plt.savefig("kplot2.png")
