#!/usr/bin/env python3
"""Surrogate, Sobol indices and a continuous worst case from the 141 solved points.

This script runs no FEM solve. It reuses design points that were already solved
and paid for, and it turns a stated limitation of the manuscript into a result.

Why this exists
---------------
The manuscript concedes, in its own claim-boundary text, that complete coverage
of the parameter-box VERTICES "does not prove the continuous interior of the
parameter domain". Alongside that it reports parameter influence as Spearman
rank correlation over the 64-point Latin hypercube:

    shell_young_pa          -0.64        pretension_mn_per_m     +0.29
    natural_volume_ratio    +0.56        shell_coupling          +0.22
    chromatin_modulus_scale -0.08        probe_shear_stress_pa   +0.06

Sixty-four points is underpowered for six inputs, and rank correlation is blind
to exactly the two things that decide whether the vertex argument holds: it
cannot see interactions, and it cannot see non-monotonicity. Yet 141 converged,
fully described design points already exist (64 LHS + 77 boundary: 64 vertices,
12 one-at-a-time boundary points, 1 nominal anchor) and are used only to be
counted. Fitted as training data they support a different sentence. Instead of
"we sampled 141 points", the claim becomes "we searched the continuous parameter
box for a sign flip, found the global worst case, and verified it with a real
solve".

Two consequences follow directly and are reported below:

  * If the endpoint is monotone in every input, then the extremes over the box
    are attained at vertices, complete vertex coverage DOES bound the continuous
    interior, and the "does not prove the interior" caveat can be upgraded to
    "supported, under a surrogate monotonicity check". If any input is
    non-monotone, the caveat is not merely cautious, it is necessary, and the
    interior worst case reported here is the case that has to be solved.
  * Total-minus-first-order Sobol indices quantify the interaction structure
    that no rank correlation can show, and the sum of first-order indices says
    directly how much of the variance the additive picture misses.

Choice of target
----------------
Default target is the ABSOLUTE paired difference, for two independent reasons.
The manuscript itself adopted "maximum_hyperosmotic_minus_control_global_mean_
resultant" as its adversarial objective because percent change loses meaning
when a group mean crosses zero (one confirmed case reports -164.35 %). And an
analytic pass over the post-processor shows the reported endpoint carries an
additive, deformation-invariant offset exactly equal to the assumed pretension
(0.05 mN/m, about 76 % of the Control value of 0.066064 mN/m). That offset
cancels in the absolute difference but sits in the denominator of the percent
change, diluting it and injecting the pretension assumption into a number that
is meant to describe mechanics. Absolute difference is the sound target;
--target percent_change is provided for comparison only.

What this is NOT
----------------
A surrogate prediction is NOT evidence. Every number below is a statement about
a model fitted to 141 points, not about the nuclear mechanics. The worst case
reported in step 6 is a HYPOTHESIS about where the parameter box is closest to a
sign flip, and it must be confirmed by a real paired solve at that point before
it enters any claim, in either direction. Step 7 prints the exact command.

Dependencies: numpy only. scikit-learn (Gaussian process) and scipy (L-BFGS-B
refinement) are used if importable and are never required; the numpy-only path
is complete, is the default when they are absent, and is the path the built-in
self-test exercises.

The self-test (--self-test) builds 141 synthetic points from a known six-input
analytic function, writes them in the published column format, reads them back
through the production loader and checks the recovered Sobol indices, the
bootstrap intervals, the per-input monotonicity verdict and the worst-case
location against ground truth that is available in closed form.

Usage:
    python3 14_surrogate_sobol_worstcase.py --project-root SUBMISSION_DIR \
        [--uq-results Figure3_UQ_pair_source_data.csv] \
        [--boundary-results FigureS6_boundary_stress_results_source_data.csv] \
        [--target absolute_difference] [--min-r2 0.9] [--ucb-k 1.0] \
        [--seed 20260810] [--markdown-out report.md] [--json-out report.json]

    python3 14_surrogate_sobol_worstcase.py --self-test
"""

from __future__ import annotations

import argparse
import csv
import itertools
import json
import math
import sys
import tempfile
from pathlib import Path

import numpy as np

PARAMETERS = (
    "chromatin_modulus_scale",
    "shell_young_pa",
    "shell_coupling",
    "pretension_mn_per_m",
    "natural_volume_ratio",
    "probe_shear_stress_pa",
)

# Bounds and scales exactly as configs/reynolds_finan_epistemic_uq.json declares
# them; read from that file when it is present, these are the fallback.
PARAMETER_BOX = {
    "chromatin_modulus_scale": {"min": 0.5, "max": 2.0, "scale": "log"},
    "shell_young_pa": {"min": 1000.0, "max": 4000.0, "scale": "log"},
    "shell_coupling": {"min": 0.0, "max": 0.7, "scale": "linear"},
    "pretension_mn_per_m": {"min": 0.025, "max": 0.1, "scale": "log"},
    "natural_volume_ratio": {"min": 0.65, "max": 0.85, "scale": "linear"},
    "probe_shear_stress_pa": {"min": 250.0, "max": 500.0, "scale": "linear"},
}

NOMINAL_POINT = {
    "chromatin_modulus_scale": 1.0,
    "shell_young_pa": 2000.0,
    "shell_coupling": 0.35,
    "pretension_mn_per_m": 0.05,
    "natural_volume_ratio": 0.7255172413793103,
    "probe_shear_stress_pa": 375.0,
}

PUBLISHED_SPEARMAN = {
    "shell_young_pa": -0.6439102564102563,
    "natural_volume_ratio": 0.5589743589743589,
    "pretension_mn_per_m": 0.2932692307692307,
    "shell_coupling": 0.22477106227106222,
    "chromatin_modulus_scale": -0.07861721611721612,
    "probe_shear_stress_pa": 0.05714285714285715,
}

TARGET_COLUMN = {
    "absolute_difference": "paired_mean_resultant_difference_mn_per_m",
    "percent_change": "paired_mean_resultant_change_percent",
}
TARGET_UNIT = {"absolute_difference": "mN/m", "percent_change": "%"}

UQ_RESULT_NAMES = (
    "Figure3_UQ_pair_source_data.csv",
    "epistemic_uq_results.csv",
)
BOUNDARY_RESULT_NAMES = (
    "FigureS6_boundary_stress_results_source_data.csv",
    "TableS5_boundary_stress_design_and_results.csv",
    "boundary_stress_results.csv",
)
CONFIG_NAMES = (
    "reynolds_finan_epistemic_uq.json",
    "epistemic_uq_manifest.json",
)
FEM_PROJECT_ROOT = Path("/root/nuclear_envelope_fem")


# ----------------------------------------------------------------------------
# Step 1 -- load the published design points and results
# ----------------------------------------------------------------------------


def to_unit_cube(value: float, spec: dict) -> float:
    """Inverse of the design generator's _transform_unit_interval."""
    lower, upper = float(spec["min"]), float(spec["max"])
    if spec["scale"] == "log":
        if lower <= 0.0 or value <= 0.0:
            raise ValueError("Log-scale parameters must be strictly positive")
        return (math.log(value) - math.log(lower)) / (math.log(upper) - math.log(lower))
    return (value - lower) / (upper - lower)


def from_unit_cube(value: float, spec: dict) -> float:
    """Exactly the design generator's _transform_unit_interval, then clamped."""
    lower, upper = float(spec["min"]), float(spec["max"])
    if spec["scale"] == "log":
        result = math.exp(math.log(lower) + value * (math.log(upper) - math.log(lower)))
    else:
        result = lower + value * (upper - lower)
    # The log round trip can overshoot by an ulp; emitted points must stay inside
    # the declared box or the solver's own bounds check rejects them.
    return min(max(result, lower), upper)


def unit_to_original(point: np.ndarray, box: dict) -> dict:
    return {
        name: from_unit_cube(float(point[index]), box[name])
        for index, name in enumerate(PARAMETERS)
    }


def _looks_like_case(record: dict) -> bool:
    keys = set(record)
    return all(name in keys for name in PARAMETERS) and bool(
        keys & set(TARGET_COLUMN.values())
    )


def _records_from_json(path: Path) -> list[dict]:
    """Find the per-case records inside a result JSON, or say where they live."""
    data = json.loads(path.read_text(encoding="utf-8"))
    found: list[list[dict]] = []

    def walk(node) -> None:
        if isinstance(node, list):
            if node and isinstance(node[0], dict) and _looks_like_case(node[0]):
                found.append(node)
                return
            for item in node:
                walk(item)
        elif isinstance(node, dict):
            for value in node.values():
                walk(value)

    walk(data)
    if not found:
        raise SystemExit(
            f"{path} holds no per-case records. In this package "
            "reynolds_finan_epistemic_uq.json is the DESIGN CONFIG -- its "
            "'samples' key is the integer LHS count (64), not a list of cases. "
            "The solved records live in the CSVs: "
            "05_Source_Data/Figure3_UQ_pair_source_data.csv (64 LHS pairs) and "
            "05_Source_Data/FigureS6_boundary_stress_results_source_data.csv "
            "(77 boundary pairs)."
        )
    return max(found, key=len)


def _read_records(path: Path) -> list[dict]:
    path = Path(path)
    if path.suffix.lower() == ".json":
        return _records_from_json(path)
    with path.open(newline="", encoding="utf-8") as handle:
        return list(csv.DictReader(handle))


