"""
Subgrid-scale model (Vreman) for high-Reynolds runs.

Why this exists even though the benchmarks do not use it
-------------------------------------------------------
The verification cases run at Re of a few hundred, fully resolved, and any
subgrid model there would only add error -- they run laminar. A road vehicle
at Re ~ 10^7 on a grid of a few million cells is a different situation: the
resolved scales stop well short of the dissipative range, and *something* has
to remove energy at the grid cutoff. Without a model that something is the
truncation error of the discretisation, which is implicit LES: it works, but
the effective dissipation is an accident of the scheme rather than a stated
physical assumption, and it cannot be reported or refined.

Vreman rather than Smagorinsky
------------------------------
Constant-coefficient Smagorinsky produces spurious eddy viscosity in laminar
and transitional regions and does not vanish at a wall, so it damps the
attached boundary layer it should leave alone. The Vreman model is built so
its operator vanishes for any flow whose velocity gradient tensor has the
structure of a simple shear or of pure rotation, which covers laminar shear
layers and near-wall flow. It costs about the same, needs no test filter, and
needs no dynamic procedure or averaging direction -- which matters here
because a car has no homogeneous direction to average over.

Known approximation
-------------------
The subgrid stress is applied as div(nu_t grad u) rather than the full
div(nu_t (grad u + grad u^T)). The transpose part is exactly zero for
constant nu_t (incompressibility) and small where nu_t varies smoothly; it is
dropped to keep this to one extra pass over the grid. This is a real
approximation and it is the first thing to revisit if wall-bounded results
look off.
"""

import numpy as np
from numba import njit, prange

FASTMATH = True

