"""
Boundary conditions for the staggered grid.

The central invariant
---------------------
On a MAC grid the velocity components normal to a domain face live *on* that
face. The all-Neumann pressure projection is built so that grad(phi).n = 0
there, which means the projection cannot change those boundary-normal
velocities -- it takes them as given and makes the interior divergence-free
with respect to them.

Everything about this module follows from that:

    boundary node values  ->  must be set BEFORE the projection, then left alone
    ghost cell values     ->  may be refreshed at any time

Overwriting a boundary node after projecting (for example by re-extrapolating
a zero-gradient outflow from an interior value the projection just changed)
silently reintroduces divergence in the outermost cell layer. The interior
still looks clean, the residual sits four or five orders of magnitude above
round-off, and it shows up as a slow drift in the integrated force rather
than as an obvious failure. Hence the split below into set_boundary_values()
and fill_ghosts(); apply() runs both and is only safe pre-projection.

Face types
----------
inlet       Prescribed velocity vector.
convective  Non-reflecting outflow, du/dt + Uc*du/dx = 0. A plain
            zero-gradient outlet reflects vortices back upstream and
            corrupts the shedding frequency the cylinder benchmark measures.
slip        Symmetry plane: zero normal velocity, zero-gradient tangential.
            Correct for a true symmetry plane, wrong as an open boundary.
farfield    Open boundary: tangential velocity held at freestream, normal
            velocity free to pass. Not interchangeable with slip -- a slip
            lid is impermeable, so a growing boundary layer displaces flow
            that has nowhere to go and the freestream accelerates. On a
            flat plate that blockage pushes u above U_inf and flips the sign
            of the computed momentum thickness. Boundary layers entrain
            fluid from above; the boundary has to let them.
noslip      Stationary wall.
moving      Wall translating in its own plane -- a rolling road. A stationary
            floor under a moving car is a wind-tunnel artefact that thickens
            the floor boundary layer and inflates drag.

Tangential ghosts use v_ghost = 2*v_wall - v_interior, placing the specified
value exactly on the wall at second order. Normal-component ghosts use
g = 2*boundary - first_interior, the mirror about the boundary plane.
"""

import numpy as np

INLET = "inlet"
CONVECTIVE = "convective"
SLIP = "slip"
NOSLIP = "noslip"
MOVING = "moving"
FARFIELD = "farfield"

FACES = ("xlo", "xhi", "ylo", "yhi", "zlo", "zhi")
_VALID = {INLET, CONVECTIVE, SLIP, NOSLIP, MOVING, FARFIELD}

# Per-face: (axis, is_low). axis 0=x, 1=y, 2=z.
_FACE_AXIS = {
    "xlo": (0, True), "xhi": (0, False),
    "ylo": (1, True), "yhi": (1, False),
    "zlo": (2, True), "zhi": (2, False),
}


def _sl(axis, index):
    """Slice tuple selecting one plane at `index` along `axis`."""
    s = [slice(None)] * 3
    s[axis] = index
    return tuple(s)


