#!/usr/bin/env python3
"""Two benchmarks with known answers for the growth plus initial-stress workflow.

Why this matters
----------------
Nothing in the package compares the solver against a problem whose answer is
known independently. numerical_acceptance.json compares meshes against each
other (0.0161 relative difference between the primary and main meshes on the
Control endpoint) and reports a maximum projected residual of 2.67e-8 against a
1e-6 gate. Both are self-consistency: a mesh sequence converges to whatever
equations were implemented, and a small residual only says Newton found the root
of the residual it was given. The workflow being trusted here is not standard --
multiplicative growth with det Fg = 0.725517, an initial stress residualised as
(F - I) . S0, a nominal dead traction on the reference surface, six rigid modes
removed by an augmented constraint -- and it is reported to six significant
figures. A custom workflow needs at least one problem with an answer known from
outside it before that precision means anything.

Case 1 -- exact, closed form, machine precision
-----------------------------------------------
Under uniform isotropic growth Fg = g I with traction-free boundaries, the
equilibrium is the homogeneous deformation F = g I, for any g > 0 and on any
domain. Then Fe = F Fg^-1 = I, Je = 1, ln Je = 0 and

    PK1 = mu F (Fg^-1 Fg^-T) + (lambda ln Je - mu) F^-T
        = (mu / g) I - (mu / g) I = 0

exactly. So the residual of the interpolant of u = (g - 1) X must vanish to
round-off. This is not a convergence statement, it is an identity, and it is a
strong test of the one kinematic assumption everything else rests on: that the
elastic response is driven by F Fg^-1 and not by F, F Fg^-T, or Fg^-1 F. Any of
those substitutions is invisible in a mesh-refinement study and fails here at
the first digit. The case is run at g = 0.725517^(1/3) = 0.898564, the core's
own natural stretch, for P1 and P2.

The same case measures what the residualised prestress does to the natural
state. (F - I) . S0 is not zero at F = g I, so with the pretension switched on
the traction-free growth solution is no longer an equilibrium. The spurious
first Piola stress is exactly (g - 1) S0, which at the shell's own growth
g = 0.725517^(0.35/3) = 0.963257 and S0 = 1000 * 0.05 / 0.12 = 416.667 Pa has
Frobenius norm 21.65 Pa, 2.9 % of the shell's mu = 740.7 Pa. That is reported,
not asserted: it is a consequence of the modelling choice, and case 1 is exact
only for the pretension-free branches.

Case 2 -- reduced-order numerical reference, NOT a closed form
--------------------------------------------------------------
A thick-walled spherical shell, inner radius 5.0 um, outer 7.5 um, no growth, no
prestress, loaded by a nominal (dead) radial traction of 200 Pa on the reference
inner surface. There is no closed-form solution for this energy, so the
reference is obtained by reducing the same continuum problem by spherical
symmetry to a two-point boundary value problem in r(R) and integrating it with
scipy.integrate.solve_bvp to a collocation tolerance of 1e-10:

    d P_R / dR + (2 / R) (P_R - P_theta) = 0
    P_i = dW / dlambda_i = mu lambda_i + (lambda ln J - mu) / lambda_i
    lambda = (dr/dR, r/R, r/R),  J = lambda_1 lambda_2 lambda_3
    P_R(Ri) = -p,  P_R(Ro) = 0

This is an independent discretization of the same energy in one dimension, not
an analytic solution, and the script says so wherever it prints. Two things
guard it: the hand-written dW/dlambda_i and d2W/dlambda_i dlambda_j are checked
against a central difference of W itself (`--case derivatives`, numpy only), and
the reference is re-solved at 1 % of the pressure and compared against the
small-strain Lame solution, which it must reproduce as the nonlinearity
vanishes. For the defaults the reference gives u_r(Ri) = 0.5708 um and
u_r(Ro) = 0.3369 um, that is 11.4 % radial strain and 5.2 % and 9.4 % away from
the Lame values, so the case does exercise the finite-deformation terms rather
than a linear problem in disguise.

The comparison then reports the L2 displacement error and the observed order
under uniform refinement. Expected discretization orders are 2 for P1 and 3 for
P2. One caveat is printed with the results, and it is why the default tolerance
is a loose 5e-2: the mesh is built from straight-sided tetrahedra whose vertices
lie exactly on the two spheres, so the domain itself carries an O(h^2) geometric
error. The volume defect column measures that term directly -- 12.7 %, 3.4 % and
0.86 % on the three default meshes, falling by 4 per level as it should. It sets
the accuracy floor and caps the observed P2 rate near 2 whenever it dominates,
and it is a property of the benchmark geometry, not of the solver. The sharp
instrument here is the observed order, not the absolute number; --base-
subdivisions 2 shifts the whole ladder onto a finer sphere at 4x the elements.

The spherical shell mesh is generated here from a subdivided icosahedron and
radially layered prisms, so this script needs none of the project's gmsh meshes,
only src/nuclear_envelope_fem on the given --project-root.

Usage:
    python3 12_benchmark_thick_sphere.py --project-root /path/to/tree --case all
    python3 12_benchmark_thick_sphere.py --case derivatives    # numpy only
"""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path
from types import SimpleNamespace

import numpy as np