def _search_roots(project_root: Path | None) -> list[Path]:
    if project_root is not None and not Path(project_root).exists():
        raise SystemExit(f"--project-root {project_root} does not exist")
    roots = []
    for candidate in (project_root, Path.cwd(), Path(__file__).resolve().parent.parent):
        if candidate is None:
            continue
        candidate = Path(candidate).resolve()
        if candidate.exists() and candidate not in roots:
            roots.append(candidate)
    return roots


def _find_result_file(roots: list[Path], names: tuple[str, ...]) -> Path | None:
    for name in names:
        for root in roots:
            direct = root / name
            if direct.is_file():
                return direct
            for match in sorted(root.rglob(name)):
                if match.is_file():
                    return match
    return None


def load_parameter_box(roots: list[Path]) -> tuple[dict, str]:
    """Prefer the published config so the normalisation matches the design."""
    path = _find_result_file(roots, CONFIG_NAMES)
    if path is None:
        return dict(PARAMETER_BOX), "built-in fallback (published config not found)"
    data = json.loads(path.read_text(encoding="utf-8"))
    parameters = data.get("parameters") or data.get("design", {}).get("parameters")
    if not parameters or any(name not in parameters for name in PARAMETERS):
        return dict(PARAMETER_BOX), f"built-in fallback ({path} lacks parameter bounds)"
    box = {
        name: {
            "min": float(parameters[name]["min"]),
            "max": float(parameters[name]["max"]),
            "scale": str(parameters[name]["scale"]),
        }
        for name in PARAMETERS
    }
    return box, str(path)


def load_published_spearman(roots: list[Path]) -> tuple[dict, str]:
    path = _find_result_file(roots, ("epistemic_uq_summary.json",))
    if path is None:
        return dict(PUBLISHED_SPEARMAN), "built-in fallback (summary not found)"
    data = json.loads(path.read_text(encoding="utf-8"))
    ranking = data.get("parameter_ranking")
    if not ranking:
        return dict(PUBLISHED_SPEARMAN), "built-in fallback (no parameter_ranking)"
    values = {
        str(entry["parameter"]): float(entry["spearman_rho"]) for entry in ranking
    }
    if any(name not in values for name in PARAMETERS):
        return dict(PUBLISHED_SPEARMAN), "built-in fallback (incomplete ranking)"
    return values, str(path)


def load_design_points(
    uq_results: Path | None,
    boundary_results: Path | None,
    project_root: Path | None,
    target: str,
) -> dict:
    """Merge the LHS and boundary records into X (n x 6 unit cube) and y."""
    roots = _search_roots(project_root)
    box, box_source = load_parameter_box(roots)
    column = TARGET_COLUMN[target]

    sources: list[tuple[str, Path]] = []
    if uq_results is not None:
        sources.append(("latin_hypercube", Path(uq_results)))
    else:
        found = _find_result_file(roots, UQ_RESULT_NAMES)
        if found is not None:
            sources.append(("latin_hypercube", found))
    if boundary_results is not None:
        sources.append(("boundary_stress", Path(boundary_results)))
    else:
        found = _find_result_file(roots, BOUNDARY_RESULT_NAMES)
        if found is not None:
            sources.append(("boundary_stress", found))
    if not sources:
        raise SystemExit(
            "No result files found. Pass --uq-results / --boundary-results, or "
            "--project-root pointing at the submission package. Expected any of "
            f"{UQ_RESULT_NAMES} and {BOUNDARY_RESULT_NAMES}."
        )

    rows: list[np.ndarray] = []
    values: list[float] = []
    labels: list[str] = []
    families: list[str] = []
    classes: list[str] = []
    counts: dict[str, int] = {}
    used: dict[str, str] = {}
    outside = 0
    duplicates = 0
    seen: set[tuple] = set()
    for family, path in sources:
        records = _read_records(path)
        kept = 0
        for record in records:
            if not _looks_like_case(record):
                continue
            raw = record.get(column, "")
            if raw in ("", None):
                continue
            try:
                point = [float(record[name]) for name in PARAMETERS]
                value = float(raw)
            except (TypeError, ValueError):
                continue
            unit = np.array(
                [to_unit_cube(point[index], box[name]) for index, name in enumerate(PARAMETERS)]
            )
            # Passing the same file twice, or a table that repeats the nominal
            # anchor, would otherwise weight those points twice in the fit.
            key = tuple(np.round(unit, 9))
            if key in seen:
                duplicates += 1
                continue
            seen.add(key)
            if np.any(unit < -1e-6) or np.any(unit > 1.0 + 1e-6):
                outside += 1
            rows.append(np.clip(unit, 0.0, 1.0))
            values.append(value)
            labels.append(str(record.get("sample", f"{family}_{kept + 1:03d}")))
            families.append(family)
            classes.append(str(record.get("design_class", family)))
            kept += 1
        counts[family] = counts.get(family, 0) + kept
        used[family] = str(path)

    if len(rows) < 20:
        raise SystemExit(
            f"Only {len(rows)} usable records with column '{column}'. "
            "Check that the result files carry the paired outcome columns."
        )

    return {
        "X": np.asarray(rows, dtype=float),
        "y": np.asarray(values, dtype=float),
        "labels": labels,
        "families": families,
        "design_classes": classes,
        "counts": counts,
        "files": used,
        "box": box,
        "box_source": box_source,
        "target": target,
        "target_column": column,
        "outside_box": outside,
        "duplicates_dropped": duplicates,
    }


# ----------------------------------------------------------------------------
# Step 2 -- surrogate
# ----------------------------------------------------------------------------


def legendre_design(points: np.ndarray, degree: int = 2) -> tuple[np.ndarray, list[tuple]]:
    """Orthonormal total-degree Legendre basis on the unit cube.

    Orthonormality on the uniform measure is what makes the coefficients square
    to variance shares, so the analytic PCE Sobol cross-check below is exact.
    """
    if degree != 2:
        raise ValueError("Only total degree 2 is implemented")
    scaled = 2.0 * np.asarray(points, dtype=float) - 1.0
    count, dimension = scaled.shape
    first = math.sqrt(3.0) * scaled
    second = math.sqrt(5.0) * 0.5 * (3.0 * scaled**2 - 1.0)

    columns = [np.ones(count)]
    index: list[tuple] = [()]
    for axis in range(dimension):
        columns.append(first[:, axis])
        index.append(((axis, 1),))
    for axis in range(dimension):
        columns.append(second[:, axis])
        index.append(((axis, 2),))
    for left, right in itertools.combinations(range(dimension), 2):
        columns.append(first[:, left] * first[:, right])
        index.append(((left, 1), (right, 1)))
    return np.column_stack(columns), index


def ridge_solve(design: np.ndarray, values: np.ndarray, penalty: float):
    regulariser = penalty * np.eye(design.shape[1])
    regulariser[0, 0] = 0.0  # the mean term must never be shrunk
    gram = design.T @ design + regulariser
    coefficients = np.linalg.solve(gram, design.T @ values)
    return coefficients, gram


def ridge_loo_rmse(design: np.ndarray, values: np.ndarray, penalty: float) -> float:
    """Closed-form leave-one-out residuals via the hat-matrix diagonal."""
    coefficients, gram = ridge_solve(design, values, penalty)
    inverse = np.linalg.inv(gram)
    leverage = np.einsum("ij,jk,ik->i", design, inverse, design)
    residual = values - design @ coefficients
    denominator = np.clip(1.0 - leverage, 1e-8, None)
    return float(np.sqrt(np.mean((residual / denominator) ** 2)))


def select_ridge_penalty(design: np.ndarray, values: np.ndarray) -> float:
    grid = np.logspace(-8.0, 2.0, 41)
    scores = [ridge_loo_rmse(design, values, float(penalty)) for penalty in grid]
    return float(grid[int(np.argmin(scores))])


class PolynomialChaosSurrogate:
    """Total-degree-2 Legendre least squares with ridge regularisation."""

    name = "numpy polynomial chaos (total-degree-2 Legendre, ridge)"
    provides_posterior_std = True

    def __init__(self, points: np.ndarray, values: np.ndarray, penalty: float | None = None):
        self.mean = float(np.mean(values))
        self.scale = float(np.std(values))
        if self.scale <= 0.0:
            self.scale = 1.0
        standardised = (values - self.mean) / self.scale
        design, self.index = legendre_design(points)
        self.penalty = select_ridge_penalty(design, standardised) if penalty is None else penalty
        self.coefficients, gram = ridge_solve(design, standardised, self.penalty)
        self.inverse_gram = np.linalg.inv(gram)
        residual = standardised - design @ self.coefficients
        effective = float(np.trace(design @ self.inverse_gram @ design.T))
        dof = max(len(values) - effective, 1.0)
        self.noise_std = float(np.sqrt(np.sum(residual**2) / dof))
        self.detail = {
            "basis_terms": int(design.shape[1]),
            "ridge_penalty": self.penalty,
            "in_sample_rmse": float(np.sqrt(np.mean(residual**2)) * self.scale),
            "effective_parameters": effective,
        }

    def set_uncertainty_floor(self, floor: float) -> None:
        """Use the honest CV error as the uncertainty floor, not the fit residual."""
        self.noise_std = float(max(self.noise_std, floor / self.scale))

    def predict(self, points: np.ndarray, return_std: bool = False):
        design, _ = legendre_design(np.atleast_2d(points))
        mean = design @ self.coefficients * self.scale + self.mean
        if not return_std:
            return mean
        leverage = np.einsum("ij,jk,ik->i", design, self.inverse_gram, design)
        std = self.noise_std * np.sqrt(1.0 + np.clip(leverage, 0.0, None)) * self.scale
        return mean, std

    def analytic_sobol(self) -> dict:
        """Exact Sobol indices of the fitted expansion, from its coefficients."""
        weights = self.coefficients[1:] ** 2
        terms = self.index[1:]
        total_variance = float(np.sum(weights))
        if total_variance <= 0.0:
            return {}
        first = np.zeros(len(PARAMETERS))
        total = np.zeros(len(PARAMETERS))
        for weight, term in zip(weights, terms, strict=True):
            axes = {axis for axis, _ in term}
            if len(axes) == 1:
                first[next(iter(axes))] += weight
            for axis in axes:
                total[axis] += weight
        return {
            "first_order": (first / total_variance).tolist(),
            "total_order": (total / total_variance).tolist(),
        }


