#!/usr/bin/env python
"""
waterfall visualisation pipeline.

Reads a directory of solver frames (.npz) and produces PNG frames plus an MP4.

    python visualize_waterfall.py output/tesla --layout dashboard --fps 30

Design notes
------------
Colour limits are computed once, globally, by scanning every frame before
rendering any of them. Per-frame autoscaling is the single most common way to
make a CFD animation lie: the colour bar silently rescales each frame, so a
wake that is decaying looks steady and a growing separation looks constant.
The scan uses robust percentiles rather than min/max so one cell of numerical
noise cannot flatten the whole scale.

Colour maps are perceptually uniform. Signed fields (pressure coefficient,
vorticity) get a diverging map centred exactly on zero, so the sign of a
structure is readable from colour alone; unsigned fields (velocity magnitude)
get a sequential map. The old inferno/jet-style ramps are avoided because
their non-uniform lightness invents banding that reads as flow structure.

3D iso-surfaces come from marching cubes on the cell-centred field. For a
wake, Q-criterion is the right variable -- iso-surfaces of velocity magnitude
or pressure in a wake are dominated by the shear layer and hide the vortex
cores that actually matter.
"""

import argparse
import json
import os
import sys

import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from matplotlib.colors import Normalize, TwoSlopeNorm

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from waterfall_core.io import list_frames, load_frame

# Signed fields need a diverging map centred on zero; unsigned need sequential.
SIGNED_FIELDS = {"p", "cp", "omega_x", "omega_y", "omega_z", "vorticity", "v", "w"}

FIELD_LABELS = {
    "u": "streamwise velocity  u/U",
    "v": "vertical velocity  v/U",
    "w": "spanwise velocity  w/U",
    "p": "pressure coefficient  Cp",
    "cp": "pressure coefficient  Cp",
    "umag": "velocity magnitude  |U|/U",
    "vorticity": "spanwise vorticity  omega_z D/U",
    "omega_z": "spanwise vorticity  omega_z D/U",
    "q": "Q criterion",
}


def _pick_diverging():
    """Prefer a perceptually uniform diverging map when matplotlib has one."""
    for name in ("berlin", "managua", "vanimo", "coolwarm"):
        try:
            plt.get_cmap(name)
            return name
        except (ValueError, KeyError):
            continue
    return "coolwarm"


DIVERGING = _pick_diverging()


# ----------------------------------------------------------------------
# field extraction
# ----------------------------------------------------------------------

def derive(arrays, meta, field):
    """Compute a named field, normalised to freestream where meaningful."""
    u_inf = float(np.linalg.norm(meta["u_inf"])) or 1.0
    rho = meta.get("rho", 1.0)

    if field == "umag":
        return np.sqrt(
            arrays["u"] ** 2 + arrays["v"] ** 2 + arrays["w"] ** 2
        ) / u_inf
    if field in ("cp", "p"):
        # Cp = (p - p_ref) / (0.5 rho U^2). The solver's pressure is
        # zero-mean by construction (the Poisson nullspace is fixed that
        # way), so the far-field value is used as the reference rather than
        # assuming zero.
        p = arrays["p"]
        p_ref = float(np.median(p[0, :, :]))
        return (p - p_ref) / (0.5 * rho * u_inf**2)
    if field == "vorticity":
        key = "omega_z" if "omega_z" in arrays else None
        if key is None:
            raise KeyError("frame has no vorticity; re-run with --save-vorticity")
        d = meta["dx"] * 0.0 + 1.0
        return arrays[key] * d / u_inf
    if field in arrays:
        a = arrays[field]
        if field in ("u", "v", "w"):
            return a / u_inf
        return a
    raise KeyError(f"unknown field {field!r}; frame has {sorted(arrays)}")


# Physically meaningful default range for pressure coefficient. Cp = 1 at a
# stagnation point and rarely drops below about -2 in external aerodynamics.
# Deriving the range from the data instead lets immersed-boundary spikes set it
# -- on the Model S case an auto-scaled Cp ran to +-7 and rendered the entire
# field flat black, hiding the actual pressure distribution completely.
CP_DEFAULT_RANGE = (-2.0, 1.0)


