"""
Benchmark 1 -- Blasius laminar flat-plate boundary layer.

Why this case
-------------
It is the only external-flow case with a genuine analytical solution, and it
tests the two things a drag number depends on most: whether viscous diffusion
is discretised correctly, and whether the no-slip wall produces the right wall
shear. A solver can get bluff-body pressure drag roughly right while being
badly wrong here, so this runs first.

Setup: uniform inflow over a no-slip bottom wall, leading edge at the inlet.
Slip top, convective outlet. Laminar, no model, no tuning.

Reference: similarity solution of f''' + (1/2) f f'' = 0 with f(0)=f'(0)=0,
f'(inf)=1, giving u/U = f'(eta), eta = y*sqrt(U/(nu*x)).

Metrics reported
----------------
u-profile L2 and max error against f'(eta) at three stations; momentum
thickness theta vs 0.664*sqrt(nu x/U); skin friction cf vs 0.664/sqrt(Re_x).
Momentum thickness is the one that matters for drag -- it is literally the
drag integral -- so it is weighted most heavily in the verdict.
"""

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

import numpy as np
from scipy.integrate import solve_ivp
from scipy.optimize import brentq

from waterfall_core import Grid, Solver, external_flow

ETA_MAX = 10.0


def blasius_profile(eta_query):
    """f'(eta) by shooting on f''(0). The known value is 0.332057; solving
    for it rather than hard-coding keeps the reference self-checking."""

    def rhs(_t, y):
        return [y[1], y[2], -0.5 * y[0] * y[2]]

    def shoot(fpp0):
        sol = solve_ivp(
            rhs, (0.0, ETA_MAX), [0.0, 0.0, fpp0],
            rtol=1e-11, atol=1e-13, dense_output=True,
        )
        return sol

    def residual(fpp0):
        return shoot(fpp0).y[1, -1] - 1.0

    fpp0 = brentq(residual, 0.1, 1.0, xtol=1e-13)
    sol = shoot(fpp0)
    eta_q = np.clip(eta_query, 0.0, ETA_MAX)
    fp = sol.sol(eta_q)[1]
    return np.clip(fp, 0.0, 1.0), fpp0


def run(nx=600, ny=192, lx=1.2, ly=0.3, u_inf=1.0, nu=2e-4,
        x_le=0.2, t_end=8.0, verbose=True):
    """
    x_le places the plate's leading edge inside the domain, with a slip
    section upstream. See BoundaryConditions.ylo_wall_range for why running
    the wall to the inlet corrupts the whole field.
    """
    grid = Grid(nx, ny, 1, lx, ly, 1.0)
    bc = external_flow(
        grid, u_inf=(u_inf, 0.0, 0.0), ground=True,
        ylo_wall_range=(x_le, lx),
    )
    solver = Solver(grid, bc, nu=nu, rho=1.0, cfl=0.4)
    solver.x_le = x_le

    if verbose:
        print(f"  {grid}")
        print(f"  nu={nu:g}  leading edge at x={x_le}  "
              f"Re_plate={u_inf*(lx-x_le)/nu:.0f}")

    elapsed = solver.run(t_end=t_end, log_every=2000 if verbose else 0)
    if verbose:
        print(f"  {solver.step_count} steps in {elapsed:.1f}s  "
              f"div={solver.divergence_norm():.2e}")
    return solver, elapsed