class GaussianProcessSurrogate:
    """scikit-learn Matern + white noise Gaussian process, used only if present."""

    name = "scikit-learn Gaussian process (Matern 5/2 + white noise)"
    provides_posterior_std = True

    def __init__(self, points: np.ndarray, values: np.ndarray, restarts: int, seed: int):
        from sklearn.gaussian_process import GaussianProcessRegressor
        from sklearn.gaussian_process.kernels import ConstantKernel, Matern, WhiteKernel

        spread = float(np.std(values)) or 1.0
        kernel = ConstantKernel(spread**2, (1e-6, 1e6)) * Matern(
            length_scale=np.ones(points.shape[1]),
            length_scale_bounds=(1e-2, 1e2),
            nu=2.5,
        ) + WhiteKernel(1e-4 * spread**2, (1e-12, 1e2))
        self.model = GaussianProcessRegressor(
            kernel=kernel,
            normalize_y=True,
            n_restarts_optimizer=int(restarts),
            random_state=int(seed),
        )
        self.model.fit(points, values)
        self.detail = {"kernel": str(self.model.kernel_)}

    def set_uncertainty_floor(self, floor: float) -> None:
        return None

    def predict(self, points: np.ndarray, return_std: bool = False):
        points = np.atleast_2d(points)
        if not return_std:
            return self.model.predict(points)
        outputs = [
            self.model.predict(points[start : start + 4096], return_std=True)
            for start in range(0, len(points), 4096)
        ]
        mean = np.concatenate([item[0] for item in outputs])
        std = np.concatenate([item[1] for item in outputs])
        return mean, std

    def analytic_sobol(self) -> dict:
        return {}


def fit_surrogate(points: np.ndarray, values: np.ndarray, args) -> object:
    if not args.force_polynomial:
        try:
            return GaussianProcessSurrogate(points, values, args.gp_restarts, args.seed)
        except ImportError:
            pass
        except Exception as error:  # a GP that fails to fit must not kill the run
            print(f"warning: Gaussian process fit failed ({error}); "
                  "falling back to the numpy polynomial chaos surrogate",
                  file=sys.stderr)
    return PolynomialChaosSurrogate(points, values)


# ----------------------------------------------------------------------------
# Step 3 -- honest cross-validation
# ----------------------------------------------------------------------------


def cross_validate(points: np.ndarray, values: np.ndarray, args, folds: int) -> dict:
    count = len(values)
    use_loo = folds <= 1 or folds >= count
    order = np.random.default_rng(args.seed).permutation(count)
    if use_loo:
        assignments = np.arange(count)
    else:
        assignments = np.empty(count, dtype=int)
        assignments[order] = np.arange(count) % folds

    predictions = np.empty(count)
    for fold in np.unique(assignments):
        test = assignments == fold
        train = ~test
        model = fit_surrogate(points[train], values[train], args)
        predictions[test] = np.asarray(model.predict(points[test]))

    residual = values - predictions
    total = float(np.sum((values - np.mean(values)) ** 2))
    r_squared = 1.0 - float(np.sum(residual**2)) / total if total > 0.0 else float("nan")
    rmse = float(np.sqrt(np.mean(residual**2)))
    spread = float(np.std(values))
    return {
        "scheme": "leave-one-out" if use_loo else f"{folds}-fold",
        "r2": r_squared,
        "rmse": rmse,
        "rmse_relative_to_spread": rmse / spread if spread > 0.0 else float("nan"),
        "max_absolute_error": float(np.max(np.abs(residual))),
        "y_spread": spread,
        "y_range": [float(np.min(values)), float(np.max(values))],
        "predictions": predictions,
    }


# ----------------------------------------------------------------------------
# Step 4 -- Sobol indices by Saltelli sampling, implemented here
# ----------------------------------------------------------------------------


def sobol_indices(
    surrogate, base_samples: int, seed: int, bootstrap: int, symmetric: bool = True
) -> dict:
    """Saltelli 2010 first order, Jansen 1999 total order. No SALib needed.

    Outputs are centred on the pooled sample mean before the products are
    formed, and by default the estimator is averaged over the (A, B) role swap.
    Both are pure variance reduction: on the self-test function they take the
    worst first-order error from 0.051 to 0.004 at the same base sample size.
    """
    rng = np.random.default_rng(seed + 101)
    dimension = len(PARAMETERS)
    matrix_a = rng.random((base_samples, dimension))
    matrix_b = rng.random((base_samples, dimension))
    values_a = np.asarray(surrogate.predict(matrix_a), dtype=float)
    values_b = np.asarray(surrogate.predict(matrix_b), dtype=float)
    centre = float(np.mean(np.concatenate([values_a, values_b])))
    values_a = values_a - centre
    values_b = values_b - centre

    first_terms = np.empty((base_samples, dimension))
    total_terms = np.empty((base_samples, dimension))
    for axis in range(dimension):
        mixed = matrix_a.copy()
        mixed[:, axis] = matrix_b[:, axis]
        values_mixed = np.asarray(surrogate.predict(mixed), dtype=float) - centre
        first_terms[:, axis] = values_b * (values_mixed - values_a)
        total_terms[:, axis] = 0.5 * (values_a - values_mixed) ** 2
        if symmetric:
            swapped = matrix_b.copy()
            swapped[:, axis] = matrix_a[:, axis]
            values_swapped = np.asarray(surrogate.predict(swapped), dtype=float) - centre
            first_terms[:, axis] = 0.5 * (
                first_terms[:, axis] + values_a * (values_swapped - values_b)
            )
            total_terms[:, axis] = 0.5 * (
                total_terms[:, axis] + 0.5 * (values_b - values_swapped) ** 2
            )

    variance = float(np.var(np.concatenate([values_a, values_b])))
    if variance <= 0.0:
        raise SystemExit("Surrogate is constant over the box; Sobol indices undefined")
    first = first_terms.mean(axis=0) / variance
    total = total_terms.mean(axis=0) / variance

    first_ci = np.full((dimension, 2), np.nan)
    total_ci = np.full((dimension, 2), np.nan)
    if bootstrap > 0:
        draws_first = np.empty((bootstrap, dimension))
        draws_total = np.empty((bootstrap, dimension))
        for draw in range(bootstrap):
            pick = rng.integers(0, base_samples, base_samples)
            resampled_variance = float(
                np.var(np.concatenate([values_a[pick], values_b[pick]]))
            )
            resampled_variance = resampled_variance if resampled_variance > 0.0 else variance
            draws_first[draw] = first_terms[pick].mean(axis=0) / resampled_variance
            draws_total[draw] = total_terms[pick].mean(axis=0) / resampled_variance
        first_ci = np.percentile(draws_first, [2.5, 97.5], axis=0).T
        total_ci = np.percentile(draws_total, [2.5, 97.5], axis=0).T

    per_sample = 2 * dimension + 2 if symmetric else dimension + 2
    return {
        "base_samples": int(base_samples),
        "estimator": (
            "Saltelli/Jansen, mean-centred"
            + (", symmetrised over the A/B swap" if symmetric else "")
        ),
        "model_evaluations": int(base_samples * per_sample),
        "variance": variance,
        "first_order": first.tolist(),
        "total_order": total.tolist(),
        "first_order_ci": first_ci.tolist(),
        "total_order_ci": total_ci.tolist(),
        "sum_first_order": float(np.sum(first)),
        "sum_total_order": float(np.sum(total)),
        "interaction_share": float(1.0 - np.sum(first)),
    }


def spearman(left: np.ndarray, right: np.ndarray) -> float:
    """Rank correlation with average ranks for ties, numpy only."""

    def rank(values: np.ndarray) -> np.ndarray:
        order = np.argsort(values, kind="mergesort")
        ranks = np.empty(len(values), dtype=float)
        ranks[order] = np.arange(1, len(values) + 1, dtype=float)
        sorted_values = values[order]
        start = 0
        while start < len(values):
            stop = start
            while stop + 1 < len(values) and sorted_values[stop + 1] == sorted_values[start]:
                stop += 1
            if stop > start:
                ranks[order[start : stop + 1]] = np.mean(
                    np.arange(start + 1, stop + 2, dtype=float)
                )
            start = stop + 1
        return ranks

    left_rank = rank(np.asarray(left, dtype=float))
    right_rank = rank(np.asarray(right, dtype=float))
    left_centred = left_rank - left_rank.mean()
    right_centred = right_rank - right_rank.mean()
    denominator = math.sqrt(float(np.sum(left_centred**2) * np.sum(right_centred**2)))
    return float(np.sum(left_centred * right_centred) / denominator) if denominator > 0 else 0.0


