#!/usr/bin/env python3
"""Consistent tangent against a central finite difference of the residual.

Why this check matters
----------------------
The package reports exactly one numerical verification of the nonlinear
machinery: projected-MINRES against augmented-direct on a coarse P2 mesh, plus a
maximum projected residual norm of 2.67e-8 against a 1e-6 gate
(numerical_acceptance.json). Both statements are about the linear solver and
about how small the residual was driven. Neither can see a wrong residual or a
wrong tangent. Newton drives to zero whatever residual it is handed, so an error
in `_region_forms` returns a perfectly converged solution of the wrong equations
with every reported diagnostic still green. The endpoints are quoted to six
figures -- 0.066064 and 0.038092 mN/m -- and that precision is only meaningful
if the residual is the derivative of the intended energy and the tangent is the
derivative of the residual.

A central difference settles the second half of that pairing exactly:

    K(u) . du  vs  (R(u + eps du) - R(u - eps du)) / (2 eps)

What it proves: the assembled tangent is the derivative of the assembled
residual, to the accuracy of the difference. Any transposed index in the four
einsums of `stress_and_tangent`, any missing q term, any sign slip in the
geometric term shows up here immediately, and an inconsistent tangent is the
usual hidden cause of the halved load steps and line-search backtracking that a
converged run never reports. What it does not prove: that the residual is the
derivative of the intended W. That is what 12_benchmark_thick_sphere.py is for.

Branches
--------
Three branches carry all the model-specific code, and this script covers each on
its own so a failure localises:

    core_no_growth        Fg = I                        det Fg = 1
    core_growth           Fg = diag(0.725517^(1/3))     det Fg = 0.725517
    shell_growth_prestress
                          Fg = diag(0.725517^(0.35/3))  det Fg = 0.893780
                          S0 = 1000 * 0.05 / 0.12 = 416.667 Pa
    shell_growth_prestress_heterogeneous
                          the same with the second-order harmonic modulation

The prestress branch is the one that needs this most. `PK1 += (F - I) . S0` with
tangent `+= I (x) S0` is the least standard term in the file, it is active only
in the shell, and 17_analytic_membrane_reference.py showed the solver and the
post-processor disagree about it: residualised in the residual, pushed forward
in full as `F S0 F^T / J` in the reported proxy. Which convention is right is a
modelling argument; whether the residual and its tangent agree with each other
is a fact, and it is checkable now.

Note on the probe load: the 375 Pa traction is a nominal stress applied to the
reference outer surface and assembled once, so it is constant in u. It cancels
identically in the central difference and contributes nothing to the tangent.
Every nonlinear term lives in `_region_forms`, which is what this script drives.
That is also why the tangent must be symmetric to machine precision here: the
hyperelastic tangent is a second derivative of W, the geometric term I (x) S0 is
the second derivative of the quadratic 0.5 (F - I) : (F - I) S0 for symmetric
S0, and a dead load adds nothing. Symmetry is asserted for the branches without
prestress and reported for the branches with it.

The mesh is a tiny self-built `MeshTet().refined(n)` block translated onto the
outer ellipsoid, so this needs none of the project's gmsh meshes or chromatin
.npz -- only `src/nuclear_envelope_fem` on the given `--project-root`.

Usage:
    python3 11_verify_tangent_fd.py --project-root /path/to/working/tree
                                    [--element-degree 2] [--refinements 2]
                                    [--tolerance 1e-6] [--json-out out.json]
"""

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)
HETEROGENEITY_AXIS = (0.0, 0.0, 1.0)
HETEROGENEITY_AMPLITUDE = 0.3
DEFAULT_EPSILONS = (1e-4, 3e-5, 1e-5, 3e-6, 1e-6, 3e-7, 1e-7, 3e-8, 1e-8)

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