NATURAL_VOLUME_RATIO = 0.725517
SHELL_COUPLING = 0.35
PRETENSION_MN_PER_M = 0.05
SHELL_THICKNESS_UM = 0.12
OUTER_AXES_UM = (7.5, 5.5, 3.5)
IDENTITY_AXIS = (1.0, 0.0, 0.0)

# `_lame_parameters` reads exactly these two attributes, so a stub keeps the
# benchmark independent of whatever else config.MaterialConfig requires.
SHELL_MATERIAL = SimpleNamespace(young_modulus_pa=2000.0, poisson_ratio=0.35)

DEFAULT_GROWTH_TOLERANCE = 1.0e-12
DEFAULT_PRESSURE_TOLERANCE = 5.0e-2
DEFAULT_DERIVATIVE_TOLERANCE = 1.0e-6


def import_nonlinear_solver(project_root: Path):
    """Import nuclear_envelope_fem.nonlinear_solver from the real working tree."""
    if project_root is None:
        raise SystemExit("--project-root is required for cases 1 and 2.")
    for candidate in (project_root / "src", project_root):
        if (candidate / "nuclear_envelope_fem" / "nonlinear_solver.py").exists():
            if str(candidate) not in sys.path:
                sys.path.insert(0, str(candidate))
            from nuclear_envelope_fem import nonlinear_solver  # noqa: E402

            return nonlinear_solver
    raise SystemExit(
        "nuclear_envelope_fem/nonlinear_solver.py not found under "
        f"{project_root} or {project_root / 'src'}. The flat copy in "
        "05_Source_Data cannot be used: it imports .tags and .config relatively "
        "and only loads as part of the package."
    )


def import_assembly_stack():
    """Import scipy and scikit-fem lazily so --case derivatives needs neither."""
    try:
        import scipy.sparse
        import scipy.sparse.linalg  # noqa: F401
        import skfem
    except ImportError as exc:
        raise SystemExit(
            f"{exc.name} is required: this benchmark assembles the project's own "
            "skfem forms and solves the augmented rigid-constraint system. "
            "Install scikit-fem==12.0.2 and scipy==1.18.0 (the versions in "
            "run_manifest.json) or run inside the project environment."
        ) from exc
    return skfem, scipy


def import_solve_bvp():
    """Import scipy.integrate lazily; only the pressure case needs it."""
    try:
        from scipy.integrate import solve_bvp
    except ImportError as exc:
        raise SystemExit(
            "scipy is required for the reduced-order reference solution "
            "(scipy.integrate.solve_bvp). Install scipy==1.18.0 or run inside "
            "the project environment."
        ) from exc
    return solve_bvp


def lame_parameters(young_pa, poisson):
    return (
        young_pa / (2.0 * (1.0 + poisson)),
        young_pa * poisson / ((1.0 + poisson) * (1.0 - 2.0 * poisson)),
    )


def strain_energy_density(stretches, mu, lmbda):
    """W = mu/2 (Fe:Fe - 3) - mu ln Je + lambda/2 (ln Je)^2 in principal stretches.

    Growth is absent from case 2, so Fe = F and the principal stretches are the
    principal stretches of F itself.
    """
    stretches = np.asarray(stretches, dtype=float)
    log_jacobian = np.sum(np.log(stretches), axis=0)
    return (
        0.5 * mu * (np.sum(stretches**2, axis=0) - 3.0)
        - mu * log_jacobian
        + 0.5 * lmbda * log_jacobian**2
    )


def principal_first_piola(stretches, mu, lmbda):
    """P_i = dW / dlambda_i = mu lambda_i + (lambda ln J - mu) / lambda_i."""
    stretches = np.asarray(stretches, dtype=float)
    log_jacobian = np.sum(np.log(stretches), axis=0)
    return mu * stretches + (lmbda * log_jacobian - mu) / stretches


def principal_stretch_tangent(stretches, mu, lmbda):
    """A_ij = d2W / dlambda_i dlambda_j.

    A_ij = mu delta_ij + lambda / (lambda_i lambda_j)
           - delta_ij (lambda ln J - mu) / lambda_i^2
    """
    stretches = np.asarray(stretches, dtype=float)
    log_jacobian = np.sum(np.log(stretches), axis=0)
    q = lmbda * log_jacobian - mu
    identity = np.eye(3).reshape((3, 3) + (1,) * (stretches.ndim - 1))
    return (
        mu * identity
        + lmbda / (stretches[:, None] * stretches[None, :])
        - identity * (q / stretches**2)[:, None]
    )