# ----------------------------------------------------------------------------
# Step 5 -- monotonicity along random one-dimensional lines
# ----------------------------------------------------------------------------


def monotonicity_scan(
    surrogate,
    lines: int,
    steps: int,
    seed: int,
    tie_tolerance: float,
    material_tolerance: float,
) -> dict:
    """Monotonicity along random axis-parallel lines, with the reversal amplitude.

    Strict monotonicity of a fitted quadratic is not the useful question: a
    curvature term far below the surrogate's own cross-validation error will
    turn a line non-monotone while meaning nothing. So each line also gets a
    reversal amplitude, min(largest fall after a peak, largest rise after a
    trough), which is exactly how far the line would have to be smoothed to
    become monotone in its better direction, and the verdict is taken at the
    resolution the surrogate actually has.
    """
    rng = np.random.default_rng(seed + 202)
    dimension = len(PARAMETERS)
    grid = np.linspace(0.0, 1.0, steps)
    material_tolerance = max(material_tolerance, tie_tolerance)
    result = {}
    for axis in range(dimension):
        base = rng.random((lines, dimension))
        block = np.repeat(base, steps, axis=0)
        block[:, axis] = np.tile(grid, lines)
        values = np.asarray(surrogate.predict(block)).reshape(lines, steps)
        differences = np.diff(values, axis=1)
        rising = np.any(differences > tie_tolerance, axis=1)
        falling = np.any(differences < -tie_tolerance, axis=1)
        strict = ~(rising & falling)

        drawdown = np.max(np.maximum.accumulate(values, axis=1) - values, axis=1)
        runup = np.max(values - np.minimum.accumulate(values, axis=1), axis=1)
        reversal = np.minimum(drawdown, runup)
        material = reversal <= material_tolerance

        net = values[:, -1] - values[:, 0]
        sign = "+" if np.mean(net) > 0 else ("-" if np.mean(net) < 0 else "0")
        result[PARAMETERS[axis]] = {
            "strict_monotone_fraction": float(np.mean(strict)),
            "material_monotone_fraction": float(np.mean(material)),
            "increasing_fraction": float(np.mean(material & (net > 0))),
            "decreasing_fraction": float(np.mean(material & (net < 0))),
            "median_reversal": float(np.median(reversal)),
            "p95_reversal": float(np.percentile(reversal, 95)),
            "max_reversal": float(np.max(reversal)),
            "max_reversal_in_cv_rmse": float(np.max(reversal) / material_tolerance)
            if material_tolerance > 0.0
            else float("inf"),
            "mean_end_to_end_change": float(np.mean(net)),
            "sign": sign,
        }
    strict_fractions = [entry["strict_monotone_fraction"] for entry in result.values()]
    material_fractions = [entry["material_monotone_fraction"] for entry in result.values()]
    result["_summary"] = {
        "lines_per_input": int(lines),
        "steps_per_line": int(steps),
        "tie_tolerance": float(tie_tolerance),
        "material_tolerance": float(material_tolerance),
        "minimum_strict_monotone_fraction": float(np.min(strict_fractions)),
        "minimum_material_monotone_fraction": float(np.min(material_fractions)),
        "all_inputs_strictly_monotone": bool(np.min(strict_fractions) >= 1.0),
        "all_inputs_materially_monotone": bool(np.min(material_fractions) >= 1.0),
        "non_monotone_inputs": [
            name
            for name, entry in result.items()
            if name != "_summary" and entry["material_monotone_fraction"] < 1.0
        ],
    }
    return result


# ----------------------------------------------------------------------------
# Step 6 -- global worst case by multi-start search
# ----------------------------------------------------------------------------


def make_objective(surrogate, ucb_k: float):
    """Score to MAXIMISE: closeness to a sign flip, conservatively inflated.

    The endpoint unloads (negative target), so the point closest to a sign flip
    is the point with the LARGEST predicted value. Adding k*sigma makes the
    search pessimistic about the claim rather than optimistic.
    """

    def objective(points: np.ndarray) -> np.ndarray:
        points = np.atleast_2d(points)
        if ucb_k == 0.0:
            return np.asarray(surrogate.predict(points), dtype=float)
        mean, std = surrogate.predict(points, return_std=True)
        return np.asarray(mean, dtype=float) + ucb_k * np.asarray(std, dtype=float)

    return objective


def finite_difference_gradient(objective, points: np.ndarray, step: float = 1e-4) -> np.ndarray:
    count, dimension = points.shape
    gradient = np.empty((count, dimension))
    for axis in range(dimension):
        forward = points.copy()
        backward = points.copy()
        forward[:, axis] = np.clip(forward[:, axis] + step, 0.0, 1.0)
        backward[:, axis] = np.clip(backward[:, axis] - step, 0.0, 1.0)
        span = forward[:, axis] - backward[:, axis]
        span = np.where(np.abs(span) < 1e-15, 1.0, span)
        gradient[:, axis] = (objective(forward) - objective(backward)) / span
    return gradient


def projected_gradient_ascent(objective, starts: np.ndarray, iterations: int = 250):
    points = np.clip(np.array(starts, dtype=float), 0.0, 1.0)
    values = objective(points)
    steps = np.full(len(points), 0.08)
    for _ in range(iterations):
        gradient = finite_difference_gradient(objective, points)
        norm = np.linalg.norm(gradient, axis=1, keepdims=True)
        direction = gradient / np.where(norm < 1e-14, 1.0, norm)
        pending = np.ones(len(points), dtype=bool)
        for _ in range(14):
            if not pending.any():
                break
            trial = np.clip(points + steps[:, None] * direction, 0.0, 1.0)
            trial_values = objective(trial)
            accepted = pending & (trial_values > values + 1e-15)
            points[accepted] = trial[accepted]
            values[accepted] = trial_values[accepted]
            pending &= ~accepted
            steps[pending] *= 0.5
        steps[~pending] *= 1.5
        if np.all(steps < 1e-10):
            break
    return points, values


def pattern_search(objective, points: np.ndarray, initial_step: float = 0.05):
    """Compass search polish; handles the box faces where the gradient is cut off."""
    points = np.clip(np.array(points, dtype=float), 0.0, 1.0)
    values = objective(points)
    steps = np.full(len(points), initial_step)
    dimension = points.shape[1]
    for _ in range(400):  # a hard cap so a slowly creeping point cannot spin forever
        if not np.any(steps > 1e-7):
            break
        improved = np.zeros(len(points), dtype=bool)
        for axis in range(dimension):
            for sign in (1.0, -1.0):
                trial = points.copy()
                trial[:, axis] = np.clip(trial[:, axis] + sign * steps, 0.0, 1.0)
                trial_values = objective(trial)
                better = trial_values > values + 1e-15
                points[better] = trial[better]
                values[better] = trial_values[better]
                improved |= better
        steps = np.where(improved, steps, steps * 0.5)
        if np.all(steps <= 1e-7):
            break
    return points, values


def worst_case_search(surrogate, args, box: dict) -> dict:
    objective = make_objective(surrogate, args.ucb_k)
    rng = np.random.default_rng(args.seed + 303)

    scan = rng.random((args.scan_points, len(PARAMETERS)))
    vertices = np.array(list(itertools.product((0.0, 1.0), repeat=len(PARAMETERS))))
    scan = np.vstack([scan, vertices])
    scan_values = objective(scan)
    ranked = np.argsort(scan_values)[::-1]

    # Starts must be spread out: the top of a random scan otherwise sits in one
    # basin and the "multi-start" search only ever finds that basin.
    chosen: list[int] = []
    for index in ranked[: max(20 * args.starts, args.starts)]:
        candidate = scan[index]
        if all(np.linalg.norm(candidate - scan[kept]) >= 0.25 for kept in chosen):
            chosen.append(int(index))
        if len(chosen) >= args.starts:
            break
    for index in ranked:
        if len(chosen) >= args.starts:
            break
        if int(index) not in chosen:
            chosen.append(int(index))
    starts = scan[np.array(chosen)]

    points, values = projected_gradient_ascent(objective, starts, args.max_iterations)
    points, values = pattern_search(objective, points)
    refined = "numpy projected gradient + compass search"
    if args.scipy_refine:
        points, values, refined = _scipy_refine(objective, points, values, refined)

    order = np.argsort(values)[::-1]
    points, values = points[order], values[order]

    candidates = []
    for point, value in zip(points, values, strict=True):
        if any(np.linalg.norm(point - np.asarray(kept["unit"])) < 0.05 for kept in candidates):
            continue
        mean, std = surrogate.predict(point[None, :], return_std=True)
        candidates.append(
            {
                "unit": point.tolist(),
                "parameters": unit_to_original(point, box),
                "objective": float(value),
                "predicted_mean": float(np.asarray(mean)[0]),
                "predicted_std": float(np.asarray(std)[0]),
                "on_boundary": bool(np.any(point < 1e-6) or np.any(point > 1.0 - 1e-6)),
                "active_bounds": int(
                    np.sum((point < 1e-6) | (point > 1.0 - 1e-6))
                ),
            }
        )
        if len(candidates) >= args.report_n:
            break

    return {
        "optimiser": refined,
        "ucb_k": float(args.ucb_k),
        "scan_points": int(len(scan)),
        "starts": int(len(starts)),
        "candidates": candidates,
        "best_scan_objective": float(np.max(scan_values)),
        "sign_flip_predicted": bool(candidates and candidates[0]["objective"] >= 0.0),
        "mean_sign_flip_predicted": bool(
            candidates and candidates[0]["predicted_mean"] >= 0.0
        ),
    }


