"""
Incompressible Navier-Stokes solver: fractional step on a staggered grid with
an immersed boundary.

Governing equations (constant density, Newtonian, isothermal):

    du/dt + div(u u) = -grad(p)/rho + nu*lap(u) + f_ibm
    div(u) = 0

Time integration
----------------
Adams-Bashforth 2 on the whole right-hand side, one pressure projection per
step. Second-order in time, one Poisson solve per step. The alternative,
RK3 with a projection per substage, is more accurate but costs three Poisson
solves; at the grid sizes that fit a 30-minute budget the spatial error
dominates, so paying 3x for temporal accuracy buys nothing. RK3 is available
via integrator="rk3" for the cheap 2D verification cases where it is free.

The first step falls back to forward Euler because AB2 needs history.

Projection
----------
    u*      = u^n + dt*(3/2 R^n - 1/2 R^{n-1})     (predictor)
    u*     += eta*(u_body - u*)                    (immersed boundary)
    lap(phi) = div(u*)/dt                          (Poisson)
    u^{n+1} = u* - dt*grad(phi)                    (correction)

phi is the full kinematic pressure, not an increment, because the projection
is applied to the un-projected predictor each step. Static pressure is
therefore rho*phi with no accumulation needed.

Ordering note: the immersed-boundary forcing is applied *before* the
projection, not after. Forcing after projection would reintroduce a divergence
error inside the body that never gets removed, and the resulting spurious
mass source inside the solid shows up as a bogus contribution to drag. The
cost of forcing first is that the final velocity inside the body is not
exactly zero -- the projection perturbs it slightly. That residual is
monitored as `ibm_slip` in the diagnostics.

Known limits, stated up front
-----------------------------
* No wall model. The first grid point off the body is in the buffer layer or
  beyond for any realistic road-vehicle Reynolds number, so skin friction is
  under-resolved. Bluff-body drag is pressure-dominated so the total is still
  useful, but do not read this solver's friction drag as physical.
* Central differencing without a subgrid model is implicit LES. Dissipation
  comes from truncation error rather than physics. Turn on the Vreman model
  (sgs="vreman") for high-Reynolds runs.
* Uniform grid. Resolution near the body cannot be increased without
  increasing it everywhere. That is the direct cost of the FFT pressure solve.
"""

import time

import numpy as np

from . import kernels
from .grid import Grid, NG
from .poisson import PoissonDCT


