"""
STL -> immersed-boundary solid fraction fields.

The solver never meshes the body. It rasterises the STL onto the background
Cartesian grid as a solid volume fraction eta in [0,1], sampled separately at
each velocity component's own face locations. That is what keeps the "low
poly, fast iteration" workflow intact: geometry changes cost one rasterisation
(seconds), not a remesh.

Rasterisation strategy
----------------------
1. Subdivide triangles until every edge is shorter than the voxel pitch, then
   bin the resulting vertices. This guarantees a watertight voxel shell with
   no leaks, which point-in-triangle tests on cell centres do not -- a thin
   triangle can pass between two sample points and open a hole.
2. Flood-fill the exterior from a domain corner. Solid is everything the fill
   cannot reach. This handles internal cavities correctly and does not require
   consistent face winding.
3. Box-average the fine mask down to the solver grid to get fractional eta.
   Fractional (rather than 0/1) occupancy is what makes drag converge smoothly
   as the body moves relative to the grid instead of jumping in staircase
   steps.

Anti-aliasing note: eta is a first-order representation of the surface. It
does not resolve a boundary layer, so wall shear is under-resolved unless the
grid is locally fine. For bluff bodies -- which is the whole target class
here -- drag is pressure-dominated and this matters much less than it would
for a streamlined body or a flat plate.
"""

import numpy as np
import trimesh
from scipy import ndimage

UNIT_SCALE = {
    "m": 1.0, "meter": 1.0, "meters": 1.0,
    "cm": 0.01, "centimeter": 0.01, "centimeters": 0.01,
    "mm": 0.001, "millimeter": 0.001, "millimeters": 0.001,
    "in": 0.0254, "inch": 0.0254, "inches": 0.0254,
    "ft": 0.3048, "foot": 0.3048, "feet": 0.3048,
}


def unit_scale(units):
    if units is None:
        raise ValueError("units must be given; SI forces cannot be computed without it")
    s = UNIT_SCALE.get(str(units).lower())
    if s is None:
        raise ValueError(
            f"Unsupported units {units!r}. Use one of {sorted(set(UNIT_SCALE))}."
        )
    return s


# ----------------------------------------------------------------------
# Loading and placement
# ----------------------------------------------------------------------

def load_mesh(path, units="m", scale=1.0):
    """Load an STL and convert to metres. `scale` is an extra multiplier for
    model-scale geometry (e.g. 0.25 for a quarter-scale tunnel model)."""
    mesh = trimesh.load(path, force="mesh")
    if mesh.is_empty or len(mesh.faces) == 0:
        raise ValueError(f"{path}: mesh is empty")
    mesh = mesh.copy()
    mesh.remove_unreferenced_vertices()
    mesh.apply_scale(unit_scale(units) * float(scale))
    return mesh


def orient_mesh(mesh, axis_order="xyz", rotate_deg=(0.0, 0.0, 0.0), flip=(False, False, False)):
    """
    Reorient geometry into solver axes (x streamwise, y up, z spanwise).

    axis_order is a permutation string saying which source axis supplies each
    solver axis: "zyx" means the STL's z becomes the solver's x. Rotations are
    applied after permutation, in x,y,z order, about the mesh centroid.
    """
    mesh = mesh.copy()
    axis_order = axis_order.lower()
    if sorted(axis_order) != ["x", "y", "z"]:
        raise ValueError(f"axis_order must be a permutation of xyz, got {axis_order!r}")

    perm = ["xyz".index(c) for c in axis_order]
    verts = mesh.vertices[:, perm]
    for i, f in enumerate(flip):
        if f:
            verts[:, i] = -verts[:, i]
    mesh.vertices = verts
    # A permutation or flip can invert handedness; fixing winding keeps
    # normals outward, which downstream force reporting depends on.
    mesh.fix_normals()

    center = mesh.bounds.mean(axis=0)
    for axis, deg in enumerate(rotate_deg):
        if abs(deg) > 1e-12:
            direction = np.zeros(3)
            direction[axis] = 1.0
            mesh.apply_transform(
                trimesh.transformations.rotation_matrix(
                    np.radians(deg), direction, center
                )
            )
    return mesh


