#!/usr/bin/env python
"""
waterfall - external aerodynamics CFD runner.

    python waterfall.py cases/tesla.yaml
    python waterfall.py cases/tesla.yaml --check      # pre-flight only
    python waterfall.py cases/tesla.yaml --preview    # + solid-fraction render

The pre-flight always runs. A 3D case is a multi-minute commitment and the
common failure -- geometry rotated onto the wrong axis, body clipping a
boundary, or a case that needs four hours -- is cheap to catch beforehand and
expensive to discover afterwards.
"""

import argparse
import json
import os
import sys
import time

import numpy as np

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

from waterfall_core import Grid, Solver, external_flow, geometry, io
from waterfall_core.config import load_config, build_case, preflight, describe
from waterfall_core.sgs import VremanSGS


def preview_geometry(grid, mesh, eta_c, path):
    """Render three orthogonal slices of the solid fraction.

    This is the check that the geometry landed where you think it did. It is
    the cheapest possible guard against the most expensive mistake -- an
    axis_order that silently puts the car sideways to the flow still produces
    a plausible-looking drag number.
    """
    import matplotlib
    matplotlib.use("Agg")
    import matplotlib.pyplot as plt

    eta = eta_c[1:-1, 1:-1, 1:-1]
    fig, axes = plt.subplots(1, 3, figsize=(16, 4.2), constrained_layout=True)
    ox, oy, oz = grid.origin

    # Slice through the *body*, not the domain. A car sits on the ground and
    # occupies the bottom fifth of a domain sized for its wake, so the domain
    # mid-height plane cuts nothing but empty air and the preview looks like
    # the geometry failed to load.
    if eta.any():
        ky = int(np.argmax(eta.sum(axis=(0, 2))))
        kz = int(np.argmax(eta.sum(axis=(0, 1))))
        kx = int(np.argmax(eta.sum(axis=(1, 2))))
    else:
        ky, kz, kx = grid.ny // 2, grid.nz // 2, grid.nx // 2

    axes[0].imshow(eta[:, :, kz].T, origin="lower", cmap="gray_r",
                   extent=[ox, ox + grid.lx, oy, oy + grid.ly], aspect="equal")
    axes[0].set_title(f"side view, z = {oz + (kz+0.5)*grid.dz:.2f} m")
    axes[0].set_xlabel("x, streamwise [m]"); axes[0].set_ylabel("y, up [m]")

    axes[1].imshow(eta[:, ky, :].T, origin="lower", cmap="gray_r",
                   extent=[ox, ox + grid.lx, oz, oz + grid.lz], aspect="equal")
    axes[1].set_title(f"plan view, y = {oy + (ky+0.5)*grid.dy:.2f} m")
    axes[1].set_xlabel("x, streamwise [m]"); axes[1].set_ylabel("z, span [m]")

    axes[2].imshow(eta[kx, :, :].T, origin="lower", cmap="gray_r",
                   extent=[oy, oy + grid.ly, oz, oz + grid.lz], aspect="equal")
    axes[2].set_title(f"frontal, x = {ox + (kx+0.5)*grid.dx:.2f} m")
    axes[2].set_xlabel("y, up [m]"); axes[2].set_ylabel("z, span [m]")

    fig.suptitle("solid fraction -- confirm orientation before running")
    fig.savefig(path, dpi=110)
    plt.close(fig)


