"""
Numba kernels for the staggered-grid incompressible solver.

Everything here operates on ghosted MAC arrays (see grid.py). Kernels loop
over interior indices only and assume ghosts already hold valid boundary
values, so there are no branches in the inner loops.

Discretisation
--------------
Convection is second-order central in divergence form. Central differencing
is used deliberately rather than an upwind scheme: upwinding adds numerical
dissipation that is indistinguishable from physical viscosity, which would
silently corrupt exactly the quantity the Blasius benchmark measures and
would damp the vortex shedding the cylinder benchmark depends on. Central
differencing is non-dissipative and, on a staggered grid, discretely
conserves kinetic energy in the inviscid limit. It has no stabilising
mechanism of its own, so it relies on the projection step and on adequate
resolution -- an under-resolved central-difference run oscillates visibly
rather than failing quietly, which is the preferable failure mode.

Diffusion is the standard second-order 7-point Laplacian, treated explicitly.
This caps the timestep at dt < h^2/(2*nu*ndim) (see solver.stable_dt). For
the Reynolds numbers of interest the convective CFL limit binds first, so the
explicit treatment costs nothing and avoids a second linear solve per step.

Immersed boundary
-----------------
Direct forcing via a solid volume-fraction field eta in [0,1] defined on each
velocity component's own face locations. The correction driving u toward the
body velocity is accumulated as it is applied, and its negated sum is the
hydrodynamic force on the body by Newton's third law. This gives forces
without any surface-pressure integration, which matters because the immersed
surface does not align with cell faces and a direct pressure integral over a
staircase boundary is badly behaved.
"""

import numpy as np
from numba import njit, prange

FASTMATH = True


# ======================================================================
# DIVERGENCE / GRADIENT
# ======================================================================

@njit(parallel=True, fastmath=FASTMATH, cache=True)
def divergence(u, v, w, dx, dy, dz, out):
    """Cell-centred divergence of the MAC velocity field. out is (nx,ny,nz)."""
    nx, ny, nz = out.shape
    for i in prange(nx):
        ii = i + 1
        for j in range(ny):
            jj = j + 1
            for k in range(nz):
                kk = k + 1
                out[i, j, k] = (
                    (u[ii + 1, jj, kk] - u[ii, jj, kk]) / dx
                    + (v[ii, jj + 1, kk] - v[ii, jj, kk]) / dy
                    + (w[ii, jj, kk + 1] - w[ii, jj, kk]) / dz
                )


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def subtract_gradient(u, v, w, phi, dt, dx, dy, dz, nx, ny, nz):
    """u -= dt * grad(phi). Boundary faces are untouched because the phi
    ghosts are zero-gradient, so the gradient there is identically zero."""
    for i in prange(nx + 1):
        ii = i + 1
        for j in range(ny):
            jj = j + 1
            for k in range(nz):
                kk = k + 1
                u[ii, jj, kk] -= dt * (phi[ii, jj, kk] - phi[ii - 1, jj, kk]) / dx

    for i in prange(nx):
        ii = i + 1
        for j in range(ny + 1):
            jj = j + 1
            for k in range(nz):
                kk = k + 1
                v[ii, jj, kk] -= dt * (phi[ii, jj, kk] - phi[ii, jj - 1, kk]) / dy

    for i in prange(nx):
        ii = i + 1
        for j in range(ny):
            jj = j + 1
            for k in range(nz + 1):
                kk = k + 1
                w[ii, jj, kk] -= dt * (phi[ii, jj, kk] - phi[ii, jj, kk - 1]) / dz


# ======================================================================
# CONVECTION + DIFFUSION  ->  right-hand side of the momentum equation
# ======================================================================