def evaluate(solver, u_inf=1.0, nu=2e-4, stations=(0.2, 0.4, 0.7)):
    """stations are distances *from the leading edge*, not from the inlet."""
    g = solver.grid
    uc, vc, _wc, _pc = solver.cell_fields()
    x, y, _ = g.cell_centers()
    x_le = getattr(solver, "x_le", 0.0)

    results = []
    for xs in stations:
        i = int(np.argmin(np.abs(x - (x_le + xs))))
        x_act = x[i] - x_le
        re_x = u_inf * x_act / nu

        prof = uc[i, :, 0]

        # Local edge velocity, not U_inf.
        #
        # The domain has a finite height, so the growing displacement
        # thickness squeezes the freestream and accelerates it by a few
        # percent. That is a real confinement effect, not a solver error, and
        # boundary-layer theory is written in terms of the *local* edge
        # velocity for exactly this reason -- it is also what a wind tunnel
        # measurement would use.
        #
        # Normalising by U_inf instead makes the integral thicknesses
        # useless: theta integrates u/U(1-u/U), so a 1% freestream excess
        # spread over the domain height contributes a negative term
        # comparable to theta itself and can flip its sign. Integrating only
        # up to the edge and dividing by U_e removes both problems.
        j_edge = int(np.argmax(prof))
        u_e = float(prof[j_edge])
        overshoot = 100.0 * (u_e / u_inf - 1.0)

        re_x_e = u_e * x_act / nu
        eta = y * np.sqrt(u_e / (nu * x_act))
        exact, _ = blasius_profile(eta)

        band = eta <= 6.0
        err = prof[band] / u_e - exact[band]
        l2 = float(np.sqrt(np.mean(err**2)))
        linf = float(np.abs(err).max())

        sl = slice(0, j_edge + 1)
        f = prof[sl] / u_e
        theta = float(np.trapezoid(f * (1.0 - f), y[sl]))
        theta_exact = 0.664 * np.sqrt(nu * x_act / u_e)

        dstar = float(np.trapezoid(1.0 - f, y[sl]))
        dstar_exact = 1.7208 * np.sqrt(nu * x_act / u_e)

        # Wall shear from the first cell centre, dy/2 off the wall.
        tau_w = nu * prof[0] / (0.5 * g.dy)
        cf = 2.0 * tau_w / u_e**2
        cf_exact = 0.664 / np.sqrt(re_x_e)

        results.append({
            "x": float(x_act),
            "Re_x": float(re_x_e),
            "u_edge": u_e,
            "edge_overshoot_pct": overshoot,
            "u_l2_error": l2,
            "u_max_error": linf,
            "theta": theta,
            "theta_exact": theta_exact,
            "theta_error_pct": float(100 * (theta - theta_exact) / theta_exact),
            "dstar": dstar,
            "dstar_exact": dstar_exact,
            "dstar_error_pct": float(100 * (dstar - dstar_exact) / dstar_exact),
            "shape_factor": float(dstar / theta) if theta > 0 else float("nan"),
            "shape_factor_exact": 2.59,
            "cf": float(cf),
            "cf_exact": float(cf_exact),
            "cf_error_pct": float(100 * (cf - cf_exact) / cf_exact),
            "pts_in_bl": int((eta <= 5.0).sum()),
        })
        _ = re_x
    return results


def main():
    print("Benchmark: Blasius flat-plate boundary layer")
    _fp, fpp0 = blasius_profile(np.array([0.0]))
    print(f"  reference f''(0) = {fpp0:.6f}  (exact 0.332057)")

    solver, elapsed = run()
    res = evaluate(solver)

    print()
    print(f"  {'x':>6} {'Re_x':>8} {'pts/BL':>7} {'L2(u)':>9} {'max(u)':>9} "
          f"{'theta%':>8} {'d*%':>8} {'H':>6} {'cf%':>8} {'blockage%':>10}")
    for r in res:
        print(f"  {r['x']:6.3f} {r['Re_x']:8.0f} {r['pts_in_bl']:7d} "
              f"{r['u_l2_error']:9.4f} {r['u_max_error']:9.4f} "
              f"{r['theta_error_pct']:+8.2f} {r['dstar_error_pct']:+8.2f} "
              f"{r['shape_factor']:6.2f} {r['cf_error_pct']:+8.2f} "
              f"{r['edge_overshoot_pct']:10.2f}")
    print("  (H = shape factor d*/theta; Blasius exact = 2.59)")

    payload = {
        "case": "blasius",
        "elapsed_s": elapsed,
        "steps": solver.step_count,
        "divergence": solver.divergence_norm(),
        "fpp0_reference": fpp0,
        "stations": res,
    }
    os.makedirs("validation", exist_ok=True)
    with open("validation/blasius.json", "w") as fh:
        json.dump(payload, fh, indent=2)
    print("\n  wrote validation/blasius.json")
    return payload


if __name__ == "__main__":
    main()