def _fluid_mask(arrays, grow=2):
    """
    True where the field is physically meaningful.

    Excludes the solid *and a few cells around it*. Masking only eta > 0.5 is
    not enough: direct forcing leaves large spurious gradients in the cells
    immediately outside the body, so pressure and Q there run orders of
    magnitude above anything in the flow. Those cells dominate any percentile
    taken over the domain, which is how a colour scale ends up ten times too
    wide and an iso-level ends up above every real structure.
    """
    eta = arrays.get("eta")
    if eta is None:
        return None
    solid = eta > 0.05
    if grow > 0 and solid.any():
        from scipy import ndimage
        solid = ndimage.binary_dilation(solid, iterations=grow)
    return ~solid


def scan_limits(frames, field, percentile=99.0, max_scan=40):
    """
    Global colour limits from a subsample of frames.

    Percentiles rather than min/max, over fluid cells only.
    """
    if field in ("cp", "p"):
        return CP_DEFAULT_RANGE

    idx = np.linspace(0, len(frames) - 1, min(max_scan, len(frames))).astype(int)
    lo, hi = [], []
    for i in idx:
        arrays, meta = load_frame(frames[i])
        d = derive(arrays, meta, field).astype(float)
        mask = _fluid_mask(arrays)
        if mask is not None:
            d = np.where(mask, d, np.nan)
        lo.append(np.nanpercentile(d, 100 - percentile))
        hi.append(np.nanpercentile(d, percentile))
    lo, hi = float(np.min(lo)), float(np.max(hi))
    if field in SIGNED_FIELDS:
        m = max(abs(lo), abs(hi))
        return -m, m
    return lo, hi


def make_norm(field, vmin, vmax):
    if field in ("cp", "p"):
        # Asymmetric on purpose: Cp is bounded above by 1 but can go well
        # below -1, so forcing a symmetric scale wastes most of the colour
        # range on values that cannot occur.
        return TwoSlopeNorm(vmin=vmin, vcenter=0.0, vmax=vmax), DIVERGING
    if field in SIGNED_FIELDS:
        m = max(abs(vmin), abs(vmax)) or 1.0
        return TwoSlopeNorm(vmin=-m, vcenter=0.0, vmax=m), DIVERGING
    return Normalize(vmin=vmin, vmax=vmax), "viridis"


# ----------------------------------------------------------------------
# panels
# ----------------------------------------------------------------------

def _axes_extent(meta, plane):
    ox, oy, oz = meta["origin"]
    nx, ny, nz = meta["nx"], meta["ny"], meta["nz"]
    dx, dy, dz = meta["dx"], meta["dy"], meta["dz"]
    if plane == "xy":
        return [ox, ox + nx * dx, oy, oy + ny * dy], "x [m]", "y [m]"
    if plane == "xz":
        return [ox, ox + nx * dx, oz, oz + nz * dz], "x [m]", "z [m]"
    return [oy, oy + ny * dy, oz, oz + nz * dz], "y [m]", "z [m]"


def _slice_data(vol, plane, index):
    if plane == "xy":
        return vol[:, :, index].T
    if plane == "xz":
        return vol[:, index, :].T
    return vol[index, :, :].T


def panel_slice(ax, arrays, meta, field, plane, index, norm, cmap,
                show_body=True):
    vol = derive(arrays, meta, field)
    img = _slice_data(vol, plane, index)
    extent, xl, yl = _axes_extent(meta, plane)

    im = ax.imshow(img, origin="lower", extent=extent, cmap=cmap, norm=norm,
                   interpolation="bilinear", aspect="equal")
    if show_body and "eta" in arrays:
        body = _slice_data(arrays["eta"], plane, index)
        # Outline rather than fill: filling hides the near-wall field, which
        # is where separation is decided.
        ax.contour(body, levels=[0.5], extent=extent, colors="k",
                   linewidths=1.0)
    ax.set_xlabel(xl)
    ax.set_ylabel(yl)
    return im