def import_nonlinear_solver(project_root: Path):
    """Import nuclear_envelope_fem.nonlinear_solver from the real working tree."""
    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, tags  # noqa: E402

            return nonlinear_solver, tags
    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 --help works outside the solver env."""
    try:
        import scipy.sparse  # noqa: F401
        import skfem
    except ImportError as exc:
        raise SystemExit(
            f"{exc.name} is required: this check assembles the project's own "
            "skfem forms. 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


def build_block(skfem, refinements, centre, size_um):
    """Tetrahedral block placed at a generic point of the outer ellipsoid.

    The block must sit away from the origin: `_tangential_initial_stress_field`
    builds its normal from x_i / a_i^2 and is undefined there.
    """
    mesh = skfem.MeshTet().refined(refinements)
    points = centre[:, None] + size_um * (np.asarray(mesh.p, dtype=float) - 0.5)
    return skfem.MeshTet(np.ascontiguousarray(points), mesh.t)


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 smooth_vector_field(seed, centre, size_um, wavenumber):
    """Three-mode trigonometric field with a bounded gradient on the block."""
    rng = np.random.default_rng(seed)
    directions = rng.standard_normal((3, 3, 3))
    weights = rng.standard_normal((3, 3))
    phases = rng.uniform(0.0, 2.0 * np.pi, size=(3, 3))

    def field(points):
        local = (np.asarray(points, dtype=float) - centre[:, None]) / size_um
        value = np.zeros_like(local)
        for component in range(3):
            for mode in range(3):
                argument = wavenumber * np.einsum(
                    "d,d...->...", directions[component, mode], local
                )
                value[component] += weights[component, mode] * np.sin(
                    argument + phases[component, mode]
                )
        return value

    return field


def nodal_interpolant(basis, field):
    """Interpolate a vector field onto vector P1/P2 tetrahedral degrees of freedom."""
    mesh = basis.mesh
    values = basis.zeros()
    values[basis.nodal_dofs] = field(mesh.p)
    if basis.edge_dofs.size > 0:
        values[basis.edge_dofs] = field(mesh.p[:, mesh.edges].mean(axis=1))
    if basis.facet_dofs.size > 0 or basis.interior_dofs.size > 0:
        raise SystemExit("This check supports vector P1 and P2 tetrahedra only.")
    return values


def growth_tensor(solver, tags, volume_ratio, region):
    """Build Fg through the project's own `_target_axis_scales` and `_growth_tensor`."""
    cfg = SimpleNamespace(
        stimulus=SimpleNamespace(
            axis_scale_factors=None,
            volume_ratio=volume_ratio,
            shell_eigenstrain_fraction=SHELL_COUPLING,
        )
    )
    physical_tag = tags.CORE_VOLUME if region == "core" else tags.SHELL_VOLUME
    return solver._growth_tensor(cfg, physical_tag, 1.0, solver._target_axis_scales(cfg))


def branch_definitions(solver, tags):
    prestress_pa = 1000.0 * PRETENSION_MN_PER_M / SHELL_THICKNESS_UM
    return (
        {
            "name": "core_no_growth",
            "material": CORE_MATERIAL,
            "growth": growth_tensor(solver, tags, 1.0, "core"),
            "prestress_pa": 0.0,
            "heterogeneity_amplitude": 0.0,
            "require_symmetry": True,
        },
        {
            "name": "core_growth",
            "material": CORE_MATERIAL,
            "growth": growth_tensor(solver, tags, NATURAL_VOLUME_RATIO, "core"),
            "prestress_pa": 0.0,
            "heterogeneity_amplitude": 0.0,
            "require_symmetry": True,
        },
        {
            "name": "shell_growth_prestress",
            "material": SHELL_MATERIAL,
            "growth": growth_tensor(solver, tags, NATURAL_VOLUME_RATIO, "shell"),
            "prestress_pa": prestress_pa,
            "heterogeneity_amplitude": 0.0,
            "require_symmetry": False,
        },
        {
            "name": "shell_growth_prestress_heterogeneous",
            "material": SHELL_MATERIAL,
            "growth": growth_tensor(solver, tags, NATURAL_VOLUME_RATIO, "shell"),
            "prestress_pa": prestress_pa,
            "heterogeneity_amplitude": HETEROGENEITY_AMPLITUDE,
            "require_symmetry": False,
        },
    )