def _scipy_refine(objective, points, values, label):
    try:
        from scipy.optimize import minimize
    except ImportError:
        return points, values, label + " (scipy unavailable, no refinement)"
    bounds = [(0.0, 1.0)] * points.shape[1]
    refined_points = points.copy()
    refined_values = values.copy()
    for index, point in enumerate(points):
        result = minimize(
            lambda vector: -float(objective(vector[None, :])[0]),
            point,
            method="L-BFGS-B",
            bounds=bounds,
        )
        if -float(result.fun) > refined_values[index]:
            refined_points[index] = np.clip(result.x, 0.0, 1.0)
            refined_values[index] = -float(result.fun)
    return refined_points, refined_values, label + " + scipy L-BFGS-B"


# ----------------------------------------------------------------------------
# Step 7 -- verification command for the true solver
# ----------------------------------------------------------------------------


def verification_command(point: dict, fem_root: Path, mesh: str, out_name: str) -> str:
    entries = ",\n    ".join(
        f'"{name}": {point[name]!r}' for name in PARAMETERS
    )
    return f"""python3 - <<'PY'
import json, sys
from pathlib import Path
ROOT = Path({str(fem_root)!r})
sys.path[:0] = [str(ROOT / "src"), str(ROOT / "scripts")]
import numpy as np
from run_extreme_finite_deformation import BASE_CONFIG
from run_reynolds_finan_epistemic_uq import (
    NOMINAL_CONTROL_NODAL_PATH, NOMINAL_HYPEROSMOTIC_NODAL_PATH,
    NOMINAL_HYPEROSMOTIC_CONFIG_PATH, _run_sample,
)

sample = {{
    "sample": {out_name!r},
    "design_class": "surrogate_predicted_interior_worst_case",
    {entries},
}}
out_dir = ROOT / "results" / "reynolds_finan_probe_optimized_20260807" / \\
    "optimization_20260810" / "epistemic_uq" / "surrogate_worst_case"
out_dir.mkdir(parents=True, exist_ok=True)
row = _run_sample(
    out_dir,
    json.loads(BASE_CONFIG.read_text(encoding="utf-8")),
    sample,
    resume=False,
    sparse_backend="cuda",
    cuda_device_index=0,
    rigid_constraint_linear_solver="projected_minres",
    nominal_control_nodal=np.load(NOMINAL_CONTROL_NODAL_PATH),
    nominal_hyperosmotic_nodal=np.load(NOMINAL_HYPEROSMOTIC_NODAL_PATH),
    nominal_volume_ratio=float(json.loads(
        NOMINAL_HYPEROSMOTIC_CONFIG_PATH.read_text(encoding="utf-8")
    )["stimulus"]["volume_ratio"]),
    mesh_name={mesh!r},
)
print(json.dumps({{key: row[key] for key in (
    "control_mean_resultant_mn_per_m", "hyperosmotic_mean_resultant_mn_per_m",
    "paired_mean_resultant_difference_mn_per_m",
    "paired_mean_resultant_change_percent", "paired_global_unloading",
)}}, indent=2))
PY"""


# ----------------------------------------------------------------------------
# Reporting
# ----------------------------------------------------------------------------


