#!/usr/bin/env python3
"""Tier-2 item 13 -- re-run the coarse-screen transition band on intermediate P2.

Why this run matters
--------------------
All 141 published parameter pairs (64 Latin-hypercube + 77 boundary-stress), i.e.
282 equilibria, were solved on `coarse_h1p25` P1 -- 710 vertices. That mesh is
measurably biased. At the nominal point the paired change is -37.91 % on coarse
P1 against -44.51 % on intermediate P2, and the least-unloading boundary vertex
`corner_101011` moved from -0.002078 to -0.002486 mN/m on refinement. The bias is
systematic and conservative -- refinement makes the unloading STRONGER, so the
coarse screen understates it -- but that is asserted from four re-runs only:

    lhs_021         -152.27 % -> -164.35 %   (-12.07 pp)   intermediate P1
    lhs_020          -28.44 % ->  -29.85 %   ( -1.41 pp)   intermediate P1
    lhs_061           -6.95 % ->   -7.77 %   ( -0.82 pp)   intermediate P1
    corner_101011     -4.74 % ->   -5.69 %   ( -0.95 pp)   intermediate P1

Three of those four are the algorithmic strong/median/weak triple of
`select_confirmation_rows`, and all four are mesh refinement at P1 only -- not a
degree refinement, which is where the remaining 6.6 pp at the nominal point sits.

Of the 64 LHS pairs, 58 are beyond -10 % and 6 fall inside the +/-10 % transition
band declared in `robustness_screen.coarse_screen_indifference_band_percent`:

    lhs_042  -9.25 %    lhs_048  -9.24 %    lhs_053  -9.10 %
    lhs_035  -8.06 %    lhs_008  -7.48 %    lhs_061  -6.95 %

Those six are the points nearest a sign flip and only one of them (lhs_061) was
ever re-run. The 6.6 pp coarse-to-fine drift at the nominal point is larger than
the whole distance from -6.95 % to zero, so the band is precisely where the
coarse screen is least able to carry the direction claim -- and precisely what
the confirmation design skipped. A drift of the observed sign would push these
cases further from zero; the point of this run is that nobody has checked.

This script re-solves every band case plus a stratified sample across the rest on
`intermediate_h0p75` P2 and reports the DRIFT DISTRIBUTION, so "coarse screening
is conservative" becomes a measured property with a stated range over many
points instead of an extrapolation from four.

How it drives the existing code
-------------------------------
`confirm_reynolds_finan_uq_cases.py` does expose a case-reconstruction path, and
this script reuses it rather than rebuilding one: it re-reads the published
per-case rows and replays them through
`run_reynolds_finan_epistemic_uq._run_sample(out_dir, base, row, ..., mesh_name=)`,
which is the same primitive `run_confirmation` and the boundary-stress runner
call. Only the selection differs -- `select_confirmation_rows` returns exactly
three roles and cannot express "the whole band".

Two constraints that the call itself does not show:

  * `_run_sample` hard-codes `element_degree=1`, so P2 is obtained by patching
    `_options` in the epistemic module's namespace (not the two-group runner's --
    the name is bound at import time). Upstream `_options` also couples
    `quadrature_order` and `initial_load_steps` to the degree, so this script
    applies all three together, exactly as upstream would for a P2 run.
  * the warm start is a P1 *vertex* field: `_project_p1_nodal_warm_start` requires
    shape (mesh.p, 3) and fills P2 edge dofs by averaging, so a P1 nominal
    displacement is valid for a P2 solve on the SAME mesh but never across meshes.
    The intermediate warm starts therefore come from the intermediate P1 nominal
    run, which is the root `confirm_reynolds_finan_uq_cases` already uses.

The published package ships no `src/`, no meshes and no `.npz`; point
`--project-root` at the real working tree.

Usage:
    python3 13_rerun_transition_band.py --project-root /path/to/project \\
        --out-root /path/to/out/transition_band --dry-run
"""

from __future__ import annotations

import argparse
import csv
import json
import sys
from pathlib import Path

import numpy as np

sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "tier1_scripts"))
from _common import (  # noqa: E402
    add_common_arguments,
    import_runner,
    patch_options,
)