@njit(parallel=True, fastmath=FASTMATH, cache=True)
def rhs_u(u, v, w, nu, dx, dy, dz, nx, ny, nz, zfac, out):
    """d(u)/dt = -div(u*U) + nu*lap(u), evaluated at x-face nodes."""
    for i in prange(nx + 1):
        ii = i + 1
        for j in range(ny):
            jj = j + 1
            for k in range(nz):
                kk = k + 1
                uc = u[ii, jj, kk]

                # d(uu)/dx
                ue = 0.5 * (uc + u[ii + 1, jj, kk])
                uw = 0.5 * (u[ii - 1, jj, kk] + uc)
                cx = (ue * ue - uw * uw) / dx

                # d(uv)/dy  -- v interpolated to the u node's y-faces
                vn = 0.5 * (v[ii - 1, jj + 1, kk] + v[ii, jj + 1, kk])
                vs = 0.5 * (v[ii - 1, jj, kk] + v[ii, jj, kk])
                un = 0.5 * (uc + u[ii, jj + 1, kk])
                us = 0.5 * (u[ii, jj - 1, kk] + uc)
                cy = (vn * un - vs * us) / dy

                # d(uw)/dz  -- vanishes identically when nz == 1
                wt = 0.5 * (w[ii - 1, jj, kk + 1] + w[ii, jj, kk + 1])
                wb = 0.5 * (w[ii - 1, jj, kk] + w[ii, jj, kk])
                ut = 0.5 * (uc + u[ii, jj, kk + 1])
                ub = 0.5 * (u[ii, jj, kk - 1] + uc)
                cz = (wt * ut - wb * ub) / dz

                lap = (
                    (u[ii + 1, jj, kk] - 2.0 * uc + u[ii - 1, jj, kk]) / (dx * dx)
                    + (u[ii, jj + 1, kk] - 2.0 * uc + u[ii, jj - 1, kk]) / (dy * dy)
                    + zfac * (u[ii, jj, kk + 1] - 2.0 * uc + u[ii, jj, kk - 1]) / (dz * dz)
                )
                out[ii, jj, kk] = -(cx + cy + zfac * cz) + nu * lap


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def rhs_v(u, v, w, nu, dx, dy, dz, nx, ny, nz, zfac, out):
    """d(v)/dt at y-face nodes."""
    for i in prange(nx):
        ii = i + 1
        for j in range(ny + 1):
            jj = j + 1
            for k in range(nz):
                kk = k + 1
                vc = v[ii, jj, kk]

                ue = 0.5 * (u[ii + 1, jj - 1, kk] + u[ii + 1, jj, kk])
                uw = 0.5 * (u[ii, jj - 1, kk] + u[ii, jj, kk])
                ve = 0.5 * (vc + v[ii + 1, jj, kk])
                vw = 0.5 * (v[ii - 1, jj, kk] + vc)
                cx = (ue * ve - uw * vw) / dx

                vn = 0.5 * (vc + v[ii, jj + 1, kk])
                vs = 0.5 * (v[ii, jj - 1, kk] + vc)
                cy = (vn * vn - vs * vs) / dy

                wt = 0.5 * (w[ii, jj - 1, kk + 1] + w[ii, jj, kk + 1])
                wb = 0.5 * (w[ii, jj - 1, kk] + w[ii, jj, kk])
                vt = 0.5 * (vc + v[ii, jj, kk + 1])
                vb = 0.5 * (v[ii, jj, kk - 1] + vc)
                cz = (wt * vt - wb * vb) / dz

                lap = (
                    (v[ii + 1, jj, kk] - 2.0 * vc + v[ii - 1, jj, kk]) / (dx * dx)
                    + (v[ii, jj + 1, kk] - 2.0 * vc + v[ii, jj - 1, kk]) / (dy * dy)
                    + zfac * (v[ii, jj, kk + 1] - 2.0 * vc + v[ii, jj, kk - 1]) / (dz * dz)
                )
                out[ii, jj, kk] = -(cx + cy + zfac * cz) + nu * lap


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def rhs_w(u, v, w, nu, dx, dy, dz, nx, ny, nz, zfac, out):
    """d(w)/dt at z-face nodes. Never called in 2D."""
    for i in prange(nx):
        ii = i + 1
        for j in range(ny):
            jj = j + 1
            for k in range(nz + 1):
                kk = k + 1
                wc = w[ii, jj, kk]

                ue = 0.5 * (u[ii + 1, jj, kk - 1] + u[ii + 1, jj, kk])
                uw = 0.5 * (u[ii, jj, kk - 1] + u[ii, jj, kk])
                we = 0.5 * (wc + w[ii + 1, jj, kk])
                ww = 0.5 * (w[ii - 1, jj, kk] + wc)
                cx = (ue * we - uw * ww) / dx

                vn = 0.5 * (v[ii, jj + 1, kk - 1] + v[ii, jj + 1, kk])
                vs = 0.5 * (v[ii, jj, kk - 1] + v[ii, jj, kk])
                wn = 0.5 * (wc + w[ii, jj + 1, kk])
                ws = 0.5 * (w[ii, jj - 1, kk] + wc)
                cy = (vn * wn - vs * ws) / dy

                wt = 0.5 * (wc + w[ii, jj, kk + 1])
                wb = 0.5 * (w[ii, jj, kk - 1] + wc)
                cz = (wt * wt - wb * wb) / dz

                lap = (
                    (w[ii + 1, jj, kk] - 2.0 * wc + w[ii - 1, jj, kk]) / (dx * dx)
                    + (w[ii, jj + 1, kk] - 2.0 * wc + w[ii, jj - 1, kk]) / (dy * dy)
                    + zfac * (w[ii, jj, kk + 1] - 2.0 * wc + w[ii, jj, kk - 1]) / (dz * dz)
                )
                out[ii, jj, kk] = -(cx + cy + zfac * cz) + nu * lap