def panel_streamlines(ax, arrays, meta, plane, index, norm, cmap,
                      density=1.6, show_body=True):
    u_inf = float(np.linalg.norm(meta["u_inf"])) or 1.0
    if plane == "xy":
        a, b = arrays["u"][:, :, index].T, arrays["v"][:, :, index].T
    elif plane == "xz":
        a, b = arrays["u"][:, index, :].T, arrays["w"][:, index, :].T
    else:
        a, b = arrays["v"][index, :, :].T, arrays["w"][index, :, :].T
    extent, xl, yl = _axes_extent(meta, plane)
    xs = np.linspace(extent[0], extent[1], a.shape[1])
    ys = np.linspace(extent[2], extent[3], a.shape[0])

    speed = np.sqrt(a**2 + b**2) / u_inf
    ax.streamplot(xs, ys, a, b, color=speed, cmap=cmap, norm=norm,
                  density=density, linewidth=0.8, arrowsize=0.6)
    if show_body and "eta" in arrays:
        body = _slice_data(arrays["eta"], plane, index)
        ax.contourf(body, levels=[0.5, 1.1], extent=extent, colors="0.15")
    ax.set_xlim(extent[0], extent[1])
    ax.set_ylim(extent[2], extent[3])
    ax.set_aspect("equal")
    ax.set_xlabel(xl)
    ax.set_ylabel(yl)


def auto_iso_level(arrays, meta, field, percentile=99.0):
    """
    Pick a Q iso-level that isolates vortex cores.

    Taking a percentile over the whole field does not work: Q is essentially
    zero through most of the domain and large only in thin shear layers and
    cores, so even the 99th percentile lands in the numerical noise floor.
    Contouring that produces sheets pinned to the far-field boundaries and a
    haze of speckle through the free stream -- structures that look like
    physics and are not.

    Restricting to strictly positive Q (rotation dominating strain, which is
    the definition of a vortical region) and going far out into the tail gives
    a level that tracks the actual cores.
    """
    v = derive(arrays, meta, field).astype(float)
    mask = _fluid_mask(arrays)
    if mask is not None:
        v = np.where(mask, v, np.nan)
    v = v[np.isfinite(v)]
    pos = v[v > 0]
    if pos.size < 64:
        return float(np.nanmax(v)) * 0.5 if v.size else 0.0
    return float(np.percentile(pos, percentile))


def panel_body(ax, arrays, meta, stride=1, color="0.45", max_faces=60000):
    """Draw the immersed body as an opaque surface in the 3D panel.

    Without it a wake iso-surface floats in empty space and there is no way to
    tell which part of the geometry shed it."""
    from skimage import measure
    from mpl_toolkits.mplot3d.art3d import Poly3DCollection

    eta = arrays.get("eta")
    if eta is None or eta.shape[2] < 2 or not (eta > 0.5).any():
        return
    if stride > 1:
        eta = eta[::stride, ::stride, ::stride]
    spacing = (meta["dx"] * stride, meta["dy"] * stride, meta["dz"] * stride)
    try:
        verts, faces, _n, _v = measure.marching_cubes(
            eta, level=0.5, spacing=spacing
        )
    except (ValueError, RuntimeError):
        return
    verts += np.array(meta["origin"])
    if len(faces) > max_faces:
        faces = faces[np.linspace(0, len(faces) - 1, max_faces).astype(int)]
    coll = Poly3DCollection(verts[faces], facecolors=color, linewidths=0)
    coll.set_edgecolor(None)
    ax.add_collection3d(coll)