PARAMETERS = (
    "chromatin_modulus_scale",
    "shell_young_pa",
    "shell_coupling",
    "pretension_mn_per_m",
    "natural_volume_ratio",
    "probe_shear_stress_pa",
)
CHANGE_COLUMN = "paired_mean_resultant_change_percent"
DIFFERENCE_COLUMN = "paired_mean_resultant_difference_mn_per_m"
CONTROL_COLUMN = "control_mean_resultant_mn_per_m"
HYPEROSMOTIC_COLUMN = "hyperosmotic_mean_resultant_mn_per_m"

# Published nominal anchor: boundary-stress `nominal_parameters` plus its coarse
# P1 result, used only when the boundary results cannot be found under the tree.
NOMINAL_ANCHOR = {
    "sample": "nominal",
    "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,
    CHANGE_COLUMN: -37.905677988124374,
    DIFFERENCE_COLUMN: -0.022876902453398566,
    CONTROL_COLUMN: 0.06035217853263504,
    HYPEROSMOTIC_COLUMN: 0.03747527607923647,
}

# Already published, for context in the summary. All four are intermediate P1.
PUBLISHED_REFERENCE_DRIFTS_PP = {
    "lhs_021": -12.073864957543435,
    "lhs_020": -1.4143411780729416,
    "lhs_061": -0.815249420376519,
    "corner_101011": -0.9466693184613462,
}