def check_reference_derivatives(mu, lmbda, tolerance, seed=0, samples=200, step=1.0e-5):
    """Central-difference the hand-written dW/dlambda and d2W/dlambda^2.

    Needs numpy only. The reference solution is only as trustworthy as these
    derivatives, so they are checked before they are integrated.
    """
    rng = np.random.default_rng(seed)
    stretches = rng.uniform(0.6, 1.6, size=(3, samples))
    analytic_first = principal_first_piola(stretches, mu, lmbda)
    analytic_second = principal_stretch_tangent(stretches, mu, lmbda)

    numeric_first = np.zeros_like(analytic_first)
    numeric_second = np.zeros_like(analytic_second)
    for i in range(3):
        shift = np.zeros_like(stretches)
        shift[i] = step
        numeric_first[i] = (
            strain_energy_density(stretches + shift, mu, lmbda)
            - strain_energy_density(stretches - shift, mu, lmbda)
        ) / (2.0 * step)
        numeric_second[:, i] = (
            principal_first_piola(stretches + shift, mu, lmbda)
            - principal_first_piola(stretches - shift, mu, lmbda)
        ) / (2.0 * step)

    scale_first = np.max(np.abs(analytic_first))
    scale_second = np.max(np.abs(analytic_second))
    first_error = float(np.max(np.abs(numeric_first - analytic_first)) / scale_first)
    second_error = float(np.max(np.abs(numeric_second - analytic_second)) / scale_second)
    symmetry = float(
        np.max(np.abs(analytic_second - np.swapaxes(analytic_second, 0, 1))) / scale_second
    )
    return {
        "first_derivative_relative_error": first_error,
        "second_derivative_relative_error": second_error,
        "second_derivative_asymmetry": symmetry,
        "tolerance": tolerance,
        "pass": bool(
            first_error <= tolerance and second_error <= tolerance and symmetry <= 1.0e-14
        ),
    }


def lame_thick_sphere_displacement(radius, inner, outer, pressure_pa, young_pa, poisson):
    """Small-strain radial displacement of a thick sphere under internal pressure.

    u(r) = p a^3 / (E (b^3 - a^3)) [ (1 - 2 nu) r + (1 + nu) b^3 / (2 r^2) ]

    Used as the initial guess for the nonlinear reference and as the limit the
    reference must reproduce as the pressure goes to zero. It is not the
    benchmark: the benchmark energy is the finite-deformation one.
    """
    radius = np.asarray(radius, dtype=float)
    factor = pressure_pa * inner**3 / (young_pa * (outer**3 - inner**3))
    return factor * (
        (1.0 - 2.0 * poisson) * radius + (1.0 + poisson) * outer**3 / (2.0 * radius**2)
    )


def spherical_shell_reference(
    inner,
    outer,
    pressure_pa,
    mu,
    lmbda,
    young_pa,
    poisson,
    *,
    nodes=201,
    collocation_tolerance=1.0e-10,
):
    """Reduced-order reference: the radial equilibrium ODE as a two-point BVP.

    Unknowns are y = [r(R), dr/dR]. The equilibrium
    dP_R/dR + (2/R)(P_R - P_theta) = 0 is differentiated through the material
    tangent, which is why A_11 and A_12 appear:

        dP_R/dR = A_11 (dr/dR)' + 2 A_12 (r/R)'
        (r/R)'  = (dr/dR - r/R) / R

    Boundary conditions are nominal (dead) tractions, P_R(Ri) = -p and
    P_R(Ro) = 0, matching how the project applies its reference-surface loads.
    """
    solve_bvp = import_solve_bvp()
    radii = np.linspace(inner, outer, nodes)
    guess_radius = radii + lame_thick_sphere_displacement(
        radii, inner, outer, pressure_pa, young_pa, poisson
    )
    guess = np.vstack([guess_radius, np.gradient(guess_radius, radii)])

    def stretches(radius, state):
        hoop = state[0] / radius
        return np.stack([state[1], hoop, hoop])

    def equilibrium(radius, state):
        principal = stretches(radius, state)
        piola = principal_first_piola(principal, mu, lmbda)
        tangent = principal_stretch_tangent(principal, mu, lmbda)
        hoop_rate = (principal[0] - principal[1]) / radius
        radial_rate = (
            -(2.0 / radius) * (piola[0] - piola[1]) - 2.0 * tangent[0, 1] * hoop_rate
        ) / tangent[0, 0]
        return np.vstack([state[1], radial_rate])

    def boundary(inner_state, outer_state):
        inner_stretch = np.array([inner_state[1], inner_state[0] / inner, inner_state[0] / inner])
        outer_stretch = np.array([outer_state[1], outer_state[0] / outer, outer_state[0] / outer])
        return np.array(
            [
                principal_first_piola(inner_stretch, mu, lmbda)[0] + pressure_pa,
                principal_first_piola(outer_stretch, mu, lmbda)[0],
            ]
        )

    solution = solve_bvp(
        equilibrium, boundary, radii, guess, tol=collocation_tolerance, max_nodes=200000
    )
    if solution.status != 0:
        raise SystemExit(f"reference BVP did not converge: {solution.message}")

    def radial_map(radius):
        radius = np.asarray(radius, dtype=float)
        return solution.sol(np.ravel(radius))[0].reshape(radius.shape)

    return {
        "radial_map": radial_map,
        "inner_displacement": float(solution.sol(inner)[0] - inner),
        "outer_displacement": float(solution.sol(outer)[0] - outer),
        "nodes": int(solution.x.size),
        "maximum_rms_residual": float(np.max(solution.rms_residuals)),
    }