class BoundaryConditions:
    def __init__(self, grid, spec, u_inf=(1.0, 0.0, 0.0), wall_velocity=None,
                 ylo_wall_range=None):
        """
        ylo_wall_range : (x_start, x_end) in metres, or None.

        Restricts the ylo wall condition to a streamwise interval; outside it
        the face reverts to slip. This exists for the flat-plate case, where
        the leading edge must sit *inside* the domain.

        Running a wall right up to the inlet forces u = U_inf and u = 0 at the
        same corner. That discontinuity radiates a pressure spike which
        convects downstream as a persistent ~13% freestream overshoot -- and
        because the overshoot is strongest near the inlet and decays with x,
        it is easy to misread as a blockage effect and chase with a taller
        domain, which does not fix it. A slip section upstream lets the flow
        approach a genuine leading edge smoothly.
        """
        self.grid = grid
        unknown = set(spec) - set(FACES)
        if unknown:
            raise ValueError(f"Unknown boundary face(s): {sorted(unknown)}")
        missing = set(FACES) - set(spec)
        if missing:
            raise ValueError(f"No condition given for face(s): {sorted(missing)}")
        for face, kind in spec.items():
            if kind not in _VALID:
                raise ValueError(
                    f"Face {face!r}: unknown type {kind!r}. "
                    f"Expected one of {sorted(_VALID)}."
                )
            if kind == CONVECTIVE and face not in ("xlo", "xhi"):
                raise ValueError("convective outflow is only supported on x faces")
        self.spec = dict(spec)

        # A single-cell-deep grid is 2D, and 2D only falls out of the 3D
        # kernels if every z-derivative vanishes. That requires zero-gradient
        # z ghosts, i.e. SLIP. A farfield z face would instead set the ghost
        # to 2*U_inf - u_interior, making the z-Laplacian 4*(U_inf - u)/dz^2
        # -- a spurious body force pulling the whole field toward freestream.
        # It is invisible in a uniform-flow test (where u == U_inf makes it
        # zero) and shows up only as wrong profiles in a real case.
        if grid.nz == 1:
            for f in ("zlo", "zhi"):
                if self.spec[f] != SLIP:
                    self.spec[f] = SLIP

        self.u_inf = np.asarray(u_inf, dtype=float)
        self.wall_velocity = wall_velocity or {}
        self.u_convect = float(np.linalg.norm(self.u_inf)) or 1.0

        # Precompute per-component masks for a partial ylo wall. Each velocity
        # component samples x at its own node positions, so the masks differ.
        self.ylo_wall_range = ylo_wall_range
        self._ylo_mask_u = None
        self._ylo_mask_w = None
        if ylo_wall_range is not None:
            x0, x1 = ylo_wall_range
            nx, nz = grid.nx, grid.nz
            xu = grid.origin[0] + (np.arange(nx + 3) - 1) * grid.dx
            xw = grid.origin[0] + (np.arange(nx + 2) - 0.5) * grid.dx
            self._ylo_mask_u = ((xu >= x0) & (xu <= x1))[:, None]
            self._ylo_mask_w = ((xw >= x0) & (xw <= x1))[:, None]
            _ = nz

    # ------------------------------------------------------------------
    # index helpers
    # ------------------------------------------------------------------

    def _n(self, axis):
        return (self.grid.nx, self.grid.ny, self.grid.nz)[axis]

    def _normal_indices(self, axis, is_low):
        """(ghost, boundary, first_interior) along `axis` for the normal
        component, whose array has n+3 entries on that axis."""
        n = self._n(axis)
        if is_low:
            return 0, 1, 2
        return n + 2, n + 1, n

    def _tangential_indices(self, axis, is_low):
        """(ghost, first_interior) along `axis` for a tangential component,
        whose array has n+2 entries on that axis."""
        n = self._n(axis)
        if is_low:
            return 0, 1
        return n + 1, n

    def _wall_v(self, face):
        return np.asarray(
            self.wall_velocity.get(face, (0.0, 0.0, 0.0)), dtype=float
        )

    # ------------------------------------------------------------------
    # stage 1: boundary node values (pre-projection only)
    # ------------------------------------------------------------------

    def set_boundary_values(self, u, v, w):
        """
        Set velocity components that sit exactly on a domain face.

        Must be called before the projection and not after. Convective faces
        are skipped: their boundary node is owned by advance_outflow().
        """
        comps = (u, v, w)
        for face, kind in self.spec.items():
            if kind == CONVECTIVE:
                continue
            axis, is_low = _FACE_AXIS[face]
            normal = comps[axis]
            _g, b, i0 = self._normal_indices(axis, is_low)

            if kind == INLET:
                normal[_sl(axis, b)] = self.u_inf[axis]
            elif kind in (SLIP, NOSLIP):
                normal[_sl(axis, b)] = 0.0
            elif kind == MOVING:
                normal[_sl(axis, b)] = self._wall_v(face)[axis]
            elif kind == FARFIELD:
                # Zero-gradient: extrapolate from the first interior node.
                # Lagged by one step, which is standard and stable.
                normal[_sl(axis, b)] = normal[_sl(axis, i0)]

    # ------------------------------------------------------------------
    # stage 2: ghosts (safe any time, including post-projection)
    # ------------------------------------------------------------------

    def fill_ghosts(self, u, v, w):
        comps = (u, v, w)
        for face, kind in self.spec.items():
            axis, is_low = _FACE_AXIS[face]

            # In 2D the kernels multiply every z term by zfac=0, so the z
            # ghosts are never read. Skipping them is not a micro-optimisation:
            # with nz=1 a z-face plane spans the entire domain, so filling it
            # costs as much as a full field update, and it is a strided
            # (stride-3) write on top of that. Measured at 6 ms/step on a
            # 230k-cell 2D grid -- larger than the convection and diffusion
            # kernels combined.
            if axis == 2 and self.grid.nz == 1:
                continue
            normal = comps[axis]
            g, b, i0 = self._normal_indices(axis, is_low)

            # Normal component: mirror about the boundary plane.
            normal[_sl(axis, g)] = (
                2.0 * normal[_sl(axis, b)] - normal[_sl(axis, i0)]
            )

            tangential_axes = [a for a in range(3) if a != axis]
            tg, ti = self._tangential_indices(axis, is_low)

            if kind == CONVECTIVE:
                # advance_outflow already convected these ghosts; touching
                # them here would double-apply the condition.
                continue

            for ta in tangential_axes:
                field = comps[ta]
                if kind == SLIP:
                    field[_sl(axis, tg)] = field[_sl(axis, ti)]
                elif kind == NOSLIP:
                    field[_sl(axis, tg)] = -field[_sl(axis, ti)]
                elif kind == INLET or kind == FARFIELD:
                    field[_sl(axis, tg)] = (
                        2.0 * self.u_inf[ta] - field[_sl(axis, ti)]
                    )
                elif kind == MOVING:
                    field[_sl(axis, tg)] = (
                        2.0 * self._wall_v(face)[ta] - field[_sl(axis, ti)]
                    )

            if face == "ylo" and self._ylo_mask_u is not None:
                self._apply_ylo_wall_mask(u, w, tg, ti)

    def _apply_ylo_wall_mask(self, u, w, tg, ti):
        """Revert the ylo tangential ghosts to slip outside the wall range."""
        u[:, tg, :] = np.where(
            self._ylo_mask_u, u[:, tg, :], u[:, ti, :]
        )
        w[:, tg, :] = np.where(
            self._ylo_mask_w, w[:, tg, :], w[:, ti, :]
        )

    def apply(self, u, v, w, dt=None):
        """Both stages. Pre-projection only -- use fill_ghosts() after."""
        self.set_boundary_values(u, v, w)
        self.fill_ghosts(u, v, w)

    # ------------------------------------------------------------------
    # outflow
    # ------------------------------------------------------------------

    def advance_outflow(self, u, v, w, dt):
        """
        March the convective outflow one timestep. Owns its face's boundary
        node and tangential ghosts; must run exactly once per step.
        """
        nx = self.grid.nx
        c = self.u_convect * dt / self.grid.dx
        if self.spec.get("xhi") == CONVECTIVE:
            u[nx + 1, :, :] -= c * (u[nx + 1, :, :] - u[nx, :, :])
            v[nx + 1, :, :] -= c * (v[nx + 1, :, :] - v[nx, :, :])
            w[nx + 1, :, :] -= c * (w[nx + 1, :, :] - w[nx, :, :])
            u[nx + 2, :, :] = u[nx + 1, :, :]
        if self.spec.get("xlo") == CONVECTIVE:
            u[1, :, :] -= c * (u[2, :, :] - u[1, :, :])
            u[0, :, :] = u[1, :, :]
            v[0, :, :] = v[1, :, :]
            w[0, :, :] = w[1, :, :]

    # ------------------------------------------------------------------
    # global mass balance
    # ------------------------------------------------------------------

    def enforce_global_mass(self, u, v, w):
        """
        Rescale the outflow so net flux through the domain boundary is zero.

        The all-Neumann Poisson problem is singular and solvable only when its
        RHS has zero mean, which is exactly the statement that as much mass
        leaves as enters. Neither convective outflow nor a far-field boundary
        guarantees that, so the imbalance is dumped uniformly onto the outlet.

        Returns the correction in m/s. If it grows past a few percent of
        freestream the domain is too short or the far field too close.
        """
        g = self.grid
        nx, ny, nz = g.nx, g.ny, g.nz
        iy, iz = slice(1, ny + 1), slice(1, nz + 1)
        ix = slice(1, nx + 1)
        da_x, da_y, da_z = g.dy * g.dz, g.dx * g.dz, g.dx * g.dy

        net = (
            (u[nx + 1, iy, iz].sum() - u[1, iy, iz].sum()) * da_x
            + (v[ix, ny + 1, iz].sum() - v[ix, 1, iz].sum()) * da_y
            + (w[ix, iy, nz + 1].sum() - w[ix, iy, 1].sum()) * da_z
        )
        outlet_area = ny * nz * da_x
        if outlet_area <= 0.0:
            return 0.0
        correction = net / outlet_area
        u[nx + 1, :, :] -= correction
        u[nx + 2, :, :] = u[nx + 1, :, :]
        return float(correction)