# ======================================================================
# IMMERSED BOUNDARY (direct forcing + force accumulation)
# ======================================================================

@njit(parallel=True, fastmath=FASTMATH, cache=True)
def apply_ibm(field, eta, target, nx_f, ny_f, nz_f):
    """
    Drive `field` toward `target` wherever solid fraction eta > 0:
        f_new = f + eta * (target - f)

    Returns sum(eta * (target - f)) over all forced nodes. Multiplying that
    sum by rho*cell_volume/dt gives the momentum per unit time added to the
    fluid; its negative is the force component on the body.

    eta is the same shape as field (ghosts included) so no index shifting is
    needed; ghost entries of eta are zero and contribute nothing.
    """
    total = 0.0
    for i in prange(nx_f):
        ii = i + 1
        acc = 0.0
        for j in range(ny_f):
            jj = j + 1
            for k in range(nz_f):
                kk = k + 1
                e = eta[ii, jj, kk]
                if e > 0.0:
                    d = e * (target - field[ii, jj, kk])
                    field[ii, jj, kk] += d
                    acc += d
        total += acc
    return total


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def zero_inside(field, eta, nx_f, ny_f, nz_f, threshold):
    """Hard-zero deep-solid nodes. Used only for output cosmetics, never
    inside the time loop -- forces come from apply_ibm, not from this."""
    for i in prange(nx_f):
        ii = i + 1
        for j in range(ny_f):
            jj = j + 1
            for k in range(nz_f):
                kk = k + 1
                if eta[ii, jj, kk] >= threshold:
                    field[ii, jj, kk] = 0.0


# ======================================================================
# DERIVED FIELDS (post-processing)
# ======================================================================