def panel_iso(ax, arrays, meta, field, level, view, cmap, norm,
              color_by="u", stride=1, max_faces=120000, show_body=True):
    """
    3D iso-surface via marching cubes.

    Faces are decimated to max_faces before rendering. Matplotlib's 3D
    renderer sorts polygons in Python, so it degrades badly past roughly
    10^5 faces -- without a cap a single frame can take minutes and the
    video never finishes.
    """
    from skimage import measure

    vol = derive(arrays, meta, field)
    if vol.shape[2] < 2:
        ax.text(0.5, 0.5, 0.5, "2D run:\nno iso-surface", ha="center")
        return

    if stride > 1:
        vol = vol[::stride, ::stride, ::stride]

    finite = np.isfinite(vol)
    if not finite.any() or level <= np.nanmin(vol) or level >= np.nanmax(vol):
        ax.text2D(0.5, 0.5, f"iso level {level:g}\noutside data range",
                  transform=ax.transAxes, ha="center")
        return

    spacing = (
        meta["dx"] * stride, meta["dy"] * stride, meta["dz"] * stride,
    )
    verts, faces, _n, _v = measure.marching_cubes(
        np.nan_to_num(vol, nan=level - 1.0), level=level, spacing=spacing
    )
    verts += np.array(meta["origin"])

    if len(faces) > max_faces:
        keep = np.linspace(0, len(faces) - 1, max_faces).astype(int)
        faces = faces[keep]

    # Colour the surface by a second field sampled at each triangle centroid,
    # so the iso-surface carries more information than its own level.
    cvol = derive(arrays, meta, color_by)
    if stride > 1:
        cvol = cvol[::stride, ::stride, ::stride]
    centroids = verts[faces].mean(axis=1)
    idx = np.floor(
        (centroids - np.array(meta["origin"])) / np.array(spacing)
    ).astype(int)
    for a in range(3):
        np.clip(idx[:, a], 0, cvol.shape[a] - 1, out=idx[:, a])
    fc = plt.get_cmap(cmap)(norm(cvol[idx[:, 0], idx[:, 1], idx[:, 2]]))

    from mpl_toolkits.mplot3d.art3d import Poly3DCollection
    coll = Poly3DCollection(verts[faces], facecolors=fc, linewidths=0)
    coll.set_edgecolor(None)
    ax.add_collection3d(coll)

    ox, oy, oz = meta["origin"]
    ax.set_xlim(ox, ox + meta["nx"] * meta["dx"])
    ax.set_ylim(oy, oy + meta["ny"] * meta["dy"])
    ax.set_zlim(oz, oz + meta["nz"] * meta["dz"])
    ax.set_box_aspect((
        meta["nx"] * meta["dx"], meta["ny"] * meta["dy"], meta["nz"] * meta["dz"]
    ))
    ax.view_init(elev=view[0], azim=view[1])
    ax.set_xlabel("x [m]")
    ax.set_ylabel("y [m]")
    ax.set_zlabel("z [m]")


# ----------------------------------------------------------------------
# frame composition
# ----------------------------------------------------------------------