class Solver:
    def __init__(
        self,
        grid,
        bc,
        nu,
        rho=1.225,
        eta_u=None,
        eta_v=None,
        eta_w=None,
        eta_c=None,
        body_velocity=(0.0, 0.0, 0.0),
        integrator="ab2",
        cfl=0.4,
        sgs=None,
        workers=-1,
    ):
        self.grid = grid
        self.bc = bc
        self.nu = float(nu)
        self.rho = float(rho)
        self.cfl = float(cfl)
        self.integrator = integrator
        self.sgs = sgs

        self.u = grid.zeros_u()
        self.v = grid.zeros_v()
        self.w = grid.zeros_w()
        self.p = grid.zeros_p()

        self.ru = grid.zeros_u()
        self.rv = grid.zeros_v()
        self.rw = grid.zeros_w()
        self.ru_old = grid.zeros_u()
        self.rv_old = grid.zeros_v()
        self.rw_old = grid.zeros_w()

        self.eta_u = eta_u if eta_u is not None else grid.zeros_u()
        self.eta_v = eta_v if eta_v is not None else grid.zeros_v()
        self.eta_w = eta_w if eta_w is not None else grid.zeros_w()
        self.eta_c = eta_c if eta_c is not None else grid.zeros_p()
        self.body_velocity = np.asarray(body_velocity, dtype=float)

        self.has_body = bool(
            self.eta_u.any() or self.eta_v.any() or self.eta_w.any()
        )

        self.poisson = PoissonDCT(grid, workers=workers)
        self._div = np.zeros((grid.nx, grid.ny, grid.nz))

        self.step_count = 0
        self.t = 0.0
        self.mass_correction = 0.0
        self.force = np.zeros(3)
        self.force_history = []
        self.time_history = []
        self._first_step = True
        self._dt_reference = None

        self.initialize_uniform(bc.u_inf)

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

    def initialize_uniform(self, u_inf):
        """Start from uniform freestream. The body is impulsively started,
        which is the standard way to reach a statistically steady wake."""
        self.u[:] = u_inf[0]
        self.v[:] = u_inf[1]
        self.w[:] = u_inf[2]
        if self.grid.two_d:
            self.w[:] = 0.0
        self.bc.apply(self.u, self.v, self.w)

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

    def stable_dt(self):
        """
        Timestep from the tighter of the convective and viscous limits.

        Convective: dt <= cfl / (|u|/dx + |v|/dy + |w|/dz), the multidimensional
        form -- using the per-axis minimum instead is optimistic and goes
        unstable on skewed flow.

        Viscous: dt <= 0.5 / (nu * sum(1/h^2)). Explicit diffusion. At the
        Reynolds numbers of interest this limit is loose; it binds only for
        the low-Re verification cases.
        """
        g = self.grid
        mu, mv, mw = kernels.max_abs_velocity(
            self.u, self.v, self.w, g.nx, g.ny, g.nz
        )
        conv_rate = mu / g.dx + mv / g.dy + (mw / g.dz if not g.two_d else 0.0)
        dt_conv = self.cfl / conv_rate if conv_rate > 1e-30 else np.inf

        inv_h2 = 1.0 / g.dx**2 + 1.0 / g.dy**2 + (0.0 if g.two_d else 1.0 / g.dz**2)
        dt_visc = 0.5 / (self.nu * inv_h2) if self.nu > 0 else np.inf

        dt = min(dt_conv, dt_visc)
        if not np.isfinite(dt):
            raise FloatingPointError(
                "No finite timestep: velocity field is zero and viscosity is zero."
            )
        return dt

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

    def _compute_rhs(self):
        g = self.grid
        # zfac switches off every z-derivative in 2D. That is what makes the
        # z ghost cells irrelevant and lets the BC layer skip them entirely
        # -- see boundary.fill_ghosts.
        zfac = 0.0 if g.two_d else 1.0
        kernels.rhs_u(
            self.u, self.v, self.w, self.nu,
            g.dx, g.dy, g.dz, g.nx, g.ny, g.nz, zfac, self.ru,
        )
        kernels.rhs_v(
            self.u, self.v, self.w, self.nu,
            g.dx, g.dy, g.dz, g.nx, g.ny, g.nz, zfac, self.rv,
        )
        if not g.two_d:
            kernels.rhs_w(
                self.u, self.v, self.w, self.nu,
                g.dx, g.dy, g.dz, g.nx, g.ny, g.nz, zfac, self.rw,
            )
        if self.sgs is not None:
            self.sgs.add(self)

    def _advance_momentum(self, dt):
        g = self.grid
        self._compute_rhs()
        if self._first_step:
            c0, c1 = 1.0, 0.0
            self._first_step = False
        else:
            c0, c1 = 1.5, -0.5

        # Interior-only updates. Boundary nodes are owned by the BC layer;
        # writing them here would overwrite the prescribed inlet value.
        su = (slice(2, g.nx + 1), slice(1, g.ny + 1), slice(1, g.nz + 1))
        sv = (slice(1, g.nx + 1), slice(2, g.ny + 1), slice(1, g.nz + 1))
        sw = (slice(1, g.nx + 1), slice(1, g.ny + 1), slice(2, g.nz + 1))

        self.u[su] += dt * (c0 * self.ru[su] + c1 * self.ru_old[su])
        self.v[sv] += dt * (c0 * self.rv[sv] + c1 * self.rv_old[sv])
        if not g.two_d:
            self.w[sw] += dt * (c0 * self.rw[sw] + c1 * self.rw_old[sw])

        self.ru, self.ru_old = self.ru_old, self.ru
        self.rv, self.rv_old = self.rv_old, self.rv
        self.rw, self.rw_old = self.rw_old, self.rw

    def _apply_ibm(self, dt):
        """Direct forcing. Returns the force on the body in Newtons."""
        g = self.grid
        if not self.has_body:
            return np.zeros(3)

        du = kernels.apply_ibm(
            self.u, self.eta_u, self.body_velocity[0], g.nx + 1, g.ny, g.nz
        )
        dv = kernels.apply_ibm(
            self.v, self.eta_v, self.body_velocity[1], g.nx, g.ny + 1, g.nz
        )
        dw = 0.0
        if not g.two_d:
            dw = kernels.apply_ibm(
                self.w, self.eta_w, self.body_velocity[2], g.nx, g.ny, g.nz + 1
            )

        # Momentum injected into the fluid per unit time is
        # rho * V_cell * sum(delta_u) / dt. Newton's third law flips the sign
        # to get the force the fluid exerts on the body.
        k = -self.rho * g.cell_volume / dt
        return np.array([du * k, dv * k, dw * k])

    def _project(self, dt):
        g = self.grid
        kernels.divergence(
            self.u, self.v, self.w, g.dx, g.dy, g.dz, self._div
        )
        self._div /= dt
        phi = self.poisson.solve(self._div, out=self.p)
        kernels.subtract_gradient(
            self.u, self.v, self.w, phi, dt,
            g.dx, g.dy, g.dz, g.nx, g.ny, g.nz,
        )

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

    def step(self, dt=None):
        if dt is None:
            dt = self.stable_dt()

        self.bc.advance_outflow(self.u, self.v, self.w, dt)
        self._advance_momentum(dt)
        self.force = self._apply_ibm(dt)

        # Boundary nodes must be final before the projection: the projection
        # takes boundary-normal velocity as given and only makes the interior
        # consistent with it. See boundary.py for why re-setting them
        # afterwards silently reintroduces divergence.
        self.bc.set_boundary_values(self.u, self.v, self.w)
        self.mass_correction = self.bc.enforce_global_mass(self.u, self.v, self.w)
        self.bc.fill_ghosts(self.u, self.v, self.w)

        self._project(dt)

        # Post-projection: ghosts only.
        self.bc.fill_ghosts(self.u, self.v, self.w)

        self.t += dt
        self.step_count += 1
        self.force_history.append(self.force.copy())
        self.time_history.append(self.t)
        return dt

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

    def divergence_norm(self):
        """Max |div(u)| over the interior, scaled to be dimensionless.
        Should sit at round-off. A rising value means the projection is not
        closing and every force number downstream is suspect."""
        g = self.grid
        kernels.divergence(self.u, self.v, self.w, g.dx, g.dy, g.dz, self._div)
        u_ref = float(np.linalg.norm(self.bc.u_inf)) or 1.0
        return float(np.abs(self._div).max() * g.dx / u_ref)

    def ibm_slip(self):
        """Max residual velocity inside the solid, normalised by freestream.
        This is the error the projection reintroduces after direct forcing;
        values above a few percent mean eta is too coarse for the geometry."""
        if not self.has_body:
            return 0.0
        u_ref = float(np.linalg.norm(self.bc.u_inf)) or 1.0
        deep = self.eta_u > 0.9
        if not deep.any():
            return 0.0
        return float(np.abs(self.u[deep] - self.body_velocity[0]).max() / u_ref)

    def filtered_force_history(self):
        """
        Force history with the Nyquist (2*dt) mode removed by a two-point
        boxcar.

        One-shot direct forcing leaves a small slip velocity inside the body
        after the projection, which the next step's forcing has to remove
        again. Combined with the AB2 history coefficient (-1/2) that sets up
        a clean odd-even oscillation: the instantaneous force alternates
        between two values every single step, typically at a few percent
        amplitude.

        This does not bias the mean -- averaging over many steps already
        cancels it, so Cd was never affected. It does wreck two other things:
        the reported force standard deviation, which is otherwise dominated
        by an artefact rather than by real unsteadiness, and the Strouhal
        number, because the FFT locks onto a large spike sitting exactly at
        the Nyquist frequency and returns it instead of the shedding peak.
        That was producing St values of ~57 in place of ~0.16.

        A two-point boxcar has an exact zero at the Nyquist frequency and
        leaves the shedding band essentially untouched (shedding is resolved
        by hundreds of steps per cycle), so it removes the artefact without
        touching the physics.
        """
        f = np.asarray(self.force_history)
        if len(f) < 3:
            return f
        return 0.5 * (f[1:] + f[:-1])

    def coefficients(self, area, u_ref=None, window=None):
        """
        Force coefficients, time-averaged over the last `window` fraction of
        the run (default: last half, which discards the impulsive-start
        transient).

        Returns dict with cd, cl, and the standard deviation of each, because
        for a shedding bluff body the fluctuation amplitude is as much a
        result as the mean and a single number hides whether the wake ever
        settled.
        """
        if u_ref is None:
            u_ref = float(np.linalg.norm(self.bc.u_inf))
        q = 0.5 * self.rho * u_ref**2
        denom = q * area
        if denom <= 0 or not self.force_history:
            return {"cd": 0.0, "cl": 0.0, "cd_std": 0.0, "cl_std": 0.0, "samples": 0}

        f = self.filtered_force_history()
        n = len(f)
        start = int(n * (1.0 - (window if window is not None else 0.5)))
        f = f[start:]
        return {
            "cd": float(f[:, 0].mean() / denom),
            "cl": float(f[:, 1].mean() / denom),
            "cd_std": float(f[:, 0].std() / denom),
            "cl_std": float(f[:, 1].std() / denom),
            "samples": int(len(f)),
            "fx_mean_N": float(f[:, 0].mean()),
            "fy_mean_N": float(f[:, 1].mean()),
            "fz_mean_N": float(f[:, 2].mean()),
        }

    def cell_fields(self):
        """Cell-centred (u, v, w, p) for output and visualisation."""
        g = self.grid
        uc = np.zeros((g.nx, g.ny, g.nz))
        vc = np.zeros_like(uc)
        wc = np.zeros_like(uc)
        kernels.cell_velocity(
            self.u, self.v, self.w, g.nx, g.ny, g.nz, uc, vc, wc
        )
        pc = self.rho * self.p[NG:-NG, NG:-NG, NG:-NG].copy()
        return uc, vc, wc, pc

    def vorticity_field(self):
        g = self.grid
        ox = np.zeros((g.nx, g.ny, g.nz))
        oy = np.zeros_like(ox)
        oz = np.zeros_like(ox)
        kernels.vorticity(
            self.u, self.v, self.w, g.dx, g.dy, g.dz, g.nx, g.ny, g.nz, ox, oy, oz
        )
        return ox, oy, oz

    def q_field(self):
        g = self.grid
        q = np.zeros((g.nx, g.ny, g.nz))
        kernels.q_criterion(
            self.u, self.v, self.w, g.dx, g.dy, g.dz, g.nx, g.ny, g.nz, q
        )
        return q

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

    def health_check(self, dt, div):
        """
        Return a failure message if the solution is going unstable, else None.

        Waiting for the divergence norm to reach O(1000) is far too late. When
        a central-difference scheme without enough dissipation goes unstable,
        the adaptive timestep collapses first: peak velocity blows up locally,
        the CFL condition shrinks dt to compensate, and simulated time stops
        advancing while the run keeps burning wall clock. The divergence norm
        stays at round-off for thousands of steps after the point of no
        return, because the projection is still doing its job perfectly on a
        field that is already garbage.

        A Model S run failed exactly this way: recoverable at step ~9750
        (17 minutes in), detected at step 27750 (44 minutes in), producing
        nothing. Watching dt is what catches it early.
        """
        if not np.isfinite(dt) or not np.isfinite(div):
            return "timestep or divergence is not finite"
        if self._dt_reference is None:
            self._dt_reference = dt
        elif dt > self._dt_reference:
            # Track the largest healthy dt seen after start-up.
            self._dt_reference = max(self._dt_reference, dt)
        if dt < 1e-3 * self._dt_reference:
            return (
                f"timestep collapsed to {dt:.3e}, {self._dt_reference/dt:.0f}x "
                f"below its healthy value -- the velocity field is blowing up "
                f"locally. At high Reynolds number this usually means there is "
                f"no subgrid model: central differencing supplies no "
                f"dissipation of its own. Set flow.sgs: vreman, or coarsen "
                f"the case, or reduce run.cfl."
            )
        if div > 1e-6:
            return (
                f"divergence norm {div:.3e} is far above round-off; the "
                f"projection is no longer closing"
            )
        return None

    def run(self, t_end=None, steps=None, callback=None, log_every=0, dt=None):
        """
        Advance the solution. Give either t_end or steps.

        callback(solver, step, dt) fires after each step; return False from it
        to stop early.
        """
        if t_end is None and steps is None:
            raise ValueError("give t_end or steps")
        t0 = time.time()
        n = 0
        while True:
            if steps is not None and n >= steps:
                break
            if t_end is not None and self.t >= t_end:
                break
            step_dt = self.step(dt)
            n += 1
            div = 0.0
            if log_every and n % log_every == 0:
                div = self.divergence_norm()
                el = time.time() - t0
                # Report a short running mean, never the instantaneous force.
                # The direct-forcing scheme carries a 2*dt oscillation (see
                # filtered_force_history), and log_every is invariably an even
                # number, so printing the instantaneous value samples the same
                # phase every time and reports one extreme of the oscillation
                # as if it were the converged force -- off by several percent
                # and perfectly steady-looking, which is the worst combination.
                recent = np.asarray(self.force_history[-min(50, len(self.force_history)):])
                fm = recent.mean(axis=0)
                print(
                    f"  step {n:6d}  t={self.t:8.4f}  dt={step_dt:.3e}  "
                    f"div={div:.2e}  F=({fm[0]:+.4g},{fm[1]:+.4g})  "
                    f"{el:6.1f}s",
                    flush=True,
                )
            # Health is checked every step, not only when logging: the whole
            # point is to fail in seconds rather than after the log interval
            # has passed thousands more times.
            problem = self.health_check(step_dt, div)
            if problem:
                raise FloatingPointError(
                    f"Solution diverged at step {n}, t={self.t:.4f}: {problem}"
                )
            if callback is not None and callback(self, n, step_dt) is False:
                break
        return time.time() - t0