@njit(parallel=True, fastmath=FASTMATH, cache=True)
def cell_velocity(u, v, w, nx, ny, nz, uc, vc, wc):
    """Interpolate MAC face velocities to cell centres for output."""
    for i in prange(nx):
        ii = i + 1
        for j in range(ny):
            jj = j + 1
            for k in range(nz):
                kk = k + 1
                uc[i, j, k] = 0.5 * (u[ii, jj, kk] + u[ii + 1, jj, kk])
                vc[i, j, k] = 0.5 * (v[ii, jj, kk] + v[ii, jj + 1, kk])
                wc[i, j, k] = 0.5 * (w[ii, jj, kk] + w[ii, jj, kk + 1])


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def vorticity(u, v, w, dx, dy, dz, nx, ny, nz, ox, oy, oz):
    """Cell-centred vorticity. Curl components are formed at edges and
    averaged to the centre, which keeps the stencil compact and second
    order."""
    for i in prange(nx):
        ii = i + 1
        for j in range(ny):
            jj = j + 1
            for k in range(nz):
                kk = k + 1
                dwdy = (
                    0.5 * (w[ii, jj + 1, kk] + w[ii, jj + 1, kk + 1])
                    - 0.5 * (w[ii, jj - 1, kk] + w[ii, jj - 1, kk + 1])
                ) / (2.0 * dy)
                dvdz = (
                    0.5 * (v[ii, jj, kk + 1] + v[ii, jj + 1, kk + 1])
                    - 0.5 * (v[ii, jj, kk - 1] + v[ii, jj + 1, kk - 1])
                ) / (2.0 * dz)
                dudz = (
                    0.5 * (u[ii, jj, kk + 1] + u[ii + 1, jj, kk + 1])
                    - 0.5 * (u[ii, jj, kk - 1] + u[ii + 1, jj, kk - 1])
                ) / (2.0 * dz)
                dwdx = (
                    0.5 * (w[ii + 1, jj, kk] + w[ii + 1, jj, kk + 1])
                    - 0.5 * (w[ii - 1, jj, kk] + w[ii - 1, jj, kk + 1])
                ) / (2.0 * dx)
                dvdx = (
                    0.5 * (v[ii + 1, jj, kk] + v[ii + 1, jj + 1, kk])
                    - 0.5 * (v[ii - 1, jj, kk] + v[ii - 1, jj + 1, kk])
                ) / (2.0 * dx)
                dudy = (
                    0.5 * (u[ii, jj + 1, kk] + u[ii + 1, jj + 1, kk])
                    - 0.5 * (u[ii, jj - 1, kk] + u[ii + 1, jj - 1, kk])
                ) / (2.0 * dy)

                ox[i, j, k] = dwdy - dvdz
                oy[i, j, k] = dudz - dwdx
                oz[i, j, k] = dvdx - dudy


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def q_criterion(u, v, w, dx, dy, dz, nx, ny, nz, q):
    """
    Q = 0.5*(|Omega|^2 - |S|^2), the standard vortex-core identifier.
    Positive Q marks regions where rotation dominates strain, which is what
    the 3D iso-surface visualisation renders as vortex structure.
    """
    for i in prange(nx):
        ii = i + 1
        for j in range(ny):
            jj = j + 1
            for k in range(nz):
                kk = k + 1
                uxp = 0.5 * (u[ii + 1, jj, kk] + u[ii + 2, jj, kk])
                uxm = 0.5 * (u[ii - 1, jj, kk] + u[ii, jj, kk])
                uyp = 0.5 * (u[ii, jj + 1, kk] + u[ii + 1, jj + 1, kk])
                uym = 0.5 * (u[ii, jj - 1, kk] + u[ii + 1, jj - 1, kk])
                uzp = 0.5 * (u[ii, jj, kk + 1] + u[ii + 1, jj, kk + 1])
                uzm = 0.5 * (u[ii, jj, kk - 1] + u[ii + 1, jj, kk - 1])

                vxp = 0.5 * (v[ii + 1, jj, kk] + v[ii + 1, jj + 1, kk])
                vxm = 0.5 * (v[ii - 1, jj, kk] + v[ii - 1, jj + 1, kk])
                vyp = 0.5 * (v[ii, jj + 1, kk] + v[ii, jj + 2, kk])
                vym = 0.5 * (v[ii, jj - 1, kk] + v[ii, jj, kk])
                vzp = 0.5 * (v[ii, jj, kk + 1] + v[ii, jj + 1, kk + 1])
                vzm = 0.5 * (v[ii, jj, kk - 1] + v[ii, jj + 1, kk - 1])

                wxp = 0.5 * (w[ii + 1, jj, kk] + w[ii + 1, jj, kk + 1])
                wxm = 0.5 * (w[ii - 1, jj, kk] + w[ii - 1, jj, kk + 1])
                wyp = 0.5 * (w[ii, jj + 1, kk] + w[ii, jj + 1, kk + 1])
                wym = 0.5 * (w[ii, jj - 1, kk] + w[ii, jj - 1, kk + 1])
                wzp = 0.5 * (w[ii, jj, kk + 1] + w[ii, jj, kk + 2])
                wzm = 0.5 * (w[ii, jj, kk - 1] + w[ii, jj, kk])

                dudx = (uxp - uxm) / (2.0 * dx)
                dudy = (uyp - uym) / (2.0 * dy)
                dudz = (uzp - uzm) / (2.0 * dz)
                dvdx = (vxp - vxm) / (2.0 * dx)
                dvdy = (vyp - vym) / (2.0 * dy)
                dvdz = (vzp - vzm) / (2.0 * dz)
                dwdx = (wxp - wxm) / (2.0 * dx)
                dwdy = (wyp - wym) / (2.0 * dy)
                dwdz = (wzp - wzm) / (2.0 * dz)

                s11, s22, s33 = dudx, dvdy, dwdz
                s12 = 0.5 * (dudy + dvdx)
                s13 = 0.5 * (dudz + dwdx)
                s23 = 0.5 * (dvdz + dwdy)
                o12 = 0.5 * (dudy - dvdx)
                o13 = 0.5 * (dudz - dwdx)
                o23 = 0.5 * (dvdz - dwdy)

                s_sq = (
                    s11 * s11 + s22 * s22 + s33 * s33
                    + 2.0 * (s12 * s12 + s13 * s13 + s23 * s23)
                )
                o_sq = 2.0 * (o12 * o12 + o13 * o13 + o23 * o23)
                q[i, j, k] = 0.5 * (o_sq - s_sq)


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def _max_abs_slabs(field, ni, nj, nk, out):
    """Per-i-slab maxima of |field| over the interior node range.

    The final reduction is deliberately left to the caller. Numba's prange
    only recognises a fixed set of reduction patterns, and a conditional
    `if a > m: m = a` across threads is not one of them -- it silently
    returns a per-thread private value instead of the global maximum. That
    failure mode is nasty: max_abs_velocity returns ~0, stable_dt hands back
    the viscous limit instead of the convective one, and the run blows up
    several steps later with no indication of where the bad number came from.
    Writing one value per slab and reducing in numpy avoids the pattern
    entirely.
    """
    for i in prange(ni):
        local = 0.0
        for j in range(nj):
            for k in range(nk):
                a = abs(field[i + 1, j + 1, k + 1])
                if a > local:
                    local = a
        out[i] = local


def _max_abs(field, ni, nj, nk):
    out = np.zeros(ni)
    _max_abs_slabs(field, ni, nj, nk, out)
    return float(out.max())


def max_abs_velocity(u, v, w, nx, ny, nz):
    """Peak |u|, |v|, |w| over interior nodes, for the CFL controller.

    Each component is reduced over its own node count -- u has nx+1 nodes in
    x but only ny and nz in the transverse directions, and reading past those
    would walk into the ghost layer of a differently-shaped array.
    """
    return (
        _max_abs(u, nx + 1, ny, nz),
        _max_abs(v, nx, ny + 1, nz),
        _max_abs(w, nx, ny, nz + 1),
    )