def icosphere(subdivisions):
    """Unit-sphere triangulation from a subdivided icosahedron."""
    golden = 0.5 * (1.0 + np.sqrt(5.0))
    vertices = []
    for first in (-1.0, 1.0):
        for second in (-1.0, 1.0):
            vertices.append((0.0, first, second * golden))
            vertices.append((first, second * golden, 0.0))
            vertices.append((second * golden, 0.0, first))
    vertices = [tuple(np.asarray(v) / np.linalg.norm(v)) for v in vertices]

    points = np.asarray(vertices, dtype=float)
    distances = np.linalg.norm(points[:, None, :] - points[None, :, :], axis=2)
    edge = np.min(distances[distances > 1.0e-9])
    adjacency = np.abs(distances - edge) < 1.0e-9
    faces = [
        (i, j, k)
        for i in range(12)
        for j in range(i + 1, 12)
        for k in range(j + 1, 12)
        if adjacency[i, j] and adjacency[j, k] and adjacency[i, k]
    ]
    if len(faces) != 20:
        raise SystemExit(f"icosahedron construction produced {len(faces)} faces, expected 20")

    for _ in range(subdivisions):
        midpoints: dict[tuple[int, int], int] = {}

        def midpoint(first, second):
            key = (min(first, second), max(first, second))
            if key not in midpoints:
                point = np.asarray(vertices[first]) + np.asarray(vertices[second])
                vertices.append(tuple(point / np.linalg.norm(point)))
                midpoints[key] = len(vertices) - 1
            return midpoints[key]

        refined = []
        for a, b, c in faces:
            ab, bc, ca = midpoint(a, b), midpoint(b, c), midpoint(c, a)
            refined += [(a, ab, ca), (b, bc, ab), (c, ca, bc), (ab, bc, ca)]
        faces = refined

    return np.asarray(vertices, dtype=float), np.asarray(faces, dtype=np.int64)


def spherical_shell_cells(inner, outer, subdivisions, layers):
    """Points and tetrahedra of a thick spherical shell, radially layered prisms.

    Every vertex lies exactly on its own sphere; the facets between them are flat,
    which is the O(h^2) geometric error the convergence table reports. Each prism
    is split into three tetrahedra after sorting its surface vertex indices, so
    neighbouring prisms pick the same diagonal on the quadrilateral they share and
    the mesh stays conforming.
    """
    directions, faces = icosphere(subdivisions)
    radii = np.linspace(inner, outer, layers + 1)
    surface_count = directions.shape[0]
    points = np.concatenate([radius * directions for radius in radii], axis=0).T

    cells = []
    for face in np.sort(faces, axis=1):
        a, b, c = (int(index) for index in face)
        for layer in range(layers):
            low, high = layer * surface_count, (layer + 1) * surface_count
            a0, b0, c0 = low + a, low + b, low + c
            a1, b1, c1 = high + a, high + b, high + c
            cells += [(a0, b0, c0, c1), (a0, b0, c1, b1), (a0, b1, c1, a1)]
    cells = np.asarray(cells, dtype=np.int64).T

    corners = points[:, cells]
    edges = corners[:, 1:, :] - corners[:, :1, :]
    volumes = np.linalg.det(np.moveaxis(edges, -1, 0)) / 6.0
    if np.min(np.abs(volumes)) <= 0.0:
        raise SystemExit("generated shell mesh contains a degenerate tetrahedron")
    flipped = volumes < 0.0
    cells[1, flipped], cells[2, flipped] = cells[2, flipped], cells[1, flipped]
    return np.ascontiguousarray(points), np.ascontiguousarray(cells)


def build_shell_mesh(skfem, inner, outer, subdivisions, layers):
    points, cells = spherical_shell_cells(inner, outer, subdivisions, layers)
    return skfem.MeshTet(points, cells)


def build_basis(skfem, mesh, element_degree):
    element = skfem.ElementTetP1() if element_degree == 1 else skfem.ElementTetP2()
    quadrature_order = 3 if element_degree == 2 else 2
    return skfem.Basis(mesh, skfem.ElementVector(element), intorder=quadrature_order)


def homogeneous_dilation(basis, stretch):
    """Interpolant of u = (g - 1) X, exact for vector P1 and P2 tetrahedra."""
    mesh = basis.mesh
    values = basis.zeros()
    values[basis.nodal_dofs] = (stretch - 1.0) * mesh.p
    if basis.edge_dofs.size > 0:
        values[basis.edge_dofs] = (stretch - 1.0) * mesh.p[:, mesh.edges].mean(axis=1)
    if basis.facet_dofs.size > 0 or basis.interior_dofs.size > 0:
        raise SystemExit("This benchmark supports vector P1 and P2 tetrahedra only.")
    return values


def assemble_inner_pressure(skfem, basis, inner, outer, pressure_pa, quadrature_order):
    """Dead nominal traction -p N on the reference inner surface."""
    from skfem.helpers import dot

    mesh = basis.mesh
    boundary = mesh.boundary_facets()
    centroids = mesh.p[:, mesh.facets[:, boundary]].mean(axis=1)
    radii = np.linalg.norm(centroids, axis=0)
    selected = boundary[radii < 0.5 * (inner + outer)]
    if selected.size == 0:
        raise SystemExit("no inner-surface facets found on the generated shell mesh")

    facet_basis = skfem.FacetBasis(
        mesh, basis.elem, facets=selected, intorder=max(2, int(quadrature_order))
    )
    nominal_stress = -pressure_pa * np.eye(3)

    @skfem.LinearForm
    def traction(virtual_displacement, w):
        return dot(virtual_displacement, np.einsum("ij,j...->i...", nominal_stress, w.n))

    load = np.asarray(skfem.asm(traction, facet_basis), dtype=float)
    return load, float(np.sum(np.asarray(facet_basis.dx, dtype=float)))


def minimum_jacobian(basis, displacement):
    from skfem.helpers import det, grad, identity

    interpolated = basis.interpolate(displacement)
    return float(np.min(np.asarray(det(grad(interpolated) + identity(interpolated)))))


