"""
Benchmark 2 -- flow past a circular cylinder.

Why this case
-------------
It is the reference bluff-body flow, and it tests the two things the flat
plate cannot: whether the immersed boundary produces correct pressure drag,
and whether the scheme sustains vortex shedding at the right frequency.
Shedding is a genuinely hard test -- an over-dissipative scheme damps it out
entirely and reports a plausible-looking steady wake, which is exactly the
failure mode a drag-only check would miss.

Regimes covered:
  Re=40   steady, symmetric recirculation. Tests pressure drag and separation
          without any unsteady complication.
  Re=100  laminar periodic shedding. Tests Strouhal number and lift amplitude.
  Re=200  shedding, still 2D. Tests Reynolds-number trend.

Reference values (consensus of experiment and high-order DNS):
  Re=40   Cd 1.48-1.56, recirculation length L/D 2.13-2.35, sep. angle ~126 deg
  Re=100  Cd 1.32-1.38, St 0.164-0.167, Cl_rms 0.30-0.34
  Re=200  Cd 1.30-1.34, St 0.192-0.197, Cl_rms 0.60-0.70

Above Re ~ 190 the real wake becomes three-dimensional, so a 2D computation
at Re=200 is expected to overpredict Cd and Cl_rms slightly. That is a
property of the case, not a solver defect, and it is flagged in the report.
"""

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

import numpy as np

from waterfall_core import Grid, Solver, external_flow
from waterfall_core import shapes

D = 1.0
U = 1.0


def build(re, cells_per_d=24, lx_d=30.0, ly_d=20.0, x_c_d=8.0):
    nx = int(lx_d * cells_per_d)
    ny = int(ly_d * cells_per_d)
    grid = Grid(nx, ny, 1, lx_d * D, ly_d * D, 1.0)
    # Far field on all open sides. Blockage D/(ly_d*D) = 5% at the default
    # height, which inflates Cd by roughly 3-5%; reported alongside the result
    # rather than silently corrected.
    bc = external_flow(grid, u_inf=(U, 0.0, 0.0), sides="farfield")
    center = (x_c_d * D, 0.5 * ly_d * D)
    eu, ev, ew, ec = shapes.circle(grid, center, D, samples=4)
    nu = U * D / re
    solver = Solver(
        grid, bc, nu=nu, rho=1.0,
        eta_u=eu, eta_v=ev, eta_w=ew, eta_c=ec, cfl=0.4,
    )
    solver.center = center
    return solver


def strouhal(times, cl, d=D, u=U):
    """Shedding frequency from the lift signal.

    Uses the second half of the record only -- the impulsive start produces a
    strong transient whose spectral leakage would otherwise dominate. Returns
    (St, peak_power_fraction); a low power fraction means the signal is not
    cleanly periodic and the St value should not be trusted.
    """
    t = np.asarray(times)
    y = np.asarray(cl)
    n = len(t)
    if n < 64:
        return float("nan"), 0.0
    half = n // 2
    t, y = t[half:], y[half:]
    y = y - y.mean()
    if np.allclose(y, 0.0):
        return 0.0, 0.0

    dt = float(np.mean(np.diff(t)))
    win = np.hanning(len(y))
    spec = np.abs(np.fft.rfft(y * win)) ** 2
    freqs = np.fft.rfftfreq(len(y), dt)
    spec[0] = 0.0
    k = int(np.argmax(spec))
    total = spec.sum()
    return float(freqs[k] * d / u), float(spec[k] / total if total > 0 else 0.0)


def recirculation_length(solver):
    """
    Distance from the rear stagnation point to where centreline u returns to
    zero, normalised by D. Only meaningful for a steady wake (Re=40).
    """
    g = solver.grid
    uc, _v, _w, _p = solver.cell_fields()
    x, y, _ = g.cell_centers()
    cx, cy = solver.center
    j = int(np.argmin(np.abs(y - cy)))
    line = uc[:, j, 0]
    eta = solver.eta_c[1:-1, 1:-1, 1:-1][:, j, 0]
    rear = cx + 0.5 * D

    # Step past any cell still substantially inside the body. The solid
    # fraction is diffuse, so the first cell past the nominal rear stagnation
    # point can still be partly forced and carry a small positive velocity --
    # starting the search there reports a recirculation length of exactly
    # zero and looks like the bubble is missing entirely.
    i = int(np.searchsorted(x, rear))
    while i < len(x) and eta[i] > 0.1:
        i += 1
    if i >= len(x) or line[i] >= 0.0:
        return 0.0

    for k in range(i, len(x)):
        if line[k] > 0.0:
            f = -line[k - 1] / (line[k] - line[k - 1])
            xz = x[k - 1] + f * (x[k] - x[k - 1])
            return float((xz - rear) / D)
    return float("nan")