UQ_RESULT_NAMES = ("epistemic_uq_results.csv", "Figure3_UQ_pair_source_data.csv")
BOUNDARY_RESULT_NAMES = (
    "boundary_stress_results.csv",
    "FigureS6_boundary_stress_results_source_data.csv",
)


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    add_common_arguments(parser)
    parser.add_argument(
        "--uq-results",
        type=Path,
        default=None,
        help=(
            "Published per-case UQ rows. Default: search the tree for "
            "epistemic_uq_results.csv, then Figure3_UQ_pair_source_data.csv. "
            "reynolds_finan_epistemic_uq.json is NOT it -- its `samples` key is "
            "the integer 64, not the records."
        ),
    )
    parser.add_argument(
        "--band-percent",
        type=float,
        default=10.0,
        help="Select every case with |paired percent change| <= this value.",
    )
    parser.add_argument(
        "--stratified",
        type=int,
        default=8,
        help="Additional cases sampled evenly across the ranked change distribution.",
    )
    parser.add_argument(
        "--skip-nominal",
        action="store_true",
        help="Do not add the nominal anchor to the re-run set.",
    )
    parser.add_argument(
        "--element-order",
        choices=("p1", "p2"),
        default="p2",
        help=(
            "p2 refines mesh AND degree. p1 reproduces the published three-case "
            "confirmation protocol, which was intermediate P1."
        ),
    )
    parser.add_argument(
        "--warm-start-root",
        type=Path,
        default=None,
        help=(
            "Directory holding control/ and hyperosmotic/ nominal runs on the "
            "TARGET mesh. Default: the intermediate P1 root that "
            "confirm_reynolds_finan_uq_cases.py uses."
        ),
    )
    parser.add_argument("--cuda-device-index", type=int, default=0)
    args = parser.parse_args()

    if args.band_percent < 0.0:
        raise SystemExit("--band-percent must be non-negative")
    if args.stratified < 0:
        raise SystemExit("--stratified must be non-negative")

    results_path = _find_file(
        args.project_root, UQ_RESULT_NAMES, explicit=args.uq_results
    )
    if results_path is None:
        raise SystemExit(
            "No published per-case UQ rows found under "
            f"{args.project_root}. Looked for {', '.join(UQ_RESULT_NAMES)}. "
            "Pass --uq-results explicitly."
        )
    rows = _load_cases(results_path)
    anchor = None
    if not args.skip_nominal:
        anchor = _load_nominal_anchor(args.project_root)
    selected = select_cases(
        rows,
        band_percent=args.band_percent,
        stratified=args.stratified,
        anchor=anchor,
    )

    degree = 2 if args.element_order == "p2" else 1
    discretization = f"{args.primary_mesh}_P{degree}"
    print("Transition-band re-run -- planned runs")
    print(f"  project root  : {args.project_root}")
    print(f"  output root   : {args.out_root}")
    print(f"  coarse source : {results_path}")
    print(f"  refined onto  : {discretization}")
    print(f"  band          : |change| <= {args.band_percent:g} %")
    print(f"  stratified    : {args.stratified} across the ranked distribution")
    print(f"  selected      : {len(selected)} of {len(rows)} published cases")
    _print_design(selected)
    print(
        "\n  Reference drifts already published (all intermediate P1): "
        + ", ".join(
            f"{name} {drift:+.2f} pp"
            for name, drift in PUBLISHED_REFERENCE_DRIFTS_PP.items()
        )
    )
    if args.dry_run:
        print("\n--dry-run: nothing solved.")
        return 0

    uq, confirm = _import_uq_modules(args.project_root)
    if tuple(uq.PARAMETERS) != PARAMETERS:
        raise SystemExit(
            "Upstream PARAMETERS order changed "
            f"({uq.PARAMETERS}); this script needs updating."
        )
    if degree == 2:
        # element_degree alone would leave P2 under-integrated: upstream _options
        # ties quadrature_order and initial_load_steps to the degree.
        patch_options(uq, element_degree=2, quadrature_order=3, initial_load_steps=16)

    warm_root = _resolve_warm_start_root(
        args.warm_start_root, args.project_root, confirm, args.primary_mesh
    )
    control_nodal, hyper_nodal, nominal_volume_ratio = _load_warm_start(warm_root)
    base = json.loads(uq.BASE_CONFIG.read_text(encoding="utf-8"))
    args.out_root.mkdir(parents=True, exist_ok=True)

    element_label = f"vector_P{degree}_tetrahedron"
    result_path = args.out_root / "transition_band_rerun_results.csv"
    failure_path = args.out_root / "transition_band_rerun_failures.csv"
    emitted: list[dict] = []
    failures: list[dict] = []
    for case in selected:
        print(f"\n=== {case['sample']} ({case['selection_role']}) ===", flush=True)
        sample = {"sample": case["sample"]}
        sample.update({name: float(case[name]) for name in PARAMETERS})
        sample["selection_role"] = case["selection_role"]
        try:
            refined = uq._run_sample(
                args.out_root,
                base,
                sample,
                resume=args.resume,
                sparse_backend=args.sparse_backend,
                cuda_device_index=args.cuda_device_index,
                rigid_constraint_linear_solver=args.rigid_constraint_linear_solver,
                nominal_control_nodal=control_nodal,
                nominal_hyperosmotic_nodal=hyper_nodal,
                nominal_volume_ratio=nominal_volume_ratio,
                mesh_name=args.primary_mesh,
            )
        except Exception as exc:  # noqa: BLE001 -- never silently drop a case
            print(f"  FAILED: {type(exc).__name__}: {exc}")
            failures.append(
                {
                    "sample": case["sample"],
                    "selection_role": case["selection_role"],
                    "coarse_change_percent": case[CHANGE_COLUMN],
                    "error_type": type(exc).__name__,
                    "error_message": str(exc),
                    "eligible_for_drift_summary": False,
                }
            )
            _write_csv(failure_path, failures)
            continue
        emitted.append(_drift_row(case, refined, discretization, element_label))
        _write_csv(result_path, emitted)
        print(
            f"  coarse {emitted[-1]['coarse_change_percent']:+.2f} % -> "
            f"refined {emitted[-1]['refined_change_percent']:+.2f} % "
            f"({emitted[-1]['drift_percentage_points']:+.2f} pp)",
            flush=True,
        )

    if not emitted:
        raise SystemExit("No case converged; nothing to summarise.")
    summary = summarize_drift(emitted, failures, selected, discretization, args)
    summary_path = args.out_root / "transition_band_rerun_summary.json"
    summary_path.write_text(
        json.dumps(summary, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
    )
    _print_summary(summary, emitted, failures)
    print(f"\nWrote {result_path}")
    if failures:
        print(f"Wrote {failure_path}")
    print(f"Wrote {summary_path}")
    return 0


def select_cases(
    rows: list[dict],
    *,
    band_percent: float,
    stratified: int,
    anchor: dict | None,
) -> list[dict]:
    """Band cases, then an even walk over the ranked rest, then the anchor.

    Deduplicated by sample, first role winning, so a band case that would also be
    drawn by the stratified walk is not solved twice.
    """
    ranked = sorted(rows, key=lambda row: float(row[CHANGE_COLUMN]))
    selected: list[dict] = []
    seen: set[str] = set()

    def take(row: dict, role: str) -> None:
        if row["sample"] in seen:
            return
        seen.add(row["sample"])
        selected.append({**row, "selection_role": role})

    for row in ranked:
        if abs(float(row[CHANGE_COLUMN])) <= band_percent:
            take(row, "transition_band")
    remaining = [row for row in ranked if row["sample"] not in seen]
    if stratified > 0 and remaining:
        count = min(stratified, len(remaining))
        indices = np.unique(
            np.rint(np.linspace(0.0, len(remaining) - 1, count)).astype(int)
        )
        for order, index in enumerate(indices, start=1):
            take(remaining[int(index)], f"stratified_{order:02d}_of_{len(indices)}")
    if anchor is not None:
        take(anchor, "nominal_anchor")
    return selected


def summarize_drift(
    rows: list[dict],
    failures: list[dict],
    selected: list[dict],
    discretization: str,
    args: argparse.Namespace,
) -> dict:
    drifts = np.asarray([row["drift_percentage_points"] for row in rows])
    toward_zero = [row["sample"] for row in rows if row["drifted_toward_zero"]]
    band_rows = [row for row in rows if row["selection_role"] == "transition_band"]
    band_drifts = np.asarray([row["drift_percentage_points"] for row in band_rows])
    flipped = [row["sample"] for row in rows if not row["direction_retained"]]
    return {
        "analysis": "transition_band_and_stratified_mesh_refinement_drift",
        "coarse_discretization": "coarse_h1p25_P1",
        "refined_discretization": discretization,
        "band_percent": args.band_percent,
        "planned_cases": len(selected),
        "converged_cases": len(rows),
        "failed_cases": len(failures),
        "failed_samples": [failure["sample"] for failure in failures],
        "transition_band_cases": len(band_rows),
        "drift_percentage_points": {
            "mean": float(np.mean(drifts)),
            "median": float(np.median(drifts)),
            "minimum": float(np.min(drifts)),
            "maximum": float(np.max(drifts)),
            "standard_deviation": float(np.std(drifts, ddof=0)),
        },
        "transition_band_drift_percentage_points": (
            {
                "mean": float(np.mean(band_drifts)),
                "minimum": float(np.min(band_drifts)),
                "maximum": float(np.max(band_drifts)),
            }
            if band_rows
            else None
        ),
        "any_case_drifted_toward_zero": bool(toward_zero),
        "cases_drifted_toward_zero": toward_zero,
        "cases_drifted_toward_zero_count": len(toward_zero),
        "all_cases_retain_unloading_direction": not flipped,
        "cases_with_direction_flip": flipped,
        "coarse_screen_conservative_across_this_sample": not toward_zero,
        "published_reference_drifts_percentage_points": PUBLISHED_REFERENCE_DRIFTS_PP,
        "published_reference_note": (
            "The four published re-runs are intermediate P1, i.e. mesh refinement "
            "without degree refinement."
        ),
        "interpretation": (
            "Measured coarse-to-refined drift over the transition band and a "
            "stratified sample of the design. It quantifies the screening-mesh "
            "bias on discrete points; it is not a continuous-space proof and does "
            "not promote local field or compression endpoints."
        ),
    }


def _drift_row(
    case: dict, refined: dict, discretization: str, element_label: str
) -> dict:
    coarse_change = float(case[CHANGE_COLUMN])
    refined_change = float(refined[CHANGE_COLUMN])
    coarse_difference = float(case.get(DIFFERENCE_COLUMN, "nan"))
    refined_difference = float(refined[DIFFERENCE_COLUMN])
    row = {
        "sample": case["sample"],
        "selection_role": case["selection_role"],
    }
    row.update({name: float(case[name]) for name in PARAMETERS})
    row.update(
        {
            "coarse_change_percent": coarse_change,
            "refined_change_percent": refined_change,
            "drift_percentage_points": refined_change - coarse_change,
            "drifted_toward_zero": abs(refined_change) < abs(coarse_change),
            "direction_retained": refined_change < 0.0,
            "coarse_difference_mn_per_m": coarse_difference,
            "refined_difference_mn_per_m": refined_difference,
            "difference_drift_mn_per_m": refined_difference - coarse_difference,
            "coarse_control_mean_mn_per_m": float(case.get(CONTROL_COLUMN, "nan")),
            "coarse_hyperosmotic_mean_mn_per_m": float(
                case.get(HYPEROSMOTIC_COLUMN, "nan")
            ),
            "refined_control_mean_mn_per_m": float(refined[CONTROL_COLUMN]),
            "refined_hyperosmotic_mean_mn_per_m": float(refined[HYPEROSMOTIC_COLUMN]),
            "refined_discretization": discretization,
            # _run_sample hard-codes the P1 element label; correct it here.
            "refined_element": element_label,
            "refined_maximum_final_projected_residual_norm": max(
                float(refined["control_final_projected_residual_norm"]),
                float(refined["hyperosmotic_final_projected_residual_norm"]),
            ),
            "refined_minimum_total_jacobian": min(
                float(refined["control_minimum_total_jacobian"]),
                float(refined["hyperosmotic_minimum_total_jacobian"]),
            ),
            "interpretation": "discrete_case_mesh_refinement_drift_not_continuous_proof",
        }
    )
    return row


def _print_design(selected: list[dict]) -> None:
    header = (
        f"\n  {'sample':<14}{'role':<22}"
        + "".join(f"{name[:11]:>13}" for name in PARAMETERS)
        + f"{'coarse %':>12}"
    )
    print(header)
    print("  " + "-" * (len(header) - 3))
    for case in selected:
        line = f"  {case['sample']:<14}{case['selection_role']:<22}"
        line += "".join(f"{float(case[name]):>13.5g}" for name in PARAMETERS)
        line += f"{float(case[CHANGE_COLUMN]):>12.3f}"
        print(line)


def _print_summary(summary: dict, rows: list[dict], failures: list[dict]) -> None:
    drift = summary["drift_percentage_points"]
    print("\n=== Drift distribution ===")
    print(f"  cases            : {summary['converged_cases']} converged")
    print(
        f"  drift (pp)       : mean {drift['mean']:+.3f}, "
        f"min {drift['minimum']:+.3f}, max {drift['maximum']:+.3f}"
    )
    if summary["transition_band_drift_percentage_points"] is not None:
        band = summary["transition_band_drift_percentage_points"]
        print(
            f"  band only (pp)   : mean {band['mean']:+.3f}, "
            f"min {band['minimum']:+.3f}, max {band['maximum']:+.3f}"
        )
    if summary["any_case_drifted_toward_zero"]:
        print(
            "  toward zero      : "
            + ", ".join(summary["cases_drifted_toward_zero"])
            + "\n  The screen is NOT uniformly conservative. Report the range, "
            "not the claim."
        )
    else:
        print(
            "  toward zero      : none -- refinement strengthened unloading in "
            "every converged case."
        )
    if summary["cases_with_direction_flip"]:
        print("  DIRECTION FLIPS  : " + ", ".join(summary["cases_with_direction_flip"]))
    if failures:
        print("\n=== Cases that did NOT converge (excluded from the drift stats) ===")
        for failure in failures:
            print(
                f"  {failure['sample']:<14}{failure['error_type']}: "
                f"{failure['error_message']}"
            )


def _import_uq_modules(project_root: Path):
    """Import the epistemic runner, and the confirmation script when present."""
    try:
        import_runner(_source_root(project_root))
    except ImportError as exc:
        raise SystemExit(
            f"The runner under {project_root} could not be imported: {exc}. "
            "The submission package is incomplete -- it ships no src/, no meshes "
            "and no chromatin .npz. Point --project-root at the real working tree."
        ) from exc
    try:
        import run_reynolds_finan_epistemic_uq as uq  # noqa: PLC0415
    except ImportError as exc:
        raise SystemExit(
            f"Could not import run_reynolds_finan_epistemic_uq: {exc}. "
            "It needs src/nuclear_envelope_fem, run_extreme_finite_deformation, "
            "run_vaziri_finan_two_group and scipy on the real project tree."
        ) from exc
    try:
        import confirm_reynolds_finan_uq_cases as confirm  # noqa: PLC0415
    except ImportError as exc:
        print(f"  note: confirm_reynolds_finan_uq_cases unavailable ({exc});")
        print("        falling back to the documented intermediate P1 warm-start root.")
        confirm = None
    return uq, confirm


def _source_root(project_root: Path) -> Path:
    """Directory holding the runners, passed to _common.import_runner.

    import_runner accepts a root whose scripts sit in 05_Source_Data/ or directly
    inside it, so handing it the resolved directory covers the real tree's
    scripts/ layout as well.
    """
    for candidate in (
        project_root / "scripts",
        project_root / "05_Source_Data",
        project_root,
    ):
        if (candidate / "run_reynolds_finan_heterogeneous_two_group.py").exists():
            return candidate
    return project_root


def _resolve_warm_start_root(
    override: Path | None, project_root: Path, confirm, mesh: str
) -> Path:
    if override is not None:
        return override
    if confirm is not None and mesh == "intermediate_h0p75":
        return Path(confirm.NOMINAL_INTERMEDIATE_ROOT)
    return (
        project_root
        / "results"
        / "reynolds_finan_probe_optimized_20260807"
        / "model_runs"
        / mesh
        / "p1"
    )


def _load_warm_start(root: Path) -> tuple[np.ndarray, np.ndarray, float]:
    control = root / "control" / "nonlinear_nodal_displacement.npy"
    hyper = root / "hyperosmotic" / "nonlinear_nodal_displacement.npy"
    config = root / "hyperosmotic" / "case_config.json"
    missing = [str(path) for path in (control, hyper, config) if not path.exists()]
    if missing:
        raise SystemExit(
            "Missing nominal warm start on the target mesh:\n  "
            + "\n  ".join(missing)
            + "\nPass --warm-start-root pointing at the nominal run for this mesh. "
            "The array must be vertex-valued on the SAME mesh as the target run."
        )
    volume_ratio = float(
        json.loads(config.read_text(encoding="utf-8"))["stimulus"]["volume_ratio"]
    )
    return np.load(control), np.load(hyper), volume_ratio


def _find_file(root: Path, names: tuple[str, ...], *, explicit: Path | None = None):
    if explicit is not None:
        if not explicit.exists():
            raise SystemExit(f"--uq-results not found: {explicit}")
        return explicit
    for name in names:
        direct = [
            root / name,
            root / "05_Source_Data" / name,
        ]
        for candidate in direct:
            if candidate.exists():
                return candidate
        matches = sorted(root.rglob(name))
        if matches:
            return matches[0]
    return None


def _load_cases(path: Path) -> list[dict]:
    with Path(path).open(newline="", encoding="utf-8") as handle:
        rows = list(csv.DictReader(handle))
    if not rows:
        raise SystemExit(f"{path} contains no rows")
    required = {"sample", CHANGE_COLUMN, *PARAMETERS}
    missing = sorted(required - set(rows[0]))
    if missing:
        raise SystemExit(
            f"{path} is missing required columns: {', '.join(missing)}. "
            "Expected the per-case UQ schema (epistemic_uq_results.csv or "
            "Figure3_UQ_pair_source_data.csv)."
        )
    return rows


def _load_nominal_anchor(project_root: Path) -> dict:
    path = _find_file(project_root, BOUNDARY_RESULT_NAMES)
    if path is None:
        return dict(NOMINAL_ANCHOR)
    with path.open(newline="", encoding="utf-8") as handle:
        for row in csv.DictReader(handle):
            if row.get("sample") == "nominal":
                return row
    return dict(NOMINAL_ANCHOR)


def _write_csv(path: Path, rows: list[dict]) -> None:
    if not rows:
        return
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)


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