def solve_pressurised_shell(
    solver, skfem, scipy, basis, material, load, load_steps, maximum_iterations=24
):
    """Newton with load stepping on the augmented rigid-constraint system.

    Mirrors the project's `augmented_direct` path: same residual and tangent
    forms, same six-mode constraint, same projected-residual convergence test at
    the same 1e-8 absolute and relative tolerances.
    """
    residual_form, tangent_form = solver._region_forms(
        material, np.eye(3), 0.0, OUTER_AXES_UM, 0.0, IDENTITY_AXIS
    )
    constraints = solver._finite_element_rigid_constraint_matrix(basis, scipy, np)
    displacement = basis.zeros()
    iterations = 0

    for step in range(1, load_steps + 1):
        factor = step / load_steps
        reference_norm = None
        for _ in range(maximum_iterations):
            iterations += 1
            interpolated = basis.interpolate(displacement)
            stiffness = skfem.asm(tangent_form, basis, disp=interpolated)
            residual = (
                np.asarray(skfem.asm(residual_form, basis, disp=interpolated), dtype=float)
                - factor * load
            )
            projected = solver._project_residual_off_rigid_modes(residual, constraints, scipy, np)
            residual_norm = float(np.linalg.norm(projected))
            if reference_norm is None:
                reference_norm = max(residual_norm, 1.0)
            if residual_norm <= 1.0e-8 + 1.0e-8 * reference_norm:
                break

            augmented = scipy.sparse.bmat(
                [[stiffness.tocsc(), constraints.T], [constraints, None]], format="csc"
            )
            right_hand_side = np.concatenate((-residual, -(constraints @ displacement)))
            increment = np.asarray(
                scipy.sparse.linalg.spsolve(augmented, right_hand_side)[: basis.N], dtype=float
            )
            if not np.all(np.isfinite(increment)):
                raise SystemExit(f"non-finite Newton increment at load factor {factor:.4f}")

            relaxation = 1.0
            while relaxation >= 1.0 / 128.0:
                candidate = displacement + relaxation * increment
                if minimum_jacobian(basis, candidate) > 0.0:
                    candidate_residual = (
                        np.asarray(
                            skfem.asm(residual_form, basis, disp=basis.interpolate(candidate)),
                            dtype=float,
                        )
                        - factor * load
                    )
                    candidate_projected = solver._project_residual_off_rigid_modes(
                        candidate_residual, constraints, scipy, np
                    )
                    if float(np.linalg.norm(candidate_projected)) < residual_norm * (
                        1.0 - 1.0e-4 * relaxation
                    ):
                        displacement = candidate
                        break
                relaxation *= 0.5
            else:
                raise SystemExit(f"line search failed at load factor {factor:.4f}")
        else:
            raise SystemExit(f"Newton did not converge at load factor {factor:.4f}")

    return displacement, iterations


def displacement_l2_error(basis, displacement, radial_map):
    """L2 norms of (u_h - u_ref) and u_ref over the volume, u_ref purely radial."""
    weights = np.asarray(basis.dx, dtype=float)
    coordinates = np.asarray(basis.global_coordinates(), dtype=float)
    interpolated = np.asarray(basis.interpolate(displacement), dtype=float)
    radius = np.sqrt(np.einsum("i...,i...->...", coordinates, coordinates))
    reference = (radial_map(radius) / radius - 1.0)[None, ...] * coordinates
    difference = interpolated - reference
    error = np.sqrt(np.einsum("i...,i...,...->", difference, difference, weights))
    norm = np.sqrt(np.einsum("i...,i...,...->", reference, reference, weights))
    return float(error), float(norm)


def surface_radial_displacement(basis, displacement, radius, tolerance_um):
    points = np.asarray(basis.mesh.p, dtype=float)
    nodal = displacement[basis.nodal_dofs]
    distance = np.linalg.norm(points, axis=0)
    selected = np.abs(distance - radius) < tolerance_um
    if not np.any(selected):
        raise SystemExit(f"no mesh vertices found at radius {radius}")
    radial = np.einsum(
        "i...,i...->...", nodal[:, selected], points[:, selected] / distance[selected]
    )
    return float(np.mean(radial)), float(np.std(radial))


def refinement_level(level, base_subdivisions):
    return base_subdivisions + level, 2 * 2**level