def run_pipeline(args, dataset: dict, emit) -> dict:
    points, values = dataset["X"], dataset["y"]
    box = dataset["box"]
    unit = TARGET_UNIT[dataset["target"]]
    report: dict = {
        "seed": int(args.seed),
        "target": dataset["target"],
        "target_column": dataset["target_column"],
        "target_unit": unit,
        "inputs": list(PARAMETERS),
        "parameter_box": box,
        "parameter_box_source": dataset["box_source"],
        "data": {
            "points": int(len(values)),
            "counts_by_family": dataset["counts"],
            "files": dataset["files"],
            "y_min": float(np.min(values)),
            "y_max": float(np.max(values)),
            "y_mean": float(np.mean(values)),
            "points_outside_declared_box": int(dataset["outside_box"]),
            "duplicate_points_dropped": int(dataset["duplicates_dropped"]),
        },
    }

    emit("# Surrogate, Sobol indices and continuous worst case")
    emit()
    emit(
        f"Target: `{dataset['target_column']}` ({unit}). "
        f"{len(values)} solved design points ("
        + ", ".join(f"{count} {name}" for name, count in dataset["counts"].items())
        + ")."
    )
    emit()
    emit("| source | file |")
    emit("|---|---|")
    for family, path in dataset["files"].items():
        emit(f"| {family} | `{path}` |")
    emit(f"| parameter box | `{dataset['box_source']}` |")
    emit()
    emit(
        f"Observed target range {np.min(values):+.6f} to {np.max(values):+.6f} {unit}"
        f" (mean {np.mean(values):+.6f}, sd {np.std(values):.6f})."
    )
    if dataset["outside_box"]:
        emit(
            f"WARNING: {dataset['outside_box']} point(s) lie outside the declared "
            "box and were clipped; check the bounds before trusting anything below."
        )
    if dataset["duplicates_dropped"]:
        emit(
            f"{dataset['duplicates_dropped']} duplicate design point(s) were "
            "dropped so they are not weighted twice in the fit."
        )

    emit()
    emit("## 1. Surrogate")
    emit()
    surrogate = fit_surrogate(points, values, args)
    emit(f"Fitted: **{surrogate.name}**.")
    for key, item in surrogate.detail.items():
        emit(f"- {key}: {item:.6g}" if isinstance(item, float) else f"- {key}: {item}")
    report["surrogate"] = {"name": surrogate.name, "detail": surrogate.detail}

    emit()
    emit("## 2. Cross-validation (every conclusion below depends on this)")
    emit()
    validation = cross_validate(points, values, args, args.cv_folds)
    predictions = validation.pop("predictions")
    report["cross_validation"] = validation
    emit(f"Scheme: {validation['scheme']}.")
    emit()
    emit("| metric | value |")
    emit("|---|---:|")
    emit(f"| out-of-sample R2 | {validation['r2']:.4f} |")
    emit(f"| RMSE | {validation['rmse']:.6f} {unit} |")
    emit(f"| RMSE / sd(y) | {validation['rmse_relative_to_spread']:.4f} |")
    emit(f"| max absolute CV error | {validation['max_absolute_error']:.6f} {unit} |")
    emit(f"| sd(y) | {validation['y_spread']:.6f} {unit} |")
    emit()
    worst = int(np.argmax(np.abs(values - predictions)))
    emit(
        f"Largest CV miss at `{dataset['labels'][worst]}`: observed "
        f"{values[worst]:+.6f}, predicted {predictions[worst]:+.6f} {unit}."
    )
    if hasattr(surrogate, "set_uncertainty_floor"):
        surrogate.set_uncertainty_floor(validation["rmse"])

    surrogate_is_poor = not (validation["r2"] >= args.min_r2)
    report["surrogate_accepted"] = not surrogate_is_poor
    if surrogate_is_poor:
        banner = "!" * 72
        emit()
        emit("```")
        emit(banner)
        emit(f"  SURROGATE REJECTED: out-of-sample R2 = {validation['r2']:.4f} "
             f"< --min-r2 = {args.min_r2}")
        emit("  Steps 3 to 5 below are VOID. Sobol indices, the monotonicity")
        emit("  verdict and the worst case are all statements about a model that")
        emit("  does not reproduce the 141 solved points. Do not quote any of")
        emit("  them. Add design points, or reduce the box, and refit.")
        emit(banner)
        emit("```")

    emit()
    emit("## 3. Sobol indices from the surrogate, against the published Spearman")
    emit()
    sobol = sobol_indices(
        surrogate, args.sobol_n, args.seed, args.bootstrap, not args.sobol_plain
    )
    report["sobol"] = sobol
    monotone = monotonicity_scan(
        surrogate,
        args.mono_lines,
        args.mono_steps,
        args.seed,
        args.mono_tol * float(np.std(values)),
        args.mono_material * validation["rmse"],
    )
    report["monotonicity"] = monotone

    published, published_source = load_published_spearman(
        _search_roots(args.project_root)
    )
    own_spearman = {
        name: spearman(points[:, index], values)
        for index, name in enumerate(PARAMETERS)
    }
    report["spearman"] = {
        "published_percent_change_64_lhs": published,
        "published_source": published_source,
        "recomputed_on_merged_points": own_spearman,
    }
    emit(
        f"Saltelli sampling: N = {sobol['base_samples']}, "
        f"{sobol['model_evaluations']} surrogate evaluations, estimator "
        f"{sobol['estimator']}, {args.bootstrap} bootstrap resamples for the "
        "95 % intervals."
    )
    emit()
    emit(
        "| input | S_i | 95 % CI | S_Ti | 95 % CI | S_Ti - S_i | direction | "
        "published rho | rho here |"
    )
    emit("|---|---:|---:|---:|---:|---:|:--:|---:|---:|")
    ranking = np.argsort(sobol["total_order"])[::-1]
    for index in ranking:
        name = PARAMETERS[index]
        first = sobol["first_order"][index]
        total = sobol["total_order"][index]
        first_ci = sobol["first_order_ci"][index]
        total_ci = sobol["total_order_ci"][index]
        first_text = (
            "n/a" if math.isnan(first_ci[0]) else f"[{first_ci[0]:+.3f}, {first_ci[1]:+.3f}]"
        )
        total_text = (
            "n/a" if math.isnan(total_ci[0]) else f"[{total_ci[0]:+.3f}, {total_ci[1]:+.3f}]"
        )
        emit(
            f"| {name} | {first:+.4f} | {first_text} | "
            f"{total:+.4f} | {total_text} | "
            f"{total - first:+.4f} | {monotone[name]['sign']} | "
            f"{published.get(name, float('nan')):+.3f} | {own_spearman[name]:+.3f} |"
        )
    emit()
    emit(
        f"Sum of first-order indices = {sobol['sum_first_order']:.4f}; "
        f"interaction share = {sobol['interaction_share']:+.4f} of the variance. "
        "Rank correlation cannot see that share at all."
    )
    analytic = surrogate.analytic_sobol()
    if analytic:
        gap = float(
            np.max(np.abs(np.array(analytic["first_order"]) - np.array(sobol["first_order"])))
        )
        gap_total = float(
            np.max(np.abs(np.array(analytic["total_order"]) - np.array(sobol["total_order"])))
        )
        report["sobol"]["analytic_cross_check"] = {
            "first_order": analytic["first_order"],
            "total_order": analytic["total_order"],
            "max_absolute_gap_first": gap,
            "max_absolute_gap_total": gap_total,
        }
        emit()
        emit(
            "Cross-check against the exact indices of the fitted expansion "
            f"(available in closed form for this basis): max gap {gap:.4f} "
            f"(first order), {gap_total:.4f} (total order). This validates the "
            "Saltelli estimator itself, not the surrogate."
        )
    emit()
    emit(
        "The published Spearman column is rank correlation of the PERCENT change "
        f"over the 64 LHS points ({published_source}); the last column recomputes "
        "rank correlation of the selected target over all merged points. Neither "
        "is comparable to a variance share, which is the point: the ranking can "
        "agree while the mechanism (interaction, non-monotonicity) is invisible."
    )

    published_rank = {
        name: position
        for position, name in enumerate(
            sorted(PARAMETERS, key=lambda item: -abs(published.get(item, 0.0))), start=1
        )
    }
    sobol_rank = {
        name: position
        for position, name in enumerate(
            sorted(
                PARAMETERS,
                key=lambda item: -sobol["total_order"][PARAMETERS.index(item)],
            ),
            start=1,
        )
    }
    moved = sorted(
        PARAMETERS, key=lambda name: -abs(published_rank[name] - sobol_rank[name])
    )
    report["ranking_shift"] = {
        name: {"published_rank": published_rank[name], "sobol_total_rank": sobol_rank[name]}
        for name in PARAMETERS
    }
    largest = moved[0]
    if abs(published_rank[largest] - sobol_rank[largest]) >= 1:
        emit()
        emit(
            f"Largest ranking disagreement: **{largest}** is rank "
            f"{published_rank[largest]} by published |rho| and rank "
            f"{sobol_rank[largest]} by total-order Sobol index "
            f"(S_Ti = {sobol['total_order'][PARAMETERS.index(largest)]:.4f})."
        )
        if (
            largest == "pretension_mn_per_m"
            and dataset["target"] == "absolute_difference"
            and published_rank[largest] < sobol_rank[largest]
        ):
            emit(
                "That is the expected signature of the additive pretension offset: "
                "the assumed pretension enters the reported endpoint as a "
                "deformation-invariant constant (0.05 mN/m, about 76 % of the "
                "0.066064 mN/m Control value), so it inflates the denominator of "
                "the percent change and correlates strongly with it while doing "
                "almost nothing to the absolute difference, where it cancels. The "
                "published rho of +0.29 for pretension is largely an artefact of "
                "the normalisation, not a mechanical sensitivity."
            )

    emit()
    emit("## 4. Monotonicity per input")
    emit()
    emit(
        f"{monotone['_summary']['lines_per_input']} random lines per input, "
        f"{monotone['_summary']['steps_per_line']} points per line. A line counts "
        "as monotone when its reversal amplitude (how far it would have to be "
        "smoothed to become monotone) stays under "
        f"{monotone['_summary']['material_tolerance']:.3e} {unit}, i.e. "
        f"{args.mono_material} x the cross-validation RMSE; the strict column "
        f"uses a numerical tie tolerance of "
        f"{monotone['_summary']['tie_tolerance']:.2e} {unit} instead."
    )
    emit()
    emit(
        "| input | monotone (at CV resolution) | strict | direction | "
        "max reversal | in CV RMSE | mean end-to-end change |"
    )
    emit("|---|---:|---:|:--:|---:|---:|---:|")
    for name in PARAMETERS:
        entry = monotone[name]
        emit(
            f"| {name} | {100 * entry['material_monotone_fraction']:.1f}% | "
            f"{100 * entry['strict_monotone_fraction']:.1f}% | {entry['sign']} | "
            f"{entry['max_reversal']:.6f} {unit} | "
            f"{entry['max_reversal_in_cv_rmse']:.2f} | "
            f"{entry['mean_end_to_end_change']:+.6f} {unit} |"
        )
    emit()
    if monotone["_summary"]["all_inputs_materially_monotone"]:
        emit(
            "**Every input is monotone on every sampled line, to the resolution "
            "the surrogate has.** For a function monotone in each coordinate the "
            "extremes over a box are attained at its vertices, so complete vertex "
            "coverage bounds the continuous interior and the manuscript's caveat "
            "that the design 'does not prove the continuous interior' can be "
            "upgraded to supported, stated as a surrogate monotonicity check at "
            "the CV quality above -- not as a proof. The worst case in step 5 "
            "should then land on a vertex; if it does not, this verdict is not "
            "trustworthy and the disagreement must be reported."
        )
    else:
        detail = "; ".join(
            f"{name}: {100 * (1.0 - monotone[name]['material_monotone_fraction']):.1f}% "
            f"of lines reverse, worst reversal "
            f"{monotone[name]['max_reversal_in_cv_rmse']:.2f} x CV RMSE"
            for name in monotone["_summary"]["non_monotone_inputs"]
        )
        emit(
            "**Not monotone in "
            + ", ".join(monotone["_summary"]["non_monotone_inputs"])
            + "**, by reversals larger than the surrogate's own error ("
            + detail
            + "). Vertex coverage therefore does NOT bound the interior: the "
            "extreme over the box can sit strictly inside it, and the manuscript's "
            "caveat is necessary rather than merely cautious. The interior "
            "candidate below is the case that has to be solved. Read the reversal "
            "sizes before leaning on this: a worst reversal barely above 1 x CV "
            "RMSE is a weak non-monotonicity and the solve in step 6 is what "
            "settles it."
        )
    if (
        monotone["_summary"]["all_inputs_materially_monotone"]
        and not monotone["_summary"]["all_inputs_strictly_monotone"]
    ):
        emit()
        emit(
            "Strict monotonicity fails on some lines, but only by reversals "
            "smaller than the cross-validation error, which is fitted noise "
            "rather than structure. Report the resolution alongside the verdict."
        )

    emit()
    emit("## 5. Global worst case (closest predicted approach to a sign flip)")
    emit()
    search = worst_case_search(surrogate, args, box)
    report["worst_case"] = search
    emit(
        f"Objective maximised: predicted target + {args.ucb_k} x sigma "
        f"({'GP posterior sd' if isinstance(surrogate, GaussianProcessSurrogate) else 'least-squares prediction sd with the CV error as its floor'}). "
        f"Optimiser: {search['optimiser']}, {search['starts']} starts drawn from "
        f"{search['scan_points']} scanned points."
    )
    emit()
    header = "| rank | " + " | ".join(PARAMETERS) + " | predicted | sigma | score | bounds active |"
    emit(header)
    emit("|---:|" + "---:|" * (len(PARAMETERS) + 4))
    for rank, candidate in enumerate(search["candidates"], start=1):
        cells = " | ".join(
            f"{candidate['parameters'][name]:.6g}" for name in PARAMETERS
        )
        emit(
            f"| {rank} | {cells} | {candidate['predicted_mean']:+.6f} | "
            f"{candidate['predicted_std']:.6f} | {candidate['objective']:+.6f} | "
            f"{candidate['active_bounds']}/6 |"
        )
    emit()
    best_observed = int(np.argmax(values))
    interior = [
        name
        for index, name in enumerate(PARAMETERS)
        if 1e-6 < search["candidates"][0]["unit"][index] < 1.0 - 1e-6
    ]
    report["worst_case"]["interior_coordinates"] = interior
    report["worst_case"]["consistent_with_monotonicity"] = bool(
        set(interior) <= set(monotone["_summary"]["non_monotone_inputs"])
    )
    emit(
        f"Least-unloading SOLVED point: `{dataset['labels'][best_observed]}` at "
        f"{values[best_observed]:+.6f} {unit}. Surrogate optimum: "
        f"{search['candidates'][0]['predicted_mean']:+.6f} {unit} "
        f"({search['candidates'][0]['objective']:+.6f} {unit} with the "
        f"{args.ucb_k}-sigma margin), "
        + (
            "at a box vertex."
            if not interior
            else "interior in " + ", ".join(interior) + "."
        )
    )
    emit()
    if not interior:
        emit(
            "The optimum is a vertex, so on this surrogate the solved boundary "
            "design already contains the worst case and the vertex argument "
            "holds. Confirm by comparing the prediction against the solved value "
            "at that vertex before saying so."
        )
    elif report["worst_case"]["consistent_with_monotonicity"]:
        emit(
            "The optimum is interior exactly in the coordinate(s) step 4 flagged "
            "as non-monotone, which is the internally consistent outcome: the "
            "vertex design cannot have found this point. That is the whole "
            "argument for solving it."
        )
    else:
        emit(
            "WARNING: the optimum is interior in "
            + ", ".join(interior)
            + " but step 4 reports those inputs as monotone. The two steps "
            "disagree, which usually means the optimum sits on a nearly flat "
            "ridge. Treat the location as poorly determined and report the "
            "predicted value, not the coordinates, until the solve settles it."
        )
    emit()
    if search["mean_sign_flip_predicted"]:
        emit(
            "The surrogate MEAN crosses zero inside the box: it predicts a point "
            "where the paired difference is no longer unloading. This is the "
            "single most consequential output of this script and it is not "
            "evidence -- solve it."
        )
    elif search["sign_flip_predicted"]:
        emit(
            f"The surrogate mean stays negative, but the {args.ucb_k}-sigma upper "
            "bound reaches zero. The claim of universal unloading is not safe by "
            "the margin of the surrogate's own uncertainty -- solve it."
        )
    else:
        emit(
            "No sign flip anywhere in the box, mean or upper bound. Under this "
            "surrogate the unloading direction is uniform over the continuous "
            "parameter domain, not merely over the 141 sampled points."
        )

    emit()
    emit("## 6. Required verification -- a surrogate prediction is not evidence")
    emit()
    best = search["candidates"][0]
    command = verification_command(
        best["parameters"], Path(args.fem_root), args.verify_mesh, "surrogate_worst_case"
    )
    report["verification"] = {
        "statement": (
            "The worst case above is a surrogate hypothesis. It carries no "
            "evidential weight until a real paired solve is run at exactly these "
            "parameters. Report the solved value, not the prediction."
        ),
        "sample": {"sample": "surrogate_worst_case", **best["parameters"]},
        "mesh": args.verify_mesh,
        "command": command,
    }
    emit(
        "Nothing above is a statement about nuclear mechanics. It is a statement "
        "about a model fitted to 141 points. Run the real paired solve at the "
        "winning point before the worst case enters any sentence of the "
        "manuscript, and report the solved difference rather than the predicted "
        "one. If the solved value disagrees with the prediction by more than the "
        f"CV RMSE ({validation['rmse']:.6f} {unit}), the surrogate is "
        "extrapolating and the search must be repeated with the new point added."
    )
    emit()
    emit("```bash")
    emit(command)
    emit("```")
    emit()
    emit(
        f"Confirmation mesh `{args.verify_mesh}`; rerun with "
        "`mesh_name=\"intermediate_h0p75\"` to repeat the manuscript's own "
        "coarse-to-intermediate confirmation step at this point."
    )
    return report


