"""Shared helpers for the Tier-1 supplementary runs.

These scripts drive the EXISTING project code without editing it. Anything the
runner does not expose on its command line (the homogeneous-core switch and the
tangent stability modes) is applied by patching `_options` at import time, so the
upstream source tree stays untouched and the runs remain diff-able.

The submission zip does not contain `src/nuclear_envelope_fem/`, the gmsh meshes
or the chromatin `.npz`; point `--project-root` at the real working tree.
"""

from __future__ import annotations

import argparse
import csv
import sys
from pathlib import Path

SUMMARY_NAME = "heterogeneous_two_group_summary.csv"
ENDPOINT_COLUMN = "mean_shell_resultant_mn_per_m"


def add_common_arguments(parser: argparse.ArgumentParser) -> None:
    parser.add_argument(
        "--project-root",
        type=Path,
        required=True,
        help="Working tree containing src/nuclear_envelope_fem, meshes and results.",
    )
    parser.add_argument(
        "--out-root",
        type=Path,
        required=True,
        help="Directory to receive the new run folders.",
    )
    parser.add_argument(
        "--primary-mesh",
        default="intermediate_h0p75",
        help="Mesh used as the converged reference for acceptance.",
    )
    parser.add_argument(
        "--sparse-backend",
        choices=("cpu", "cuda", "auto"),
        default="cpu",
    )
    parser.add_argument(
        "--rigid-constraint-linear-solver",
        choices=("augmented_direct", "projected_minres"),
        default="augmented_direct",
    )
    parser.add_argument(
        "--resume",
        action="store_true",
        help="Reuse converged equilibria already present in the output folder.",
    )
    parser.add_argument(
        "--dry-run",
        action="store_true",
        help="Print the planned runs and exit without solving.",
    )


def import_runner(project_root: Path):
    """Import the upstream two-group runner from the real project tree."""
    source_dir = project_root / "05_Source_Data"
    if not (source_dir / "run_reynolds_finan_heterogeneous_two_group.py").exists():
        source_dir = project_root
    if not (source_dir / "run_reynolds_finan_heterogeneous_two_group.py").exists():
        raise SystemExit(
            f"run_reynolds_finan_heterogeneous_two_group.py not found under {project_root}"
        )
    if str(source_dir) not in sys.path:
        sys.path.insert(0, str(source_dir))
    import run_reynolds_finan_heterogeneous_two_group as runner  # noqa: E402

    return runner


def patch_options(runner, **overrides):
    """Wrap runner._options so fields with no CLI flag can still be set.

    Returns the original callable so a caller can restore it if needed.
    """
    original = runner._options

    def patched(*args, **kwargs):
        options = original(*args, **kwargs)
        for field, value in overrides.items():
            if not hasattr(options, field):
                raise SystemExit(
                    f"NonlinearSolveOptions has no field {field!r}; "
                    "the upstream API changed and this script needs updating."
                )
            object.__setattr__(options, field, value)
        return options

    runner._options = patched
    return original


def read_endpoints(run_dir: Path) -> dict[tuple[str, str], dict[str, float]]:
    """Read one run folder's summary CSV keyed by (discretization, condition)."""
    summary = Path(run_dir) / SUMMARY_NAME
    if not summary.exists():
        raise FileNotFoundError(f"missing {summary}")
    rows: dict[tuple[str, str], dict[str, float]] = {}
    with summary.open(encoding="utf-8") as handle:
        for row in csv.DictReader(handle):
            key = (row["discretization"], row["condition"])
            rows[key] = row
    return rows


def paired_change(rows, discretization: str) -> tuple[float, float, float, float]:
    """Return (control, hyperosmotic, absolute difference, percent change)."""
    control = float(rows[(discretization, "Control")][ENDPOINT_COLUMN])
    hyper = float(rows[(discretization, "Hyperosmotic")][ENDPOINT_COLUMN])
    difference = hyper - control
    percent = 100.0 * difference / control if control != 0.0 else float("nan")
    return control, hyper, difference, percent


def discretizations(rows) -> list[str]:
    seen: list[str] = []
    for discretization, _ in rows:
        if discretization not in seen:
            seen.append(discretization)
    return seen