def run_growth_case(solver, skfem, args, tolerance):
    """Case 1: exact homogeneous growth solution, verified to machine precision."""
    mu, lmbda = solver._lame_parameters(SHELL_MATERIAL)
    core_stretch = NATURAL_VOLUME_RATIO ** (1.0 / 3.0)
    shell_stretch = NATURAL_VOLUME_RATIO ** (SHELL_COUPLING / 3.0)
    growth = np.diag(np.full(3, core_stretch))

    print("=" * 78)
    print("Case 1 -- uniform isotropic growth, exact traction-free equilibrium")
    print("=" * 78)
    print(f"  g = {core_stretch:.9f}, det Fg = {np.linalg.det(growth):.9f} "
          f"(target {NATURAL_VOLUME_RATIO})")
    print(f"  mu = {mu:.4f} Pa, lambda = {lmbda:.4f} Pa")

    deformation = np.diag(np.full(3, core_stretch))
    first_piola, _ = solver._neo_hookean_first_piola_tangent(deformation, growth, mu, lmbda)
    elastic_jacobian = float(np.linalg.det(deformation) / np.linalg.det(growth))
    pointwise = float(np.max(np.abs(first_piola)) / mu)
    print(f"  pointwise  max|PK1| / mu = {pointwise:.3e}, |Je - 1| = "
          f"{abs(elastic_jacobian - 1.0):.3e}")

    subdivisions, layers = refinement_level(0, args.base_subdivisions)
    mesh = build_shell_mesh(skfem, args.inner_radius_um, args.outer_radius_um, subdivisions, layers)
    report = {
        "stretch": core_stretch,
        "determinant_growth": float(np.linalg.det(growth)),
        "pointwise_relative_stress": pointwise,
        "elastic_jacobian_error": abs(elastic_jacobian - 1.0),
        "tetrahedra": int(mesh.t.shape[1]),
        "assembled": {},
    }
    failures = []

    for degree in (1, 2):
        basis = build_basis(skfem, mesh, degree)
        residual_form, _ = solver._region_forms(
            SHELL_MATERIAL, growth, 0.0, OUTER_AXES_UM, 0.0, IDENTITY_AXIS
        )
        exact = homogeneous_dilation(basis, core_stretch)
        residual = np.asarray(
            skfem.asm(residual_form, basis, disp=basis.interpolate(exact)), dtype=float
        )
        undeformed = np.asarray(
            skfem.asm(residual_form, basis, disp=basis.interpolate(basis.zeros())), dtype=float
        )
        ratio = float(np.linalg.norm(residual, np.inf) / np.linalg.norm(undeformed, np.inf))
        passed = ratio <= tolerance
        if not passed:
            failures.append(f"P{degree}")
        print(f"  assembled  P{degree}: |R(u = (g-1)X)|_inf / |R(0)|_inf = {ratio:.3e}  "
              f"({basis.N} dofs)  {'PASS' if passed else 'FAIL'}")
        report["assembled"][f"P{degree}"] = {
            "dofs": int(basis.N),
            "relative_residual": ratio,
            "pass": bool(passed),
        }

    prestress_pa = 1000.0 * PRETENSION_MN_PER_M / SHELL_THICKNESS_UM
    sample = np.asarray(OUTER_AXES_UM, dtype=float) / np.sqrt(3.0)
    initial_stress = solver._tangential_initial_stress_matrix(sample, OUTER_AXES_UM, prestress_pa)
    spurious, _ = solver._prestress_residual_and_tangent(
        np.diag(np.full(3, shell_stretch)), initial_stress
    )
    predicted = abs(shell_stretch - 1.0) * prestress_pa * np.sqrt(2.0)
    print()
    print("  reported, not asserted: with the pretension on, (F - I) . S0 does not")
    print("  vanish at F = g I, so the exact growth solution is no longer an")
    print(f"  equilibrium. At the shell growth g = {shell_stretch:.6f} and "
          f"S0 = {prestress_pa:.3f} Pa")
    print(f"  the spurious PK1 is {np.linalg.norm(spurious):.4f} Pa "
          f"(predicted (g-1)|S0| = {predicted:.4f} Pa),")
    print(f"  {100.0 * np.linalg.norm(spurious) / mu:.2f} % of mu.")
    report["prestress_perturbation"] = {
        "shell_stretch": shell_stretch,
        "prestress_pa": prestress_pa,
        "spurious_first_piola_norm_pa": float(np.linalg.norm(spurious)),
        "predicted_norm_pa": float(predicted),
        "fraction_of_mu": float(np.linalg.norm(spurious) / mu),
    }

    report["pass"] = not failures
    print()
    print(f"  Case 1 {'PASS' if not failures else 'FAIL ' + ', '.join(failures)} "
          f"against tolerance {tolerance:.1e}")
    return report