# ----------------------------------------------------------------------------
# Self-test: synthetic 141 points from a known analytic function
# ----------------------------------------------------------------------------

SELF_TEST = {
    "linear": np.array([0.10, -0.70, 0.25, 0.0, 0.60, 0.08]),
    "curvature": 0.55,      # coefficient of the concave term in pretension
    "centre": 0.37,         # interior maximiser in pretension, in unit coordinates
    "interaction": 0.45,    # shell_young_pa x natural_volume_ratio
    "scale": 0.03,
    "offset": -0.05,
}


def self_test_function(unit_points: np.ndarray) -> np.ndarray:
    """Known 6-input function on the unit cube, with exact ANOVA structure."""
    unit_points = np.atleast_2d(np.asarray(unit_points, dtype=float))
    centred = unit_points - 0.5
    centre = SELF_TEST["centre"]
    squared = (unit_points[:, 3] - centre) ** 2
    mean_squared = ((1.0 - centre) ** 3 + centre**3) / 3.0
    value = centred @ SELF_TEST["linear"]
    value = value - SELF_TEST["curvature"] * (squared - mean_squared)
    value = value + SELF_TEST["interaction"] * centred[:, 1] * centred[:, 4]
    return SELF_TEST["offset"] + SELF_TEST["scale"] * value


def self_test_truth() -> dict:
    centre = SELF_TEST["centre"]
    second = ((1.0 - centre) ** 3 + centre**3) / 3.0
    fourth = ((1.0 - centre) ** 5 + centre**5) / 5.0
    variance_quadratic = fourth - second**2

    parts = np.zeros(len(PARAMETERS))
    for axis in range(len(PARAMETERS)):
        parts[axis] = SELF_TEST["linear"][axis] ** 2 / 12.0
    parts[3] = SELF_TEST["curvature"] ** 2 * variance_quadratic
    interaction = SELF_TEST["interaction"] ** 2 / 144.0
    total_variance = float(np.sum(parts) + interaction)

    first = parts / total_variance
    total = parts.copy()
    total[1] += interaction
    total[4] += interaction
    total = total / total_variance

    best_point, best_value = None, -np.inf
    for bits in itertools.product((0.0, 1.0), repeat=len(PARAMETERS)):
        point = np.array(bits)
        point[3] = centre  # the only coordinate with an interior optimum
        value = float(self_test_function(point[None, :])[0])
        if value > best_value:
            best_point, best_value = point, value
    return {
        "first_order": first.tolist(),
        "total_order": total.tolist(),
        "interaction_share": float(interaction / total_variance),
        "argmax_unit": best_point.tolist(),
        "argmax_value": best_value,
        "monotone": [True, True, True, False, True, True],
        "signs": ["+", "-", "+", "0", "+", "+"],
    }


def _write_synthetic_csv(path: Path, records: list[dict], columns: list[str]) -> None:
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=columns)
        writer.writeheader()
        writer.writerows(records)


def build_synthetic_dataset(directory: Path, seed: int) -> tuple[Path, Path]:
    """64 LHS + 77 boundary points, in the published column format."""
    box = PARAMETER_BOX
    rng = np.random.default_rng(seed)
    samples = 64
    unit = np.empty((samples, len(PARAMETERS)))
    for column in range(len(PARAMETERS)):
        unit[:, column] = (rng.permutation(samples) + rng.random(samples)) / samples

    def record(name: str, unit_point: np.ndarray, design_class: str) -> dict:
        original = unit_to_original(unit_point, box)
        value = float(self_test_function(unit_point[None, :])[0])
        control = 0.066064
        return {
            "sample": name,
            "design_class": design_class,
            **{key: repr(item) for key, item in original.items()},
            "mesh": "coarse_h1p25",
            "control_mean_resultant_mn_per_m": repr(control),
            "hyperosmotic_mean_resultant_mn_per_m": repr(control + value),
            "paired_mean_resultant_difference_mn_per_m": repr(value),
            "paired_mean_resultant_change_percent": repr(100.0 * value / control),
            "paired_global_unloading": str(value < 0.0),
        }

    columns = [
        "sample",
        "design_class",
        *PARAMETERS,
        "mesh",
        "control_mean_resultant_mn_per_m",
        "hyperosmotic_mean_resultant_mn_per_m",
        "paired_mean_resultant_difference_mn_per_m",
        "paired_mean_resultant_change_percent",
        "paired_global_unloading",
    ]

    lhs_records = [
        record(f"lhs_{index + 1:03d}", unit[index], "latin_hypercube")
        for index in range(samples)
    ]

    boundary_records = []
    for bits in itertools.product((0.0, 1.0), repeat=len(PARAMETERS)):
        name = "corner_" + "".join(str(int(bit)) for bit in bits)
        boundary_records.append(record(name, np.array(bits), "hyperrectangle_vertex"))
    nominal_unit = np.array(
        [to_unit_cube(NOMINAL_POINT[name], box[name]) for name in PARAMETERS]
    )
    for axis, name in enumerate(PARAMETERS):
        for side, value in (("low", 0.0), ("high", 1.0)):
            point = nominal_unit.copy()
            point[axis] = value
            boundary_records.append(
                record(f"axis_{axis + 1:02d}_{side}", point, "one_at_a_time_boundary")
            )
    boundary_records.append(record("nominal", nominal_unit, "nominal_anchor"))

    uq_path = directory / "Figure3_UQ_pair_source_data.csv"
    boundary_path = directory / "FigureS6_boundary_stress_results_source_data.csv"
    _write_synthetic_csv(uq_path, lhs_records, columns)
    _write_synthetic_csv(boundary_path, boundary_records, columns)
    return uq_path, boundary_path