def place_in_domain(mesh, grid, ground_clearance=None, x_fraction=0.3, center_span=True):
    """
    Translate the mesh into the domain.

    x_fraction sets where the body's nose sits as a fraction of domain length;
    0.3 leaves 30% of the domain upstream, which is enough for the inlet not
    to feel the body, and 70% downstream for the wake to develop before it
    meets the outlet.

    ground_clearance, if given, is the gap in metres between the lowest point
    of the body and the ylo plane. Otherwise the body is centred vertically.
    """
    mesh = mesh.copy()
    lo, hi = mesh.bounds
    size = hi - lo
    o = grid.origin

    tx = o[0] + x_fraction * grid.lx - lo[0]

    if ground_clearance is None:
        ty = o[1] + 0.5 * (grid.ly - size[1]) - lo[1]
    else:
        ty = o[1] + ground_clearance - lo[1]

    if center_span:
        tz = o[2] + 0.5 * (grid.lz - size[2]) - lo[2]
    else:
        tz = o[2] - lo[2]

    mesh.apply_translation([tx, ty, tz])
    return mesh


# ----------------------------------------------------------------------
# Rasterisation
# ----------------------------------------------------------------------

def _surface_voxels(mesh, origin, pitch, shape):
    """Indices of fine voxels the surface passes through."""
    # Subdividing to below the pitch guarantees at least one sample vertex per
    # voxel the surface crosses -- this is what makes the shell leak-free.
    v, _f = trimesh.remesh.subdivide_to_size(
        mesh.vertices, mesh.faces, max_edge=pitch * 0.5
    )
    idx = np.floor((v - origin) / pitch).astype(np.int64)
    np.clip(idx, 0, np.array(shape) - 1, out=idx)
    return idx


def solid_mask_fine(mesh, grid, refine=2):
    """
    Boolean solid occupancy on a grid `refine` times finer than the solver
    grid in each direction.
    """
    shape = (grid.nx * refine, grid.ny * refine, max(grid.nz * refine, 1))
    pitch = np.array([grid.dx / refine, grid.dy / refine, grid.dz / refine])
    if grid.two_d:
        # In 2D the geometry is extruded through the single cell layer, so
        # occupancy is decided purely in the x-y plane.
        shape = (grid.nx * refine, grid.ny * refine, 1)

    occupied = np.zeros(shape, dtype=bool)

    # Uniform pitch is required by the index arithmetic below; grids are
    # uniform by construction but the pitch may differ per axis, so index
    # each axis independently.
    v, _ = trimesh.remesh.subdivide_to_size(
        mesh.vertices, mesh.faces, max_edge=float(pitch[:2].min()) * 0.5
    )
    rel = (v - grid.origin) / pitch
    idx = np.floor(rel).astype(np.int64)
    if grid.two_d:
        idx[:, 2] = 0
    inside = np.all((idx >= 0) & (idx < np.array(shape)), axis=1)
    idx = idx[inside]
    occupied[idx[:, 0], idx[:, 1], idx[:, 2]] = True

    # Flood-fill the exterior. Padding by one guarantees a connected outside
    # even when the body touches a domain face (a car on the ground does).
    padded = np.pad(occupied, 1, mode="constant", constant_values=False)
    free = ~padded
    labels, _n = ndimage.label(free)
    outside_label = labels[0, 0, 0]
    if outside_label == 0:
        # Corner is on the shell -- degenerate, treat everything as fluid
        # rather than silently filling the domain with solid.
        return np.zeros(shape, dtype=bool)
    exterior = labels == outside_label
    solid = ~exterior[1:-1, 1:-1, 1:-1]
    return solid


def _boxdown(fine, refine, shape):
    """Average a fine boolean mask down to `shape` -> fractional occupancy."""
    nx, ny, nz = shape
    rz = refine if fine.shape[2] > 1 else 1
    trimmed = fine[: nx * refine, : ny * refine, : nz * rz]
    return (
        trimmed.reshape(nx, refine, ny, refine, nz, rz)
        .mean(axis=(1, 3, 5))
        .astype(np.float64)
    )