def run_pressure_case(solver, skfem, scipy, args, tolerance):
    """Case 2: thick-walled sphere against the reduced-order ODE reference."""
    mu, lmbda = solver._lame_parameters(SHELL_MATERIAL)
    young = SHELL_MATERIAL.young_modulus_pa
    poisson = SHELL_MATERIAL.poisson_ratio

    print("=" * 78)
    print("Case 2 -- thick-walled sphere against a reduced-order ODE reference")
    print("=" * 78)
    print(f"  Ri = {args.inner_radius_um} um, Ro = {args.outer_radius_um} um, "
          f"p = {args.pressure_pa} Pa (dead nominal traction)")
    print(f"  E = {young} Pa, nu = {poisson}, mu = {mu:.4f} Pa, lambda = {lmbda:.4f} Pa")
    print("  the reference is a 1-D BVP for the same energy, integrated numerically;")
    print("  it is not a closed-form solution.")

    derivatives = check_reference_derivatives(mu, lmbda, DEFAULT_DERIVATIVE_TOLERANCE)
    print(f"  reference derivatives: dW/dl {derivatives['first_derivative_relative_error']:.2e}, "
          f"d2W/dl2 {derivatives['second_derivative_relative_error']:.2e} "
          f"({'PASS' if derivatives['pass'] else 'FAIL'})")
    if not derivatives["pass"]:
        return {"pass": False, "derivatives": derivatives}

    reference = spherical_shell_reference(
        args.inner_radius_um,
        args.outer_radius_um,
        args.pressure_pa,
        mu,
        lmbda,
        young,
        poisson,
    )
    print(f"  reference solved on {reference['nodes']} nodes, max rms residual "
          f"{reference['maximum_rms_residual']:.2e}")
    print(f"  u_r(Ri) = {reference['inner_displacement']:.6f} um, "
          f"u_r(Ro) = {reference['outer_displacement']:.6f} um "
          f"({100.0 * reference['inner_displacement'] / args.inner_radius_um:.2f} % strain)")

    small_pressure = 0.01 * args.pressure_pa
    small = spherical_shell_reference(
        args.inner_radius_um,
        args.outer_radius_um,
        small_pressure,
        mu,
        lmbda,
        young,
        poisson,
    )
    linear_inner = float(
        lame_thick_sphere_displacement(
            args.inner_radius_um,
            args.inner_radius_um,
            args.outer_radius_um,
            small_pressure,
            young,
            poisson,
        )
    )
    linear_gap = abs(small["inner_displacement"] - linear_inner) / abs(linear_inner)
    print(f"  small-pressure limit: at {small_pressure:.2f} Pa the reference is "
          f"{linear_gap:.2e} from Lame")
    if linear_gap > 1.0e-2:
        print("  WARNING: the reference does not recover the small-strain solution")

    analytic_volume = 4.0 / 3.0 * np.pi * (args.outer_radius_um**3 - args.inner_radius_um**3)
    rows = []
    for level in range(args.levels):
        subdivisions, layers = refinement_level(level, args.base_subdivisions)
        mesh = build_shell_mesh(
            skfem, args.inner_radius_um, args.outer_radius_um, subdivisions, layers
        )
        basis = build_basis(skfem, mesh, args.element_degree)
        load, inner_area = assemble_inner_pressure(
            skfem,
            basis,
            args.inner_radius_um,
            args.outer_radius_um,
            args.pressure_pa,
            3 if args.element_degree == 2 else 2,
        )
        displacement, iterations = solve_pressurised_shell(
            solver, skfem, scipy, basis, SHELL_MATERIAL, load, args.load_steps
        )
        error, norm = displacement_l2_error(basis, displacement, reference["radial_map"])
        volume = float(np.sum(np.asarray(basis.dx, dtype=float)))
        spacing = (volume / mesh.t.shape[1]) ** (1.0 / 3.0)
        inner_mean, inner_spread = surface_radial_displacement(
            basis, displacement, args.inner_radius_um, 1.0e-6
        )
        outer_mean, outer_spread = surface_radial_displacement(
            basis, displacement, args.outer_radius_um, 1.0e-6
        )
        rows.append(
            {
                "level": level,
                "tetrahedra": int(mesh.t.shape[1]),
                "dofs": int(basis.N),
                "spacing_um": float(spacing),
                "volume_defect": float(1.0 - volume / analytic_volume),
                "inner_area_defect": float(
                    1.0 - inner_area / (4.0 * np.pi * args.inner_radius_um**2)
                ),
                "relative_l2_error": float(error / norm),
                "newton_iterations": int(iterations),
                "inner_displacement_um": inner_mean,
                "inner_spread_um": inner_spread,
                "outer_displacement_um": outer_mean,
                "outer_spread_um": outer_spread,
            }
        )

    for index in range(1, len(rows)):
        previous, current = rows[index - 1], rows[index]
        current["observed_order"] = float(
            np.log(previous["relative_l2_error"] / current["relative_l2_error"])
            / np.log(previous["spacing_um"] / current["spacing_um"])
        )

    expected_order = 2 if args.element_degree == 1 else 3
    print()
    print(f"  vector P{args.element_degree}, {args.load_steps} load steps")
    print("  level    tets     dofs    h (um)   vol defect   rel L2 err   order   Newton")
    for row in rows:
        order = row.get("observed_order")
        print(
            f"  {row['level']:>5} {row['tetrahedra']:>7} {row['dofs']:>8} "
            f"{row['spacing_um']:>9.4f} {row['volume_defect']:>12.2e} "
            f"{row['relative_l2_error']:>12.3e} "
            f"{'    -  ' if order is None else f'{order:>7.3f}'} {row['newton_iterations']:>8}"
        )
    print(f"  expected {expected_order} for the discretization error alone; the straight-sided")
    print("  tetrahedra add their own O(h^2) geometric term, sized by the volume defect")
    print("  column, which sets the accuracy floor and caps the observed P2 rate near 2")
    print("  when it dominates. --base-subdivisions 2 quarters it at 4x the elements.")

    print()
    print(f"  radial displacement at the two surfaces (um), reference "
          f"{reference['inner_displacement']:.6f} and {reference['outer_displacement']:.6f}")
    print("  level     inner FEM   rel diff   spread     outer FEM   rel diff   spread")
    for row in rows:
        inner_gap = abs(row["inner_displacement_um"] - reference["inner_displacement"]) / abs(
            reference["inner_displacement"]
        )
        outer_gap = abs(row["outer_displacement_um"] - reference["outer_displacement"]) / abs(
            reference["outer_displacement"]
        )
        row["inner_relative_difference"] = float(inner_gap)
        row["outer_relative_difference"] = float(outer_gap)
        print(
            f"  {row['level']:>5} {row['inner_displacement_um']:>13.6f} {inner_gap:>10.3e} "
            f"{row['inner_spread_um']:>8.1e} {row['outer_displacement_um']:>13.6f} "
            f"{outer_gap:>10.3e} {row['outer_spread_um']:>8.1e}"
        )
    print("  spread is the standard deviation over the surface vertices, i.e. how much")
    print("  the mesh breaks the spherical symmetry the reference assumes.")

    finest = rows[-1]["relative_l2_error"]
    passed = finest <= tolerance
    print()
    print(f"  Case 2 {'PASS' if passed else 'FAIL'}: finest relative L2 error "
          f"{finest:.3e} against tolerance {tolerance:.1e}, with a geometric defect of "
          f"{rows[-1]['volume_defect']:.2e} underneath it")
    return {
        "pass": bool(passed),
        "derivatives": derivatives,
        "reference": {
            "inner_displacement_um": reference["inner_displacement"],
            "outer_displacement_um": reference["outer_displacement"],
            "nodes": reference["nodes"],
            "maximum_rms_residual": reference["maximum_rms_residual"],
            "small_pressure_lame_relative_gap": float(linear_gap),
        },
        "expected_order": expected_order,
        "levels": rows,
        "tolerance": tolerance,
    }