def render_frame(path, args, norm, cmap, iso_norm, iso_cmap, out_png):
    arrays, meta = load_frame(path)
    nz = meta["nz"]
    two_d = nz == 1

    kz = nz // 2 if args.slice_index is None else args.slice_index
    ky = meta["ny"] // 2

    if args.layout == "slice" or (two_d and args.layout in ("iso", "dashboard")):
        fig, ax = plt.subplots(figsize=(11, 5.5), constrained_layout=True)
        im = panel_slice(ax, arrays, meta, args.field, "xy", kz, norm, cmap)
        fig.colorbar(im, ax=ax, label=FIELD_LABELS.get(args.field, args.field),
                     shrink=0.85)
        axes_for_title = ax

    elif args.layout == "stream":
        fig, ax = plt.subplots(figsize=(11, 5.5), constrained_layout=True)
        panel_streamlines(ax, arrays, meta, "xy", kz, norm, cmap,
                          density=args.stream_density)
        sm = plt.cm.ScalarMappable(norm=norm, cmap=cmap)
        fig.colorbar(sm, ax=ax, label="velocity magnitude  |U|/U", shrink=0.85)
        axes_for_title = ax

    elif args.layout == "iso":
        fig = plt.figure(figsize=(10, 7), constrained_layout=True)
        ax = fig.add_subplot(111, projection="3d")
        panel_iso(ax, arrays, meta, args.iso_field, args.iso_level,
                  (args.elev, args.azim), iso_cmap, iso_norm,
                  color_by=args.field, stride=args.iso_stride)
        sm = plt.cm.ScalarMappable(norm=iso_norm, cmap=iso_cmap)
        fig.colorbar(sm, ax=ax, label=FIELD_LABELS.get(args.field, args.field),
                     shrink=0.6)
        axes_for_title = ax

    else:  # dashboard
        fig = plt.figure(figsize=(15, 8.5), constrained_layout=True)
        gs = fig.add_gridspec(2, 2, width_ratios=[1.25, 1])
        ax0 = fig.add_subplot(gs[0, 0])
        im = panel_slice(ax0, arrays, meta, args.field, "xy", kz, norm, cmap)
        ax0.set_title(f"{FIELD_LABELS.get(args.field, args.field)}  "
                      f"(z-mid slice)", fontsize=10)
        fig.colorbar(im, ax=ax0, shrink=0.85)

        ax1 = fig.add_subplot(gs[1, 0])
        panel_streamlines(ax1, arrays, meta, "xy", kz, norm, cmap,
                          density=args.stream_density)
        ax1.set_title("streamlines, coloured by speed", fontsize=10)

        ax2 = fig.add_subplot(gs[:, 1], projection="3d")
        panel_iso(ax2, arrays, meta, args.iso_field, args.iso_level,
                  (args.elev, args.azim), iso_cmap, iso_norm,
                  color_by=args.field, stride=args.iso_stride)
        ax2.set_title(f"{args.iso_field} iso-surface at {args.iso_level:g}",
                      fontsize=10)
        axes_for_title = ax0

    fx, fy, fz = meta.get("force", (0, 0, 0))
    fig.suptitle(
        f"waterfall   t = {meta['time']:.4f} s   step {meta['step']}   "
        f"F = ({fx:+.3g}, {fy:+.3g}, {fz:+.3g}) N",
        fontsize=11,
    )
    _ = axes_for_title
    fig.savefig(out_png, dpi=args.dpi)
    plt.close(fig)


def _prepare_frame(img):
    """
    Make a PNG safe for H.264: drop alpha, force even dimensions.

    Both matter. matplotlib writes RGBA, and yuv420p wants three channels.
    More subtly, libx264 with yuv420p requires *even* width and height,
    because chroma is subsampled 2x1 in each direction -- an odd dimension
    makes ffmpeg abort, and because imageio streams frames over a pipe the
    symptom is a bare "broken pipe" with an empty stderr rather than a
    message naming the real problem. Figure size times DPI lands on an odd
    pixel count often enough that this is not an edge case; cropping one row
    or column is invisible and unconditional.
    """
    if img.ndim == 3 and img.shape[2] == 4:
        img = img[:, :, :3]
    h, w = img.shape[:2]
    return img[: h - (h % 2), : w - (w % 2)]


def stitch(png_dir, out_mp4, fps):
    import imageio.v2 as imageio
    import imageio_ffmpeg  # noqa: F401  (ensures a bundled encoder exists)

    names = sorted(n for n in os.listdir(png_dir) if n.endswith(".png"))
    if not names:
        raise SystemExit(f"no PNGs in {png_dir}")
    # macro_block_size=1 stops the encoder rescaling frames to a multiple of
    # 16, which would blur the axis labels and colour-bar ticks.
    with imageio.get_writer(out_mp4, fps=fps, codec="libx264",
                            quality=8, macro_block_size=1,
                            ffmpeg_log_level="error") as w:
        for n in names:
            w.append_data(_prepare_frame(imageio.imread(os.path.join(png_dir, n))))
    return out_mp4


# ----------------------------------------------------------------------