# c = 2.5 * Cs^2 with Cs = 0.17, the value Vreman derives and validates.
VREMAN_C = 0.07


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def vreman_nu_t(u, v, w, dx, dy, dz, nx, ny, nz, c, nu_t):
    """Eddy viscosity at cell centres. nu_t is ghosted (nx+2, ny+2, nz+2)."""
    for i in prange(nx):
        ii = i + 1
        for j in range(ny):
            jj = j + 1
            for k in range(nz):
                kk = k + 1

                # alpha[m][n] = d(u_n)/d(x_m), at the cell centre
                a00 = (u[ii + 1, jj, kk] - u[ii, jj, kk]) / dx
                a11 = (v[ii, jj + 1, kk] - v[ii, jj, kk]) / dy
                a22 = (w[ii, jj, kk + 1] - w[ii, jj, kk]) / dz

                ucn = 0.5 * (u[ii, jj + 1, kk] + u[ii + 1, jj + 1, kk])
                ucs = 0.5 * (u[ii, jj - 1, kk] + u[ii + 1, jj - 1, kk])
                a10 = (ucn - ucs) / (2.0 * dy)
                uct = 0.5 * (u[ii, jj, kk + 1] + u[ii + 1, jj, kk + 1])
                ucb = 0.5 * (u[ii, jj, kk - 1] + u[ii + 1, jj, kk - 1])
                a20 = (uct - ucb) / (2.0 * dz)

                vce = 0.5 * (v[ii + 1, jj, kk] + v[ii + 1, jj + 1, kk])
                vcw = 0.5 * (v[ii - 1, jj, kk] + v[ii - 1, jj + 1, kk])
                a01 = (vce - vcw) / (2.0 * dx)
                vct = 0.5 * (v[ii, jj, kk + 1] + v[ii, jj + 1, kk + 1])
                vcb = 0.5 * (v[ii, jj, kk - 1] + v[ii, jj + 1, kk - 1])
                a21 = (vct - vcb) / (2.0 * dz)

                wce = 0.5 * (w[ii + 1, jj, kk] + w[ii + 1, jj, kk + 1])
                wcw = 0.5 * (w[ii - 1, jj, kk] + w[ii - 1, jj, kk + 1])
                a02 = (wce - wcw) / (2.0 * dx)
                wcn = 0.5 * (w[ii, jj + 1, kk] + w[ii, jj + 1, kk + 1])
                wcs = 0.5 * (w[ii, jj - 1, kk] + w[ii, jj - 1, kk + 1])
                a12 = (wcn - wcs) / (2.0 * dy)

                aa = (
                    a00 * a00 + a01 * a01 + a02 * a02
                    + a10 * a10 + a11 * a11 + a12 * a12
                    + a20 * a20 + a21 * a21 + a22 * a22
                )
                if aa < 1e-30:
                    nu_t[ii, jj, kk] = 0.0
                    continue

                d0, d1, d2 = dx * dx, dy * dy, dz * dz
                b00 = d0 * a00 * a00 + d1 * a10 * a10 + d2 * a20 * a20
                b11 = d0 * a01 * a01 + d1 * a11 * a11 + d2 * a21 * a21
                b22 = d0 * a02 * a02 + d1 * a12 * a12 + d2 * a22 * a22
                b01 = d0 * a00 * a01 + d1 * a10 * a11 + d2 * a20 * a21
                b02 = d0 * a00 * a02 + d1 * a10 * a12 + d2 * a20 * a22
                b12 = d0 * a01 * a02 + d1 * a11 * a12 + d2 * a21 * a22

                bb = (
                    b00 * b11 - b01 * b01
                    + b00 * b22 - b02 * b02
                    + b11 * b22 - b12 * b12
                )
                # bb is non-negative analytically; round-off can make it
                # slightly negative in near-uniform flow, where the model
                # should be off anyway.
                nu_t[ii, jj, kk] = c * np.sqrt(bb / aa) if bb > 0.0 else 0.0


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def add_sgs_u(u, nu_t, dx, dy, dz, nx, ny, nz, zfac, out):
    """out += div(nu_t grad u) 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]
                ne = nu_t[ii, jj, kk]
                nw = nu_t[ii - 1, jj, kk]
                tx = (
                    ne * (u[ii + 1, jj, kk] - uc)
                    - nw * (uc - u[ii - 1, jj, kk])
                ) / (dx * dx)

                nn = 0.25 * (
                    nu_t[ii - 1, jj, kk] + nu_t[ii, jj, kk]
                    + nu_t[ii - 1, jj + 1, kk] + nu_t[ii, jj + 1, kk]
                )
                ns = 0.25 * (
                    nu_t[ii - 1, jj - 1, kk] + nu_t[ii, jj - 1, kk]
                    + nu_t[ii - 1, jj, kk] + nu_t[ii, jj, kk]
                )
                ty = (
                    nn * (u[ii, jj + 1, kk] - uc)
                    - ns * (uc - u[ii, jj - 1, kk])
                ) / (dy * dy)

                nt = 0.25 * (
                    nu_t[ii - 1, jj, kk] + nu_t[ii, jj, kk]
                    + nu_t[ii - 1, jj, kk + 1] + nu_t[ii, jj, kk + 1]
                )
                nb = 0.25 * (
                    nu_t[ii - 1, jj, kk - 1] + nu_t[ii, jj, kk - 1]
                    + nu_t[ii - 1, jj, kk] + nu_t[ii, jj, kk]
                )
                tz = (
                    nt * (u[ii, jj, kk + 1] - uc)
                    - nb * (uc - u[ii, jj, kk - 1])
                ) / (dz * dz)

                out[ii, jj, kk] += tx + ty + zfac * tz


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def add_sgs_v(v, nu_t, dx, dy, dz, nx, ny, nz, zfac, out):
    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]
                nn = nu_t[ii, jj, kk]
                ns = nu_t[ii, jj - 1, kk]
                ty = (
                    nn * (v[ii, jj + 1, kk] - vc)
                    - ns * (vc - v[ii, jj - 1, kk])
                ) / (dy * dy)

                ne = 0.25 * (
                    nu_t[ii, jj - 1, kk] + nu_t[ii, jj, kk]
                    + nu_t[ii + 1, jj - 1, kk] + nu_t[ii + 1, jj, kk]
                )
                nw = 0.25 * (
                    nu_t[ii - 1, jj - 1, kk] + nu_t[ii - 1, jj, kk]
                    + nu_t[ii, jj - 1, kk] + nu_t[ii, jj, kk]
                )
                tx = (
                    ne * (v[ii + 1, jj, kk] - vc)
                    - nw * (vc - v[ii - 1, jj, kk])
                ) / (dx * dx)

                nt = 0.25 * (
                    nu_t[ii, jj - 1, kk] + nu_t[ii, jj, kk]
                    + nu_t[ii, jj - 1, kk + 1] + nu_t[ii, jj, kk + 1]
                )
                nb = 0.25 * (
                    nu_t[ii, jj - 1, kk - 1] + nu_t[ii, jj, kk - 1]
                    + nu_t[ii, jj - 1, kk] + nu_t[ii, jj, kk]
                )
                tz = (
                    nt * (v[ii, jj, kk + 1] - vc)
                    - nb * (vc - v[ii, jj, kk - 1])
                ) / (dz * dz)

                out[ii, jj, kk] += tx + ty + zfac * tz


@njit(parallel=True, fastmath=FASTMATH, cache=True)
def add_sgs_w(w, nu_t, dx, dy, dz, nx, ny, nz, out):
    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]
                nt = nu_t[ii, jj, kk]
                nb = nu_t[ii, jj, kk - 1]
                tz = (
                    nt * (w[ii, jj, kk + 1] - wc)
                    - nb * (wc - w[ii, jj, kk - 1])
                ) / (dz * dz)

                ne = 0.25 * (
                    nu_t[ii, jj, kk - 1] + nu_t[ii, jj, kk]
                    + nu_t[ii + 1, jj, kk - 1] + nu_t[ii + 1, jj, kk]
                )
                nw = 0.25 * (
                    nu_t[ii - 1, jj, kk - 1] + nu_t[ii - 1, jj, kk]
                    + nu_t[ii, jj, kk - 1] + nu_t[ii, jj, kk]
                )
                tx = (
                    ne * (w[ii + 1, jj, kk] - wc)
                    - nw * (wc - w[ii - 1, jj, kk])
                ) / (dx * dx)

                nn = 0.25 * (
                    nu_t[ii, jj, kk - 1] + nu_t[ii, jj, kk]
                    + nu_t[ii, jj + 1, kk - 1] + nu_t[ii, jj + 1, kk]
                )
                ns = 0.25 * (
                    nu_t[ii, jj - 1, kk - 1] + nu_t[ii, jj - 1, kk]
                    + nu_t[ii, jj, kk - 1] + nu_t[ii, jj, kk]
                )
                ty = (
                    nn * (w[ii, jj + 1, kk] - wc)
                    - ns * (wc - w[ii, jj - 1, kk])
                ) / (dy * dy)

                out[ii, jj, kk] += tx + ty + tz


class VremanSGS:
    """Plugged into Solver(sgs=...). Called once per RHS evaluation."""

    def __init__(self, c=VREMAN_C, damp_in_solid=True):
        self.c = float(c)
        self.damp_in_solid = damp_in_solid
        self.nu_t = None

    def add(self, solver):
        g = solver.grid
        if self.nu_t is None:
            self.nu_t = g.zeros_p()
        zfac = 0.0 if g.two_d else 1.0

        vreman_nu_t(
            solver.u, solver.v, solver.w,
            g.dx, g.dy, g.dz, g.nx, g.ny, g.nz, self.c, self.nu_t,
        )
        if self.damp_in_solid:
            # Velocity gradients across the immersed surface are an artefact
            # of the forcing, not turbulence. Left alone they generate a
            # large eddy viscosity in a thin shell around the body, which
            # thickens the effective boundary layer and inflates drag.
            self.nu_t *= (1.0 - solver.eta_c)

        add_sgs_u(solver.u, self.nu_t, g.dx, g.dy, g.dz,
                  g.nx, g.ny, g.nz, zfac, solver.ru)
        add_sgs_v(solver.v, self.nu_t, g.dx, g.dy, g.dz,
                  g.nx, g.ny, g.nz, zfac, solver.rv)
        if not g.two_d:
            add_sgs_w(solver.w, self.nu_t, g.dx, g.dy, g.dz,
                      g.nx, g.ny, g.nz, solver.rw)

    def max_ratio(self, solver):
        """Peak nu_t/nu. Above roughly 100 the model is doing more work than
        the resolved scales and the grid is too coarse for the case."""
        if self.nu_t is None:
            return 0.0
        return float(self.nu_t.max() / solver.nu)