def run_self_test(args) -> int:
    truth = self_test_truth()
    lines: list[str] = []

    def emit(text: str = "") -> None:
        print(text)
        lines.append(text)

    with tempfile.TemporaryDirectory(prefix="surrogate_selftest_") as directory:
        uq_path, boundary_path = build_synthetic_dataset(Path(directory), args.seed)
        dataset = load_design_points(uq_path, boundary_path, None, "absolute_difference")
        emit("# Self-test on a synthetic 141-point dataset with known ground truth")
        emit()
        emit(
            "Six-input analytic function on the unit cube, written into the real "
            "column format and read back through the production loader, so the "
            "log/linear normalisation is exercised end to end. Ground truth: "
            "five additive terms, one concave interior term in pretension "
            f"(maximiser at u = {SELF_TEST['centre']}), and one "
            "shell_young_pa x natural_volume_ratio interaction."
        )
        emit()
        emit(f"Loaded {len(dataset['y'])} points "
             + ", ".join(f"{count} {name}" for name, count in dataset["counts"].items())
             + ". The published Spearman column below belongs to the real study "
             "and is meaningless against synthetic data; ignore it here.")

        # The known argmax is a property of the mean surface, so the confidence
        # margin must be switched off for the location check to be meaningful.
        args.ucb_k = 0.0
        args.force_polynomial = True
        report = run_pipeline(args, dataset, emit)

    checks = []
    validation = report["cross_validation"]
    checks.append(("out-of-sample R2 >= 0.999", validation["r2"] >= 0.999,
                   f"{validation['r2']:.6f}"))

    first_gap = float(
        np.max(np.abs(np.array(report["sobol"]["first_order"]) - np.array(truth["first_order"])))
    )
    total_gap = float(
        np.max(np.abs(np.array(report["sobol"]["total_order"]) - np.array(truth["total_order"])))
    )
    checks.append(("max |S_i - true| <= 0.02", first_gap <= 0.02, f"{first_gap:.5f}"))
    checks.append(("max |S_Ti - true| <= 0.02", total_gap <= 0.02, f"{total_gap:.5f}"))
    interaction_gap = abs(report["sobol"]["interaction_share"] - truth["interaction_share"])
    checks.append(
        ("interaction share matches within 0.02", interaction_gap <= 0.02,
         f"{report['sobol']['interaction_share']:.5f} vs {truth['interaction_share']:.5f}")
    )

    # The absolute tolerances above are arbitrary; this one is not. It asks
    # whether the reported bootstrap intervals actually cover the truth, so the
    # uncertainty the script prints is tested, not just its point estimates.
    outside = 0
    for axis in range(len(PARAMETERS)):
        for key, true_values in (
            ("first_order_ci", truth["first_order"]),
            ("total_order_ci", truth["total_order"]),
        ):
            low, high = report["sobol"][key][axis]
            if not (low - 0.005 <= true_values[axis] <= high + 0.005):
                outside += 1
    checks.append(
        ("truth inside the reported 95 % bootstrap intervals (<= 2 of 12 outside)",
         outside <= 2, f"{outside} of 12 outside")
    )

    monotone_ok = True
    monotone_detail = []
    for axis, name in enumerate(PARAMETERS):
        fraction = report["monotonicity"][name]["material_monotone_fraction"]
        expected = truth["monotone"][axis]
        good = fraction >= 0.99 if expected else fraction <= 0.01
        sign_good = (
            report["monotonicity"][name]["sign"] == truth["signs"][axis]
            if expected
            else True
        )
        monotone_ok = monotone_ok and good and sign_good
        monotone_detail.append(f"{name}={fraction:.2f}")
    checks.append(
        ("monotonicity per input matches", monotone_ok, ", ".join(monotone_detail))
    )

    best = np.array(report["worst_case"]["candidates"][0]["unit"])
    location_gap = float(np.max(np.abs(best - np.array(truth["argmax_unit"]))))
    checks.append(
        ("worst case within 0.03 of the true argmax", location_gap <= 0.03,
         f"max coordinate error {location_gap:.5f}")
    )
    value_gap = abs(
        report["worst_case"]["candidates"][0]["predicted_mean"] - truth["argmax_value"]
    )
    checks.append(
        ("worst-case value within 1e-4", value_gap <= 1e-4, f"{value_gap:.3e}")
    )

    emit()
    emit("## Self-test verdict")
    emit()
    emit("| check | result | detail |")
    emit("|---|:--:|---|")
    for label, passed, detail in checks:
        emit(f"| {label} | {'PASS' if passed else 'FAIL'} | {detail} |")
    emit()
    emit("| input | S_i recovered | S_i true | S_Ti recovered | S_Ti true |")
    emit("|---|---:|---:|---:|---:|")
    for axis, name in enumerate(PARAMETERS):
        emit(
            f"| {name} | {report['sobol']['first_order'][axis]:+.4f} | "
            f"{truth['first_order'][axis]:+.4f} | "
            f"{report['sobol']['total_order'][axis]:+.4f} | "
            f"{truth['total_order'][axis]:+.4f} |"
        )
    emit()
    emit("| coordinate | recovered argmax | true argmax |")
    emit("|---|---:|---:|")
    for axis, name in enumerate(PARAMETERS):
        emit(f"| {name} | {best[axis]:.4f} | {truth['argmax_unit'][axis]:.4f} |")

    passed = all(item[1] for item in checks)
    emit()
    emit(f"**Self-test {'PASSED' if passed else 'FAILED'}.**")

    if args.markdown_out:
        Path(args.markdown_out).write_text("\n".join(lines) + "\n", encoding="utf-8")
        print(f"\nWrote {args.markdown_out}", file=sys.stderr)
    if args.json_out:
        payload = {"self_test": {"checks": [
            {"check": label, "passed": bool(ok), "detail": detail}
            for label, ok, detail in checks
        ], "ground_truth": truth}, "report": _json_ready(report)}
        Path(args.json_out).write_text(json.dumps(payload, indent=2), encoding="utf-8")
        print(f"Wrote {args.json_out}", file=sys.stderr)
    return 0 if passed else 1


def _json_ready(value):
    if isinstance(value, dict):
        return {key: _json_ready(item) for key, item in value.items()}
    if isinstance(value, (list, tuple)):
        return [_json_ready(item) for item in value]
    if isinstance(value, (np.floating, np.integer)):
        return value.item()
    if isinstance(value, np.ndarray):
        return value.tolist()
    if isinstance(value, Path):
        return str(value)
    return value


def main() -> int:
    parser = argparse.ArgumentParser(
        description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter
    )
    parser.add_argument("--project-root", type=Path, default=None)
    parser.add_argument("--uq-results", type=Path, default=None)
    parser.add_argument("--boundary-results", type=Path, default=None)
    parser.add_argument(
        "--target", choices=tuple(TARGET_COLUMN), default="absolute_difference"
    )
    parser.add_argument("--min-r2", type=float, default=0.9)
    parser.add_argument("--cv-folds", type=int, default=0, help="0 or 1 means leave-one-out")
    parser.add_argument("--sobol-n", type=int, default=16384)
    parser.add_argument("--sobol-plain", action="store_true",
                        help="one-sided Saltelli estimator instead of the A/B-symmetrised one")
    parser.add_argument("--bootstrap", type=int, default=200)
    parser.add_argument("--mono-lines", type=int, default=512)
    parser.add_argument("--mono-steps", type=int, default=33)
    parser.add_argument("--mono-tol", type=float, default=1e-6,
                        help="numerical tie tolerance as a fraction of sd(y)")
    parser.add_argument("--mono-material", type=float, default=1.0,
                        help="a reversal counts only above this multiple of the CV RMSE")
    parser.add_argument("--ucb-k", type=float, default=1.0,
                        help="conservative margin: maximise mean + k*sigma")
    parser.add_argument("--scan-points", type=int, default=20000)
    parser.add_argument("--starts", type=int, default=24)
    parser.add_argument("--max-iterations", type=int, default=250)
    parser.add_argument("--report-n", type=int, default=5)
    parser.add_argument("--gp-restarts", type=int, default=8)
    parser.add_argument("--force-polynomial", action="store_true",
                        help="skip scikit-learn even if it is importable")
    parser.add_argument("--scipy-refine", action="store_true",
                        help="polish with scipy L-BFGS-B if scipy is importable")
    parser.add_argument("--fem-root", type=Path, default=FEM_PROJECT_ROOT)
    parser.add_argument("--verify-mesh", default="coarse_h1p25")
    parser.add_argument("--seed", type=int, default=20260810)
    parser.add_argument("--markdown-out", type=Path, default=None)
    parser.add_argument("--json-out", type=Path, default=None)
    parser.add_argument("--self-test", action="store_true",
                        help="run the whole pipeline on synthetic data with known truth")
    args = parser.parse_args()

    if args.self_test:
        return run_self_test(args)

    lines: list[str] = []

    def emit(text: str = "") -> None:
        print(text)
        lines.append(text)

    dataset = load_design_points(
        args.uq_results, args.boundary_results, args.project_root, args.target
    )
    report = run_pipeline(args, dataset, emit)

    if args.markdown_out:
        Path(args.markdown_out).write_text("\n".join(lines) + "\n", encoding="utf-8")
        print(f"\nWrote {args.markdown_out}", file=sys.stderr)
    if args.json_out:
        Path(args.json_out).write_text(
            json.dumps(_json_ready(report), indent=2), encoding="utf-8"
        )
        print(f"Wrote {args.json_out}", file=sys.stderr)
    return 0 if report.get("surrogate_accepted", False) else 2


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