def main():
    p = argparse.ArgumentParser(
        description="Render waterfall solver output to PNG frames and MP4.",
        formatter_class=argparse.ArgumentDefaultsHelpFormatter,
    )
    p.add_argument("frames_dir", help="directory containing frame_*.npz")
    p.add_argument("-o", "--outdir", default=None,
                   help="output directory (default: <frames_dir>/viz)")
    p.add_argument("--layout", default="dashboard",
                   choices=["slice", "stream", "iso", "dashboard"])
    p.add_argument("--field", default="umag",
                   help="scalar to colour by: umag, u, v, w, cp, vorticity, q")
    p.add_argument("--iso-field", default="q",
                   help="field for the 3D iso-surface (q recommended)")
    p.add_argument("--iso-level", type=float, default=None,
                   help="iso level; default is an automatic percentile")
    p.add_argument("--iso-stride", type=int, default=1,
                   help="downsample factor before marching cubes")
    p.add_argument("--slice-index", type=int, default=None,
                   help="z index for 2D panels (default: mid-plane)")
    p.add_argument("--elev", type=float, default=22.0, help="3D elevation, deg")
    p.add_argument("--azim", type=float, default=-60.0, help="3D azimuth, deg")
    p.add_argument("--stream-density", type=float, default=1.6)
    p.add_argument("--fps", type=int, default=30)
    p.add_argument("--dpi", type=int, default=110)
    p.add_argument("--cmap", default=None, help="override the colour map")
    p.add_argument("--vmin", type=float, default=None)
    p.add_argument("--vmax", type=float, default=None)
    p.add_argument("--stride", type=int, default=1, help="use every Nth frame")
    p.add_argument("--no-video", action="store_true")
    args = p.parse_args()

    frames = list_frames(args.frames_dir)[:: args.stride]
    if not frames:
        raise SystemExit(f"no frames found in {args.frames_dir}")

    outdir = args.outdir or os.path.join(args.frames_dir, "viz")
    png_dir = os.path.join(outdir, "png")
    os.makedirs(png_dir, exist_ok=True)

    print(f"{len(frames)} frames -> {outdir}")

    if args.vmin is None or args.vmax is None:
        vmin, vmax = scan_limits(frames, args.field)
        vmin = args.vmin if args.vmin is not None else vmin
        vmax = args.vmax if args.vmax is not None else vmax
    else:
        vmin, vmax = args.vmin, args.vmax
    norm, cmap = make_norm(args.field, vmin, vmax)
    if args.cmap:
        cmap = args.cmap
    print(f"  colour limits [{vmin:.4g}, {vmax:.4g}] with '{cmap}'")

    iso_norm, iso_cmap = norm, cmap
    iso_level = args.iso_level
    if args.layout in ("iso", "dashboard") and iso_level is None:
        arrays, meta = load_frame(frames[len(frames) // 2])
        if meta["nz"] > 1:
            try:
                iso_level = auto_iso_level(arrays, meta, args.iso_field)
            except KeyError:
                iso_level = 0.0
        else:
            iso_level = 0.0
        print(f"  auto iso level ({args.iso_field}) = {iso_level:.4g}")
    args.iso_level = iso_level if iso_level is not None else 0.0

    for n, f in enumerate(frames):
        out_png = os.path.join(png_dir, f"f_{n:05d}.png")
        render_frame(f, args, norm, cmap, iso_norm, iso_cmap, out_png)
        if (n + 1) % 10 == 0 or n == len(frames) - 1:
            print(f"  rendered {n+1}/{len(frames)}", flush=True)

    manifest = {
        "frames": len(frames),
        "field": args.field,
        "layout": args.layout,
        "vmin": vmin, "vmax": vmax, "cmap": cmap,
        "iso_field": args.iso_field, "iso_level": args.iso_level,
        "view": [args.elev, args.azim],
        "fps": args.fps,
    }
    with open(os.path.join(outdir, "manifest.json"), "w") as fh:
        json.dump(manifest, fh, indent=2)

    if not args.no_video:
        mp4 = os.path.join(outdir, f"{args.layout}_{args.field}.mp4")
        stitch(png_dir, mp4, args.fps)
        print(f"  wrote {mp4}")


if __name__ == "__main__":
    main()
