#!/usr/bin/env python
"""
waterfall UI -- local web app.

    python ui.py

Opens http://127.0.0.1:8010 in your browser. Drag an STL onto the page, check
the orientation preview, press Run.

Why a local web app rather than a desktop GUI: it needs no toolkit install, it
renders the MP4s and PNGs the pipeline already produces without any extra
work, and it runs the solver as an ordinary subprocess so a crash in a run can
never take the interface down with it. The server binds to loopback only.
"""

import argparse
import asyncio
import io
import json
import os
import shutil
import sys
import threading
import webbrowser

import numpy as np
from fastapi import FastAPI, UploadFile, File, HTTPException
from fastapi.responses import HTMLResponse, JSONResponse, FileResponse
from fastapi.staticfiles import StaticFiles
from pydantic import BaseModel

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

from waterfall_core import geometry
from waterfall_core.runs import RunManager, OUTPUT_ROOT, slugify

HERE = os.path.dirname(os.path.abspath(__file__))
GEOM_DIR = os.path.join(HERE, "geometries")
UI_DIR = os.path.join(HERE, "ui")

app = FastAPI(title="waterfall")
manager = RunManager()

os.makedirs(GEOM_DIR, exist_ok=True)
os.makedirs(OUTPUT_ROOT, exist_ok=True)

# StaticFiles is used for results rather than a hand-rolled handler because it
# implements HTTP range requests, which the browser needs to seek in an MP4.
app.mount("/files", StaticFiles(directory=OUTPUT_ROOT), name="files")
app.mount("/static", StaticFiles(directory=UI_DIR), name="static")


# ----------------------------------------------------------------------
# geometry inspection
# ----------------------------------------------------------------------

def guess_orientation(size):
    """
    Guess which STL axis is streamwise, vertical and spanwise.

    Heuristic: the longest axis is streamwise, and of the remaining two the
    shorter is vertical. That is right for cars, wings and most vehicle
    components, and wrong often enough that the UI always shows the preview
    and lets you override it -- an axis_order error still produces a
    plausible-looking drag number, so it must be checked by eye.
    """
    order = list(np.argsort(size)[::-1])       # longest first
    stream = order[0]
    rest = [a for a in range(3) if a != stream]
    vertical = rest[0] if size[rest[0]] <= size[rest[1]] else rest[1]
    span = [a for a in rest if a != vertical][0]
    return "".join("xyz"[a] for a in (stream, vertical, span))


@app.post("/api/upload")
async def upload(file: UploadFile = File(...)):
    if not file.filename.lower().endswith((".stl", ".obj", ".ply")):
        raise HTTPException(400, "Need an .stl, .obj or .ply file")
    dest = os.path.join(GEOM_DIR, os.path.basename(file.filename))
    with open(dest, "wb") as fh:
        shutil.copyfileobj(file.file, fh)
    return await inspect_path(dest)


class InspectReq(BaseModel):
    path: str
    units: str = "mm"
    scale: float = 1.0
    axis_order: str | None = None


async def inspect_path(path, units="mm", scale=1.0, axis_order=None):
    """Load, orient and measure a geometry. Runs off the event loop because
    a 300k-face mesh takes seconds and would otherwise freeze the UI."""
    def work():
        mesh = geometry.load_mesh(path, units=units, scale=scale)
        raw = mesh.bounds[1] - mesh.bounds[0]
        order = axis_order or guess_orientation(raw)
        oriented = geometry.orient_mesh(mesh, axis_order=order)
        size = oriented.bounds[1] - oriented.bounds[0]
        return {
            "path": os.path.relpath(path, HERE).replace("\\", "/"),
            "name": os.path.basename(path),
            "units": units,
            "scale": scale,
            "axis_order": order,
            "guessed_axis_order": guess_orientation(raw),
            "faces": int(len(oriented.faces)),
            "watertight": bool(oriented.is_watertight),
            "raw_size": [float(v) for v in raw],
            "size_m": [float(v) for v in size],
            "length_m": float(size[0]),
            "height_m": float(size[1]),
            "width_m": float(size[2]),
        }

    return await asyncio.get_running_loop().run_in_executor(None, work)


@app.post("/api/inspect")
async def inspect(req: InspectReq):
    p = os.path.join(HERE, req.path) if not os.path.isabs(req.path) else req.path
    if not os.path.exists(p):
        raise HTTPException(404, f"not found: {req.path}")
    return await inspect_path(p, req.units, req.scale, req.axis_order)