# ----------------------------------------------------------------------
# Presets
# ----------------------------------------------------------------------

def external_flow(grid, u_inf, ground=False, road_speed=None, sides=FARFIELD,
                  spanwise=None, ylo_wall_range=None):
    """
    Standard external-aerodynamics setup: inlet upstream, convective outflow
    downstream, open far field elsewhere.

    ground=True replaces the ylo far field with a wall. Pass road_speed for a
    rolling road matching the freestream -- the correct ground-vehicle
    condition. Leave it None for a fixed floor.

    sides defaults to farfield, not slip. Slip is only correct on a true
    symmetry plane; as an open boundary it blocks entrainment and accelerates
    the freestream. Pass sides="slip" deliberately, e.g. for a half-car model.

    spanwise overrides the z faces alone -- useful for a spanwise symmetry
    plane with an open top.
    """
    if sides not in (SLIP, FARFIELD):
        raise ValueError(f"sides must be 'slip' or 'farfield', got {sides!r}")
    z_kind = spanwise if spanwise is not None else sides

    if ground:
        y_lo = MOVING if road_speed is not None else NOSLIP
    else:
        y_lo = sides

    spec = {
        "xlo": INLET,
        "xhi": CONVECTIVE,
        "ylo": y_lo,
        "yhi": sides,
        "zlo": z_kind,
        "zhi": z_kind,
    }
    wall_velocity = {}
    if ground and road_speed is not None:
        wall_velocity["ylo"] = (float(road_speed), 0.0, 0.0)
    return BoundaryConditions(
        grid, spec, u_inf=u_inf, wall_velocity=wall_velocity,
        ylo_wall_range=ylo_wall_range,
    )