def main():
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("config")
    ap.add_argument("--check", action="store_true",
                    help="run pre-flight checks and exit")
    ap.add_argument("--preview", action="store_true",
                    help="also write a solid-fraction preview image")
    ap.add_argument("--force", action="store_true",
                    help="run even if pre-flight reports errors")
    ap.add_argument("--override", nargs="*", default=[],
                    help="config overrides, e.g. domain.cells_per_length=24")
    args = ap.parse_args()

    cfg = load_config(args.config)
    for ov in args.override:
        key, _, val = ov.partition("=")
        section, _, name = key.partition(".")
        if section not in cfg or not name:
            raise SystemExit(f"bad override {ov!r}")
        try:
            val = json.loads(val)
        except json.JSONDecodeError:
            pass
        cfg[section][name] = val

    t0 = time.time()
    grid, mesh, plan = build_case(cfg)

    outdir = cfg["run"]["output"]
    os.makedirs(outdir, exist_ok=True)

    # Rasterise before reporting so the reference area quoted in the
    # pre-flight is the one Cd is actually divided by. The cheap silhouette
    # estimate in mesh_report fills enclosed holes -- right for a closed car
    # body, wrong for a part with genuine through-gaps such as a bumper with
    # brake ducts, where it inflates the reference area and depresses Cd by
    # the same factor with nothing on screen to show it happened.
    print("rasterising geometry ...")
    eu, ev, ew, ec, area_mask = geometry.solid_fractions(
        grid, mesh, refine=2, return_area=True
    )
    solid_frac = float(ec[1:-1, 1:-1, 1:-1].mean())

    ok, checks = preflight(cfg, grid, mesh, plan, strict=not args.force,
                           frontal_area=area_mask)

    print(describe(cfg, grid, mesh, plan, checks))
    print(f"  rasterised frontal area {area_mask:.4f} m^2; "
          f"solid fills {100*solid_frac:.2f}% of domain cells")
    print(f"(pre-flight took {time.time()-t0:.1f}s)")

    if solid_frac <= 0:
        print("  ERROR: geometry rasterised to nothing. Check geometry.units "
              "and geometry.axis_order.")
        ok = False

    if args.preview or args.check:
        pv = os.path.join(outdir, "geometry_preview.png")
        preview_geometry(grid, mesh, ec, pv)
        print(f"  wrote {pv}")

    if args.check:
        return 0 if ok else 1
    if not ok:
        print("\npre-flight failed. Fix the errors above or pass --force.")
        return 1

    fcfg, rcfg = cfg["flow"], cfg["run"]
    u_inf = (fcfg["speed"], 0.0, 0.0)
    bc = external_flow(
        grid, u_inf=u_inf,
        ground=fcfg["ground"],
        road_speed=fcfg["speed"] if fcfg["rolling_road"] else None,
        sides=fcfg["sides"],
    )
    sgs_name = fcfg.get("sgs")
    if sgs_name in (None, "none", False):
        sgs = None
        if plan["reynolds"] > 1e5:
            print("  WARNING: no subgrid model at Re = "
                  f"{plan['reynolds']:.2g}. Central differencing has no "
                  "dissipation of its own; expect the timestep to collapse.")
    elif str(sgs_name).lower() == "vreman":
        sgs = VremanSGS()
    else:
        raise SystemExit(f"unknown flow.sgs {sgs_name!r}; use 'vreman' or null")

    solver = Solver(
        grid, bc, nu=fcfg["viscosity"], rho=fcfg["density"],
        eta_u=eu, eta_v=ev, eta_w=ew, eta_c=ec, cfl=rcfg["cfl"],
        sgs=sgs,
    )

    frames_dir = os.path.join(outdir, "frames")
    n_frames = int(rcfg["frames"])
    t_end = plan["t_end"]
    next_frame = [0]
    frame_times = np.linspace(0.0, t_end, n_frames) if n_frames > 0 else []

    def on_step(s, _n, _dt):
        if next_frame[0] < len(frame_times) and s.t >= frame_times[next_frame[0]]:
            io.save_frame(frames_dir, next_frame[0], s,
                          fields=rcfg["save_fields"])
            next_frame[0] += 1
        return True

    print(f"\nrunning to t = {t_end:.3f} s ...")
    elapsed = solver.run(t_end=t_end, callback=on_step,
                         log_every=rcfg["log_every"])

    area = checks["mesh"]["frontal_area_m2"]
    coef = solver.coefficients(area, u_ref=fcfg["speed"], window=0.5)

    summary = {
        "config": args.config,
        "geometry": checks["mesh"],
        "grid": {"nx": grid.nx, "ny": grid.ny, "nz": grid.nz,
                 "cells": grid.n_cells, "dx": grid.dx,
                 "cells_per_length": plan["cells_per_length"]},
        "flow": dict(fcfg),
        "reynolds": plan["reynolds"],
        "reference_area_m2": area,
        "blockage_pct": checks["blockage_pct"],
        "cd": coef["cd"], "cd_std": coef["cd_std"],
        "cl": coef["cl"], "cl_std": coef["cl_std"],
        "drag_force_N": coef["fx_mean_N"],
        "lift_force_N": coef["fy_mean_N"],
        "downforce_N": -coef["fy_mean_N"] if coef["fy_mean_N"] < 0 else 0.0,
        "steps": solver.step_count,
        "sim_time_s": solver.t,
        "wall_time_s": elapsed,
        "estimated_wall_time_s": plan["estimated_seconds"],
        "final_divergence": solver.divergence_norm(),
        "ibm_slip": solver.ibm_slip(),
        "sgs": sgs_name,
        "max_nu_t_over_nu": (sgs.max_ratio(solver) if sgs else 0.0),
        "frames_written": next_frame[0],
        "averaging_window": "last 50% of the run",
    }
    with open(os.path.join(outdir, "summary.json"), "w") as fh:
        json.dump(summary, fh, indent=2)

    print("\n" + "=" * 68)
    print(f"Cd = {coef['cd']:.4f}  (rms fluctuation {coef['cd_std']:.4f})")
    print(f"Cl = {coef['cl']:.4f}  (rms fluctuation {coef['cl_std']:.4f})")
    print(f"drag  {coef['fx_mean_N']:.2f} N   lift {coef['fy_mean_N']:.2f} N   "
          f"ref area {area:.4f} m^2")
    print(f"{solver.step_count:,} steps in {elapsed/60:.1f} min "
          f"(estimated {plan['estimated_seconds']/60:.1f})")
    print(f"final divergence {summary['final_divergence']:.2e}, "
          f"IBM slip {summary['ibm_slip']:.3f}")
    print(f"wrote {outdir}/summary.json and {next_frame[0]} frames")
    print("=" * 68)
    print(f"\nvisualise with:\n  python visualize_waterfall.py {frames_dir}")
    return 0


if __name__ == "__main__":
    sys.exit(main())