@app.get("/api/preview")
async def preview(path: str, units: str = "mm", scale: float = 1.0,
                  axis_order: str = "xyz"):
    """
    Three orthographic silhouettes of the oriented geometry, as a PNG.

    This is the single most important control in the UI. The axis guess is a
    bounding-box heuristic and it is wrong for whole classes of part -- a
    bumper or a wing is *wider* than it is long, so "longest axis is
    streamwise" points the flow along the span. That mistake does not produce
    an error; it produces a confident, plausible, wrong drag number. Telling
    the user to check the preview is useless unless there is a preview.

    Vertices are scattered rather than rasterised, which is approximate but
    renders a 300k-face mesh in well under a second -- fast enough to redraw
    on every change to units, scale or axis order.
    """
    full = os.path.join(HERE, path) if not os.path.isabs(path) else path
    if not os.path.exists(full):
        raise HTTPException(404, "not found")

    def work():
        import matplotlib
        matplotlib.use("Agg")
        import matplotlib.pyplot as plt

        mesh = geometry.load_mesh(full, units=units, scale=scale)
        mesh = geometry.orient_mesh(mesh, axis_order=axis_order)
        v = mesh.vertices
        step = max(1, len(v) // 60000)
        v = v[::step]

        fig, axes = plt.subplots(1, 3, figsize=(12, 3.4), constrained_layout=True)
        fig.patch.set_facecolor("#161b23")
        views = [
            (0, 1, "side view", "x  flow → [m]", "y  up [m]"),
            (0, 2, "plan view", "x  flow → [m]", "z  span [m]"),
            (2, 1, "front view (what the flow sees)", "z  span [m]", "y  up [m]"),
        ]
        for ax, (a, b, title, xl, yl) in zip(axes, views):
            ax.scatter(v[:, a], v[:, b], s=0.4, c="#4aa8ff", alpha=0.35,
                       linewidths=0, rasterized=True)
            ax.set_title(title, color="#e6edf5", fontsize=10)
            ax.set_xlabel(xl, color="#8b98a9", fontsize=8)
            ax.set_ylabel(yl, color="#8b98a9", fontsize=8)
            ax.set_aspect("equal")
            ax.set_facecolor("#0e1116")
            ax.tick_params(colors="#8b98a9", labelsize=7)
            for sp in ax.spines.values():
                sp.set_color("#2a323e")

        buf = io.BytesIO()
        fig.savefig(buf, format="png", dpi=96, facecolor=fig.get_facecolor())
        plt.close(fig)
        buf.seek(0)
        return buf.read()

    png = await asyncio.get_running_loop().run_in_executor(None, work)
    from fastapi.responses import Response
    return Response(png, media_type="image/png",
                    headers={"Cache-Control": "no-store"})


@app.get("/api/geometries")
def geometries():
    out = []
    for n in sorted(os.listdir(GEOM_DIR)):
        if n.lower().endswith((".stl", ".obj", ".ply")):
            p = os.path.join(GEOM_DIR, n)
            out.append({"name": n, "size": os.path.getsize(p),
                        "path": f"geometries/{n}"})
    return out


# ----------------------------------------------------------------------
# runs
# ----------------------------------------------------------------------

CONFIG_TEMPLATE = """# generated by the waterfall UI
geometry:
  path: {path}
  units: {units}
  scale: {scale}
  axis_order: {axis_order}
  rotate_deg: [0.0, 0.0, 0.0]
  flip: [false, false, false]

domain:
  cells_per_length: {cells}
  upstream: {upstream}
  downstream: {downstream}
  side: {side}
  above: {above}
  ground_clearance: {clearance}
  max_cells: {max_cells}

flow:
  speed: {speed}
  density: 1.225
  viscosity: 1.5e-5
  ground: {ground}
  rolling_road: {rolling_road}
  sides: farfield
  sgs: {sgs}

run:
  convective_times: {convective_times}
  cfl: 0.4
  frames: {frames}
  save_fields: [vorticity, q]
  output: {output}
  time_budget_s: 1800
  log_every: 100
"""


class RunReq(BaseModel):
    path: str
    label: str = ""
    units: str = "mm"
    scale: float = 1.0
    axis_order: str = "xyz"
    cells_per_length: int = 32
    speed: float = 30.0
    upstream: float = 1.5
    downstream: float = 4.0
    side: float = 2.0
    above: float = 3.0
    ground: bool = True
    rolling_road: bool = True
    ground_clearance: float | None = 0.12
    convective_times: float = 6.0
    frames: int = 120
    max_cells: int = 2500000
    sgs: str = "vreman"
    check_only: bool = False


def build_config(req: RunReq, out_dir):
    return CONFIG_TEMPLATE.format(
        path=req.path, units=req.units, scale=req.scale,
        axis_order=req.axis_order, cells=req.cells_per_length,
        upstream=req.upstream, downstream=req.downstream,
        side=req.side, above=req.above,
        clearance="null" if req.ground_clearance is None else req.ground_clearance,
        max_cells=req.max_cells, speed=req.speed,
        ground=str(req.ground).lower(),
        rolling_road=str(req.rolling_road).lower(),
        sgs="null" if req.sgs in ("none", "null", "") else req.sgs,
        convective_times=req.convective_times, frames=req.frames,
        output=out_dir.replace("\\", "/"),
    )


@app.post("/api/run")
def start_run(req: RunReq):
    label = req.label or slugify(req.path)
    run = manager.create(label, "")
    cfg = build_config(req, run.dir)
    with open(os.path.join(run.dir, "config.yaml"), "w") as fh:
        fh.write(cfg)
    # --preview on every run, not just pre-flight checks: it costs one
    # matplotlib figure and it is what gives each project a thumbnail in the
    # gallery. A results grid of identical blank cards is nearly useless for
    # finding an old run.
    args = ("--check",) if req.check_only else ("--preview",)
    manager.start(run.id, extra_args=args)
    return {"id": run.id}


@app.get("/api/runs")
def list_runs():
    manager.scan()
    return manager.list()


@app.get("/api/runs/{run_id}")
def get_run(run_id: str):
    d = manager.get(run_id)
    if d is None:
        raise HTTPException(404, "no such run")
    return d


@app.post("/api/runs/{run_id}/stop")
def stop_run(run_id: str):
    return {"stopped": manager.stop(run_id)}


@app.delete("/api/runs/{run_id}")
def delete_run(run_id: str):
    if not manager.delete(run_id):
        raise HTTPException(404, "no such run")
    return {"deleted": run_id}


class VizReq(BaseModel):
    layout: str = "dashboard"
    field: str = "cp"
    iso_field: str = "q"
    fps: int = 24
    stride: int = 2


@app.post("/api/runs/{run_id}/visualize")
def visualize(run_id: str, req: VizReq):
    run = manager.runs.get(run_id)
    if not run:
        raise HTTPException(404, "no such run")
    frames = os.path.join(run.dir, "frames")
    if not os.path.isdir(frames):
        raise HTTPException(400, "this run has no frames to visualise")

    import subprocess
    cmd = [
        sys.executable, "-u", "visualize_waterfall.py", frames,
        "--layout", req.layout, "--field", req.field,
        "--iso-field", req.iso_field, "--fps", str(req.fps),
        "--stride", str(req.stride),
    ]
    log = os.path.join(run.dir, "visualize.log")

    def go():
        with open(log, "w", buffering=1, errors="replace") as fh:
            subprocess.run(cmd, stdout=fh, stderr=subprocess.STDOUT, cwd=HERE)

    threading.Thread(target=go, daemon=True).start()
    return {"started": True, "log": f"/files/{run_id}/visualize.log"}


@app.get("/api/estimate")
def estimate(cells_per_length: int, length_m: float, speed: float,
             upstream: float = 1.5, downstream: float = 4.0,
             side: float = 2.0, above: float = 3.0,
             height_m: float = 1.0, width_m: float = 1.0,
             convective_times: float = 6.0, max_cells: int = 2500000):
    """
    Live runtime/cell-count estimate for the settings panel.

    Mirrors config.estimate_runtime so the number the UI shows is the number
    the pre-flight will print -- a UI that estimates differently from the tool
    it drives is worse than one that does not estimate at all.
    """
    from waterfall_core.config import DEFAULT_COST_NS

    lx = length_m * (1.0 + upstream + downstream)
    ly = height_m * (1.0 + above)
    lz = width_m * (1.0 + 2.0 * side)
    h = length_m / max(cells_per_length, 1)
    nx, ny, nz = (max(int(round(v / h)), 8) for v in (lx, ly, lz))
    cells = nx * ny * nz
    if cells > max_cells:
        f = (max_cells / cells) ** (1 / 3)
        nx, ny, nz = (max(int(v * f), 8) for v in (nx, ny, nz))
        h = lx / nx
        cells = nx * ny * nz

    dx, dy, dz = lx / nx, ly / ny, lz / nz
    u_peak = 2.5 * speed
    dt = 0.4 / (u_peak / dx + 0.5 * u_peak / dy + 0.5 * u_peak / dz)
    t_end = convective_times * lx / speed
    steps = int(np.ceil(t_end / dt))
    seconds = steps * cells * DEFAULT_COST_NS * 1e-9
    return {
        "cells": cells, "nx": nx, "ny": ny, "nz": nz,
        "steps": steps, "seconds": seconds, "dt": dt,
        "reynolds": speed * length_m / 1.5e-5,
        "cells_per_length": length_m / dx,
        "within_budget": seconds <= 1800,
    }


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

@app.get("/", response_class=HTMLResponse)
def index():
    with open(os.path.join(UI_DIR, "index.html"), encoding="utf-8") as fh:
        return fh.read()


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--port", type=int, default=8010)
    ap.add_argument("--host", default="127.0.0.1")
    ap.add_argument("--no-browser", action="store_true")
    args = ap.parse_args()

    url = f"http://{args.host}:{args.port}"
    print(f"waterfall UI -> {url}")
    if not args.no_browser:
        threading.Timer(1.2, lambda: webbrowser.open(url)).start()

    import uvicorn
    uvicorn.run(app, host=args.host, port=args.port, log_level="warning")


if __name__ == "__main__":
    main()