def run_derivative_case(tolerance):
    mu, lmbda = lame_parameters(
        SHELL_MATERIAL.young_modulus_pa, SHELL_MATERIAL.poisson_ratio
    )
    outcome = check_reference_derivatives(mu, lmbda, tolerance)
    print("=" * 78)
    print("Reference strain-energy derivatives against a central difference of W")
    print("=" * 78)
    print(f"  mu = {mu:.4f} Pa, lambda = {lmbda:.4f} Pa, 200 random stretch triples in [0.6, 1.6]")
    print(f"  dW/dlambda_i        relative error {outcome['first_derivative_relative_error']:.3e}")
    print(f"  d2W/dlambda_i dl_j  relative error {outcome['second_derivative_relative_error']:.3e}")
    print(f"  d2W asymmetry       {outcome['second_derivative_asymmetry']:.3e}")
    print(f"  {'PASS' if outcome['pass'] else 'FAIL'} against tolerance {tolerance:.1e}")
    return outcome


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--project-root",
        type=Path,
        default=None,
        help="Working tree containing src/nuclear_envelope_fem; not needed by --case derivatives.",
    )
    parser.add_argument(
        "--case",
        choices=("1", "2", "derivatives", "all"),
        default="all",
        help="1 = exact growth solution, 2 = thick sphere, derivatives = numpy-only "
        "check of the reference strain-energy derivatives.",
    )
    parser.add_argument("--element-degree", type=int, choices=(1, 2), default=1)
    parser.add_argument("--inner-radius-um", type=float, default=5.0)
    parser.add_argument("--outer-radius-um", type=float, default=7.5)
    parser.add_argument("--pressure-pa", type=float, default=200.0)
    parser.add_argument(
        "--levels",
        type=int,
        default=None,
        help="Refinement levels for case 2; default 3 for P1 and 2 for P2.",
    )
    parser.add_argument(
        "--base-subdivisions",
        type=int,
        default=1,
        help="Icosphere subdivisions of the coarsest level; 2 quarters the "
        "geometric defect at 4x the element count.",
    )
    parser.add_argument("--load-steps", type=int, default=4)
    parser.add_argument(
        "--tolerance",
        type=float,
        default=None,
        help="PASS threshold; default 1e-12 for case 1, 5e-2 relative L2 for case 2 "
        "(the geometric defect of the straight-sided mesh sets that floor), 1e-6 for "
        "the derivative check.",
    )
    parser.add_argument("--json-out", default=None)
    args = parser.parse_args()

    if args.levels is None:
        args.levels = 3 if args.element_degree == 1 else 2

    report: dict = {"case": args.case, "element_degree": args.element_degree}
    failures = []

    if args.case in ("derivatives", "all"):
        tolerance = args.tolerance if args.tolerance is not None else DEFAULT_DERIVATIVE_TOLERANCE
        outcome = run_derivative_case(tolerance)
        report["derivatives"] = outcome
        if not outcome["pass"]:
            failures.append("derivatives")
        print()

    if args.case in ("1", "2", "all"):
        solver = import_nonlinear_solver(args.project_root)
        skfem, scipy = import_assembly_stack()

    if args.case in ("1", "all"):
        tolerance = args.tolerance if args.tolerance is not None else DEFAULT_GROWTH_TOLERANCE
        outcome = run_growth_case(solver, skfem, args, tolerance)
        report["growth"] = outcome
        if not outcome["pass"]:
            failures.append("case 1")
        print()

    if args.case in ("2", "all"):
        tolerance = args.tolerance if args.tolerance is not None else DEFAULT_PRESSURE_TOLERANCE
        outcome = run_pressure_case(solver, skfem, scipy, args, tolerance)
        report["pressure"] = outcome
        if not outcome["pass"]:
            failures.append("case 2")
        print()

    print("=" * 78)
    print("FAIL: " + ", ".join(failures) if failures else "PASS: all requested cases")
    report["failures"] = failures

    if args.json_out:
        Path(args.json_out).write_text(json.dumps(report, indent=2), encoding="utf-8")
        print(f"Wrote {args.json_out}")
    return 1 if failures else 0


if __name__ == "__main__":
    raise SystemExit(main())