def run_case(re, t_end, cells_per_d=24, verbose=True, **kw):
    solver = build(re, cells_per_d=cells_per_d, **kw)
    if verbose:
        print(f"\n  Re={re}  {solver.grid}  D/dx={D/solver.grid.dx:.0f}")

    # Break the symmetry so shedding can start. A perfectly symmetric
    # discrete problem stays symmetric forever, and the wake would remain
    # artificially steady well above the critical Reynolds number -- the
    # solver would look over-dissipative when it is merely undisturbed.
    if re > 50:
        g = solver.grid
        rng = np.random.default_rng(re)
        solver.v[1:-1, 2:g.ny + 1, 1:-1] += 1e-3 * U * rng.standard_normal(
            (g.nx, g.ny - 1, g.nz)
        )

    elapsed = solver.run(t_end=t_end, log_every=2000 if verbose else 0)

    area = D * solver.grid.dz  # frontal area per unit span
    coef = solver.coefficients(area, u_ref=U, window=0.5)
    t = np.asarray(solver.time_history)[1:]
    f = solver.filtered_force_history()
    q = 0.5 * solver.rho * U**2 * area
    cl_sig = f[:, 1] / q
    st, power = strouhal(t, cl_sig)

    half = len(cl_sig) // 2
    cl_rms = float(np.sqrt(np.mean((cl_sig[half:] - cl_sig[half:].mean()) ** 2)))

    out = {
        "Re": re,
        "cells_per_D": cells_per_d,
        "cd": coef["cd"],
        "cd_std": coef["cd_std"],
        "cl_mean": coef["cl"],
        "cl_rms": cl_rms,
        "strouhal": st,
        "spectral_peak_fraction": power,
        "recirc_length_over_D": recirculation_length(solver),
        "steps": solver.step_count,
        "elapsed_s": elapsed,
        "divergence": solver.divergence_norm(),
        "ibm_slip": solver.ibm_slip(),
        "blockage_pct": 100.0 * D / solver.grid.ly,
    }
    if verbose:
        print(f"    Cd={out['cd']:.3f} +/- {out['cd_std']:.3f}   "
              f"Cl_rms={cl_rms:.3f}   St={st:.4f}   "
              f"Lr/D={out['recirc_length_over_D']:.2f}   "
              f"div={out['divergence']:.1e}  {elapsed:.0f}s")
    return solver, out


REFERENCE = {
    40:  {"cd": (1.48, 1.56), "recirc": (2.13, 2.35), "st": None,
          "cl_rms": (0.0, 0.01)},
    100: {"cd": (1.32, 1.38), "recirc": None, "st": (0.164, 0.167),
          "cl_rms": (0.30, 0.34)},
    200: {"cd": (1.30, 1.34), "recirc": None, "st": (0.192, 0.197),
          "cl_rms": (0.60, 0.70)},
}


def _verdict(value, rng):
    if rng is None or value is None or not np.isfinite(value):
        return "-", None
    lo, hi = rng
    mid = 0.5 * (lo + hi)
    err = 100.0 * (value - mid) / mid if mid != 0 else float("nan")
    return ("in range" if lo <= value <= hi else f"{err:+.1f}%"), err


def main():
    print("Benchmark: circular cylinder")
    cases = [(40, 110.0), (100, 170.0), (200, 170.0)]
    results = []
    for re, t_end in cases:
        _s, out = run_case(re, t_end)
        ref = REFERENCE[re]
        out["cd_verdict"], out["cd_err_pct"] = _verdict(out["cd"], ref["cd"])
        out["st_verdict"], out["st_err_pct"] = _verdict(out["strouhal"], ref["st"])
        out["recirc_verdict"], out["recirc_err_pct"] = _verdict(
            out["recirc_length_over_D"], ref["recirc"]
        )
        out["clrms_verdict"], out["clrms_err_pct"] = _verdict(
            out["cl_rms"], ref["cl_rms"]
        )
        out["reference"] = {k: v for k, v in ref.items()}
        results.append(out)

    print()
    print(f"  {'Re':>5} {'Cd':>7} {'ref Cd':>13} {'verdict':>10} "
          f"{'St':>7} {'ref St':>14} {'verdict':>10} {'Cl_rms':>8} {'Lr/D':>7}")
    for r in results:
        ref = REFERENCE[r["Re"]]
        cds = f"{ref['cd'][0]}-{ref['cd'][1]}"
        sts = f"{ref['st'][0]}-{ref['st'][1]}" if ref["st"] else "steady"
        print(f"  {r['Re']:5d} {r['cd']:7.3f} {cds:>13} {r['cd_verdict']:>10} "
              f"{r['strouhal']:7.4f} {sts:>14} {r['st_verdict']:>10} "
              f"{r['cl_rms']:8.3f} {r['recirc_length_over_D']:7.2f}")

    os.makedirs("validation", exist_ok=True)
    with open("validation/cylinder.json", "w") as fh:
        json.dump({"case": "cylinder", "results": results}, fh, indent=2)
    print("\n  wrote validation/cylinder.json")
    return results


if __name__ == "__main__":
    main()