def solid_fractions(grid, mesh, refine=2, return_area=False):
    """
    Returns (eta_u, eta_v, eta_w, eta_c) as ghosted arrays matching the MAC
    layout. Ghost entries are zero so boundary nodes are never force-driven.

    Face fractions come from averaging the two adjacent cell fractions, which
    is the correct second-order estimate of occupancy at the face plane.

    With return_area=True, also returns the frontal area of the rasterised
    body -- free, since the mask is already built.
    """
    fine = solid_mask_fine(mesh, grid, refine=refine)
    cell = _boxdown(fine, refine, (grid.nx, grid.ny, grid.nz))

    eta_c = grid.zeros_p()
    eta_c[1:-1, 1:-1, 1:-1] = cell

    eta_u = grid.zeros_u()
    # u node i sits between cells i-1 and i
    eta_u[2:grid.nx + 1, 1:-1, 1:-1] = 0.5 * (cell[:-1] + cell[1:])

    eta_v = grid.zeros_v()
    eta_v[1:-1, 2:grid.ny + 1, 1:-1] = 0.5 * (cell[:, :-1] + cell[:, 1:])

    eta_w = grid.zeros_w()
    if grid.nz > 1:
        eta_w[1:-1, 1:-1, 2:grid.nz + 1] = 0.5 * (cell[:, :, :-1] + cell[:, :, 1:])

    if return_area:
        return eta_u, eta_v, eta_w, eta_c, frontal_area_from_mask(
            fine, grid, refine, axis=0
        )
    return eta_u, eta_v, eta_w, eta_c


# ----------------------------------------------------------------------
# Reference quantities
# ----------------------------------------------------------------------

def projected_frontal_area(mesh, axis=0, resolution=384):
    """
    Projected area normal to `axis`, by rasterising the silhouette.

    This replaces the bounding-box product used previously. For a car the
    bounding box overstates frontal area by roughly 15-25%, and Cd divides by
    this number, so the error lands directly on the reported coefficient.

    `resolution` is the pixel count along the longer projected axis and is
    deliberately modest. The rasteriser works by subdividing triangles until
    every edge is below the pixel pitch, so cost grows as resolution^2 in the
    subdivided triangle count: at 2048 pixels a 300k-face car mesh expands
    past tens of millions of triangles and the "pre-flight check" runs longer
    than the simulation it is supposed to be checking. At 384 the area is
    converged to well under a percent, which is far finer than the solver's
    own grid resolves the body anyway.

    Prefer frontal_area_from_mask() when a rasterised solid mask already
    exists -- it is free and it measures the body the solver actually sees.
    """
    other = [a for a in range(3) if a != axis]
    pts = mesh.vertices[:, other]

    lo = pts.min(axis=0)
    hi = pts.max(axis=0)
    span = hi - lo
    if np.any(span <= 0):
        return 0.0

    pitch = span.max() / resolution
    shape = np.maximum(np.ceil(span / pitch).astype(int), 1)

    v2, _f2 = trimesh.remesh.subdivide_to_size(
        mesh.vertices, mesh.faces, max_edge=pitch * 0.9
    )
    p = v2[:, other]
    idx = np.floor((p - lo) / pitch).astype(np.int64)
    np.clip(idx, 0, shape - 1, out=idx)

    img = np.zeros(shape, dtype=bool)
    img[idx[:, 0], idx[:, 1]] = True
    img = ndimage.binary_fill_holes(img)
    return float(img.sum() * pitch * pitch)


def frontal_area_from_mask(fine_mask, grid, refine, axis=0):
    """
    Frontal area from an already-rasterised solid mask.

    This is the reference area the solver's immersed body actually presents
    to the flow, including the staircase error. Using it rather than the
    exact CAD silhouette keeps Cd self-consistent: numerator and denominator
    then refer to the same object, so a grid-refinement study converges
    cleanly instead of mixing two different bodies.
    """
    silhouette = fine_mask.any(axis=axis)
    cell = [grid.dx / refine, grid.dy / refine, grid.dz / refine]
    if grid.two_d:
        cell[2] = grid.dz
    area_per_pixel = np.prod([c for a, c in enumerate(cell) if a != axis])
    return float(silhouette.sum() * area_per_pixel)


def mesh_report(mesh, name="mesh", frontal_area=True):
    """Diagnostics printed before a run so geometry problems surface early."""
    lo, hi = mesh.bounds
    size = hi - lo
    return {
        "name": name,
        "faces": int(len(mesh.faces)),
        "vertices": int(len(mesh.vertices)),
        "watertight": bool(mesh.is_watertight),
        "winding_consistent": bool(mesh.is_winding_consistent),
        "bounds_min_m": [float(x) for x in lo],
        "bounds_max_m": [float(x) for x in hi],
        "size_m": [float(x) for x in size],
        "volume_m3": float(mesh.volume) if mesh.is_watertight else None,
        "frontal_area_m2": (
            projected_frontal_area(mesh, axis=0) if frontal_area else None
        ),
    }