def truncation_slope(epsilons, errors, optimum):
    """Fitted d log(error) / d log(eps) above the optimum; 2 for a central difference."""
    if optimum < 2:
        return float("nan")
    above = slice(0, optimum)
    return float(
        np.polyfit(np.log(epsilons[above]), np.log(errors[above]), 1)[0]
    )


def check_branch(solver, skfem, basis, branch, displacement, direction, epsilons):
    residual_form, tangent_form = solver._region_forms(
        branch["material"],
        branch["growth"],
        branch["prestress_pa"],
        OUTER_AXES_UM,
        branch["heterogeneity_amplitude"],
        HETEROGENEITY_AXIS,
    )
    stiffness = skfem.asm(tangent_form, basis, disp=basis.interpolate(displacement))
    residual = np.asarray(
        skfem.asm(residual_form, basis, disp=basis.interpolate(displacement)),
        dtype=float,
    )
    directional = np.asarray(stiffness @ direction, dtype=float)
    scale = float(np.linalg.norm(directional))
    if scale == 0.0:
        raise SystemExit(f"branch {branch['name']}: K du vanishes, nothing to compare")

    errors = []
    for epsilon in epsilons:
        plus = np.asarray(
            skfem.asm(
                residual_form, basis, disp=basis.interpolate(displacement + epsilon * direction)
            ),
            dtype=float,
        )
        minus = np.asarray(
            skfem.asm(
                residual_form, basis, disp=basis.interpolate(displacement - epsilon * direction)
            ),
            dtype=float,
        )
        difference = (plus - minus) / (2.0 * epsilon)
        errors.append(float(np.linalg.norm(difference - directional)) / scale)
    errors = np.asarray(errors, dtype=float)
    epsilons = np.asarray(epsilons, dtype=float)
    optimum = int(np.argmin(errors))

    largest = abs(stiffness).max()
    asymmetry = float(abs(stiffness - stiffness.T).max() / max(float(largest), np.finfo(float).eps))

    return {
        "name": branch["name"],
        "determinant_growth": float(np.linalg.det(branch["growth"])),
        "prestress_pa": branch["prestress_pa"],
        "heterogeneity_amplitude": branch["heterogeneity_amplitude"],
        "residual_norm": float(np.linalg.norm(residual)),
        "directional_norm": scale,
        "epsilons": epsilons.tolist(),
        "relative_errors": errors.tolist(),
        "minimum_relative_error": float(errors[optimum]),
        "optimal_epsilon": float(epsilons[optimum]),
        "v_shape": bool(0 < optimum < len(errors) - 1),
        "truncation_slope": truncation_slope(epsilons, errors, optimum),
        "relative_asymmetry": asymmetry,
        "require_symmetry": branch["require_symmetry"],
    }


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--project-root",
        type=Path,
        required=True,
        help="Working tree containing src/nuclear_envelope_fem.",
    )
    parser.add_argument("--element-degree", type=int, choices=(1, 2), default=2)
    parser.add_argument(
        "--refinements",
        type=int,
        default=2,
        help="MeshTet().refined(n); 2 gives 384 tetrahedra, which is enough.",
    )
    parser.add_argument(
        "--block-size-um",
        type=float,
        default=0.3,
        help="Edge length of the sampling block, 2.5 shell thicknesses by default.",
    )
    parser.add_argument(
        "--displacement-amplitude-um",
        type=float,
        default=0.01,
        help="Peak nodal displacement of the state u the tangent is taken at.",
    )
    parser.add_argument(
        "--wavenumber",
        type=float,
        default=2.0,
        help="Spatial frequency of the trial fields in units of the block size.",
    )
    parser.add_argument(
        "--tolerance",
        type=float,
        default=1.0e-6,
        help="Required relative error at the optimal eps.",
    )
    parser.add_argument("--seed", type=int, default=0)
    parser.add_argument("--json-out", default=None)
    args = parser.parse_args()

    solver, tags = import_nonlinear_solver(args.project_root)
    skfem = import_assembly_stack()

    centre = np.asarray(OUTER_AXES_UM, dtype=float) / np.sqrt(3.0)
    mesh = build_block(skfem, args.refinements, centre, args.block_size_um)
    basis = build_basis(skfem, mesh, args.element_degree)

    displacement = nodal_interpolant(
        basis, smooth_vector_field(args.seed, centre, args.block_size_um, args.wavenumber)
    )
    displacement *= args.displacement_amplitude_um / np.max(np.abs(displacement))
    direction = nodal_interpolant(
        basis,
        smooth_vector_field(args.seed + 1, centre, args.block_size_um, args.wavenumber),
    )
    direction /= np.max(np.abs(direction))

    print("=" * 78)
    print("Consistent tangent vs central finite difference of the residual")
    print("=" * 78)
    print(f"  project root      {args.project_root}")
    print(f"  block centre      ({centre[0]:.4f}, {centre[1]:.4f}, {centre[2]:.4f}) um "
          f"on the {OUTER_AXES_UM} ellipsoid")
    print(f"  block size        {args.block_size_um} um, {mesh.t.shape[1]} tetrahedra, "
          f"vector P{args.element_degree}, {basis.N} dofs")
    print(f"  state |u|_max     {np.max(np.abs(displacement)):.6f} um")
    print("  direction         unit max-norm, so eps is a nodal displacement in um")
    print()

    results = []
    failures = []
    for branch in branch_definitions(solver, tags):
        outcome = check_branch(
            solver, skfem, basis, branch, displacement, direction, DEFAULT_EPSILONS
        )
        results.append(outcome)

        print("-" * 78)
        print(f"branch {outcome['name']}")
        print(f"  det Fg {outcome['determinant_growth']:.6f}   "
              f"S0 {outcome['prestress_pa']:.3f} Pa   "
              f"heterogeneity {outcome['heterogeneity_amplitude']}")
        print(f"  |R(u)| {outcome['residual_norm']:.6e}   |K du| {outcome['directional_norm']:.6e}")
        print("      eps        relative error")
        for epsilon, error in zip(outcome["epsilons"], outcome["relative_errors"]):
            marker = " <- optimum" if epsilon == outcome["optimal_epsilon"] else ""
            print(f"   {epsilon:8.1e}     {error:.6e}{marker}")
        print(f"  minimum {outcome['minimum_relative_error']:.3e} at eps "
              f"{outcome['optimal_epsilon']:.1e}")
        print(f"  V-shape {outcome['v_shape']}, truncation slope "
              f"{outcome['truncation_slope']:.3f} (2 expected for a central difference)")
        print(f"  relative asymmetry of K {outcome['relative_asymmetry']:.3e}"
              + ("" if outcome["require_symmetry"] else "  (reported, not asserted)"))

        passed = outcome["minimum_relative_error"] <= args.tolerance
        if outcome["require_symmetry"] and outcome["relative_asymmetry"] > 1.0e-12:
            passed = False
            print("  FAIL asymmetric tangent without prestress")
        outcome["pass"] = bool(passed)
        if not passed:
            failures.append(outcome["name"])
        print(f"  {'PASS' if passed else 'FAIL'} against tolerance {args.tolerance:.1e}")
        if not outcome["v_shape"]:
            print("  note: no interior optimum, so one of truncation or round-off never "
                  "dominated over the sampled eps range")

    print("-" * 78)
    if failures:
        print(f"FAIL: {', '.join(failures)}")
    else:
        print("PASS: every branch reproduces its tangent by central difference")

    if args.json_out:
        report = {
            "element_degree": args.element_degree,
            "tetrahedra": int(mesh.t.shape[1]),
            "dofs": int(basis.N),
            "tolerance": args.tolerance,
            "branches": results,
            "failures": failures,
        }
        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())
