"""
Analytic solid-fraction fields for canonical shapes.

Used by the verification benchmarks. Going through an STL for a cylinder
would fold the rasteriser's error into the physics result and make a failure
ambiguous; here the geometry is exact to the level of the supersampling, so a
wrong drag coefficient can only be the solver's fault.

Fractions are computed by supersampling a signed-distance function over the
cell surrounding each velocity node. Fractional rather than 0/1 occupancy is
what lets drag converge smoothly as the body moves relative to the grid --
with a binary mask, Cd jumps in visible steps whenever a surface crosses a
cell boundary, and grid-refinement studies become uninterpretable.
"""

import numpy as np


def _fill(node_coords, sdf, spacing, samples=4):
    """
    Occupancy fraction at each node, by averaging sign(sdf) over a
    supersampled block of the cell centred on that node.

    node_coords : tuple of 1D coordinate arrays (x, y, z) for the node grid
    sdf         : callable (X, Y, Z) -> signed distance, negative inside
    spacing     : (dx, dy, dz) of the cell owned by each node
    """
    x, y, z = node_coords
    dx, dy, dz = spacing
    offs = (np.arange(samples) + 0.5) / samples - 0.5

    acc = np.zeros((x.size, y.size, z.size))
    n_sub = 0
    for ox in offs:
        for oy in offs:
            zoffs = offs if z.size > 1 else [0.0]
            for oz in zoffs:
                X, Y, Z = np.meshgrid(
                    x + ox * dx, y + oy * dy, z + oz * dz, indexing="ij"
                )
                acc += (sdf(X, Y, Z) < 0.0).astype(float)
                n_sub += 1
    return acc / n_sub


def _mac_node_coords(grid):
    """Interior node coordinates for each MAC component."""
    ox, oy, oz = grid.origin
    xc = ox + (np.arange(grid.nx) + 0.5) * grid.dx
    yc = oy + (np.arange(grid.ny) + 0.5) * grid.dy
    zc = oz + (np.arange(grid.nz) + 0.5) * grid.dz
    xf = ox + np.arange(grid.nx + 1) * grid.dx
    yf = oy + np.arange(grid.ny + 1) * grid.dy
    zf = oz + np.arange(grid.nz + 1) * grid.dz
    return (xc, yc, zc), (xf, yf, zf)


def fractions_from_sdf(grid, sdf, samples=4):
    """Returns (eta_u, eta_v, eta_w, eta_c) as ghosted MAC arrays."""
    (xc, yc, zc), (xf, yf, zf) = _mac_node_coords(grid)
    h = (grid.dx, grid.dy, grid.dz)

    eta_c = grid.zeros_p()
    eta_c[1:-1, 1:-1, 1:-1] = _fill((xc, yc, zc), sdf, h, samples)

    eta_u = grid.zeros_u()
    eta_u[1:grid.nx + 2, 1:-1, 1:-1] = _fill((xf, yc, zc), sdf, h, samples)

    eta_v = grid.zeros_v()
    eta_v[1:-1, 1:grid.ny + 2, 1:-1] = _fill((xc, yf, zc), sdf, h, samples)

    eta_w = grid.zeros_w()
    if grid.nz > 1:
        eta_w[1:-1, 1:-1, 1:grid.nz + 2] = _fill((xc, yc, zf), sdf, h, samples)

    return eta_u, eta_v, eta_w, eta_c


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

def circle(grid, center, diameter, samples=4):
    """Infinite circular cylinder with its axis along z."""
    cx, cy = center[0], center[1]
    r = 0.5 * diameter

    def sdf(X, Y, _Z):
        return np.sqrt((X - cx) ** 2 + (Y - cy) ** 2) - r

    return fractions_from_sdf(grid, sdf, samples)


def sphere(grid, center, diameter, samples=4):
    cx, cy, cz = center
    r = 0.5 * diameter

    def sdf(X, Y, Z):
        return np.sqrt((X - cx) ** 2 + (Y - cy) ** 2 + (Z - cz) ** 2) - r

    return fractions_from_sdf(grid, sdf, samples)


def box(grid, center, size, samples=4):
    """Axis-aligned rectangular block. size is the full extent per axis."""
    c = np.asarray(center, dtype=float)
    hs = 0.5 * np.asarray(size, dtype=float)

    def sdf(X, Y, Z):
        dx = np.abs(X - c[0]) - hs[0]
        dy = np.abs(Y - c[1]) - hs[1]
        dz = np.abs(Z - c[2]) - hs[2]
        outside = np.sqrt(
            np.maximum(dx, 0) ** 2 + np.maximum(dy, 0) ** 2 + np.maximum(dz, 0) ** 2
        )
        inside = np.minimum(np.maximum(np.maximum(dx, dy), dz), 0.0)
        return outside + inside

    return fractions_from_sdf(grid, sdf, samples)


def square_cylinder(grid, center, side, samples=4):
    """Square cross-section bluff body, axis along z."""
    cx, cy = center[0], center[1]
    h = 0.5 * side

    def sdf(X, Y, _Z):
        dx = np.abs(X - cx) - h
        dy = np.abs(Y - cy) - h
        outside = np.sqrt(np.maximum(dx, 0) ** 2 + np.maximum(dy, 0) ** 2)
        inside = np.minimum(np.maximum(dx, dy), 0.0)
        return outside + inside

    return fractions_from_sdf(grid, sdf, samples)
