"""
Uniform staggered (MAC) Cartesian grid.

Layout, with one ghost layer (ng=1) on every side:

    p[i,j,k]   cell centre,  x = (i-0.5)*dx,  i in [1, nx]
    u[i,j,k]   x-face,       x = (i-1)*dx,    i in [1, nx+1]
    v[i,j,k]   y-face,       y = (j-1)*dy,    j in [1, ny+1]
    w[i,j,k]   z-face,       z = (k-1)*dz,    k in [1, nz+1]

Axis convention for external aerodynamics:
    x = streamwise (inlet -> outlet)
    y = vertical   (ground plane at y=0, lift acts along +y)
    z = spanwise

Setting nz=1 collapses the solver to 2D exactly: the z Poisson eigenvalue
becomes 0 and the zero-gradient z ghost cells make every z-derivative vanish.
No separate 2D code path exists, so the 2D benchmarks exercise the same
kernels as the 3D car runs.
"""

import numpy as np

NG = 1  # ghost layers


class Grid:
    def __init__(self, nx, ny, nz, lx, ly, lz, origin=(0.0, 0.0, 0.0)):
        if min(nx, ny, nz) < 1:
            raise ValueError("nx, ny, nz must all be >= 1")
        self.nx, self.ny, self.nz = int(nx), int(ny), int(nz)
        self.lx, self.ly, self.lz = float(lx), float(ly), float(lz)
        self.origin = np.asarray(origin, dtype=float)

        self.dx = self.lx / self.nx
        self.dy = self.ly / self.ny
        # A 1-cell-deep grid is 2D; dz is arbitrary but must stay finite.
        self.dz = self.lz / self.nz if self.nz > 1 else self.lz

        self.two_d = self.nz == 1
        self.cell_volume = self.dx * self.dy * self.dz

    # -- array factories ------------------------------------------------
    # Shapes include ghosts. Face arrays carry one extra entry along their
    # own axis because n cells have n+1 faces.

    def zeros_p(self):
        return np.zeros((self.nx + 2, self.ny + 2, self.nz + 2))

    def zeros_u(self):
        return np.zeros((self.nx + 3, self.ny + 2, self.nz + 2))

    def zeros_v(self):
        return np.zeros((self.nx + 2, self.ny + 3, self.nz + 2))

    def zeros_w(self):
        return np.zeros((self.nx + 2, self.ny + 2, self.nz + 3))

    # -- coordinates ----------------------------------------------------

    def cell_centers(self):
        x = self.origin[0] + (np.arange(self.nx) + 0.5) * self.dx
        y = self.origin[1] + (np.arange(self.ny) + 0.5) * self.dy
        z = self.origin[2] + (np.arange(self.nz) + 0.5) * self.dz
        return x, y, z

    def face_centers_x(self):
        """Coordinates of u nodes (x-faces), interior only."""
        x = self.origin[0] + np.arange(self.nx + 1) * self.dx
        y = self.origin[1] + (np.arange(self.ny) + 0.5) * self.dy
        z = self.origin[2] + (np.arange(self.nz) + 0.5) * self.dz
        return x, y, z

    def face_centers_y(self):
        x = self.origin[0] + (np.arange(self.nx) + 0.5) * self.dx
        y = self.origin[1] + np.arange(self.ny + 1) * self.dy
        z = self.origin[2] + (np.arange(self.nz) + 0.5) * self.dz
        return x, y, z

    def face_centers_z(self):
        x = self.origin[0] + (np.arange(self.nx) + 0.5) * self.dx
        y = self.origin[1] + (np.arange(self.ny) + 0.5) * self.dy
        z = self.origin[2] + np.arange(self.nz + 1) * self.dz
        return x, y, z

    @property
    def n_cells(self):
        return self.nx * self.ny * self.nz

    def __repr__(self):
        return (
            f"Grid({self.nx}x{self.ny}x{self.nz} = {self.n_cells:,} cells, "
            f"L=({self.lx:.3g},{self.ly:.3g},{self.lz:.3g}), "
            f"h=({self.dx:.4g},{self.dy:.4g},{self.dz:.4g})"
            f"{', 2D' if self.two_d else ''})"
        )
