#!/usr/bin/env python3
"""Tier-1 item 1 -- decompose the predicted unloading into its two sources.

Why this run matters
--------------------
The assumed lamina pretension is 0.05 mN/m, which is 76 % of the 0.066 mN/m
Control endpoint (0.05 mN/m over a 0.12 um shell is a 417 Pa membrane stress
against 551 Pa for the full endpoint). Source-code inspection confirmed the
pretension is a residualised initial stress -- `PK1 += (F - I) . S0` -- so it
contributes exactly zero at u = 0 and is pushed forward into the Cauchy stress
that the reported proxy is built from. The headline number therefore carries the
full pretension contribution, and the epistemic design never tested pretension
below 0.025 mN/m, so this decomposition has never been done.

Three configurations, six equilibria:

  probe_only        pretension = 0      probe = 375 Pa
  pretension_only   pretension = 0.05   probe = 0 Pa
  full              pretension = 0.05   probe = 375 Pa   (reproduces the paper)

The collector then reports how much of the change comes from pretension
relaxation, how much from the probe response of a smaller body, and how much is
the coupling term that neither run captures alone.

Note on `probe_only`: with no pretension and no probe the shell would carry no
load at all, so the pretension-free case is the one that isolates the probe.
Note on `pretension_only`: the runner labels a zero-probe run
"pure_osmotic_natural_volume" internally; that is expected, not an error.
"""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent))
from _common import (  # noqa: E402
    add_common_arguments,
    import_runner,
    paired_change,
    read_endpoints,
)

CASES = (
    {
        "name": "probe_only",
        "pretension_mn_per_m": 0.0,
        "probe_shear_stress_pa": 375.0,
        "role": "Isolates the probe response of the shrunken geometry.",
    },
    {
        "name": "pretension_only",
        "pretension_mn_per_m": 0.05,
        "probe_shear_stress_pa": 0.0,
        "role": "Isolates relaxation of the assumed homeostatic pretension.",
    },
    {
        "name": "full",
        "pretension_mn_per_m": 0.05,
        "probe_shear_stress_pa": 375.0,
        "role": "Reproduces the published configuration.",
    },
)


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    add_common_arguments(parser)
    parser.add_argument(
        "--pretension-mn-per-m",
        type=float,
        default=0.05,
        help="Nominal pretension used by the pretension_only and full cases.",
    )
    parser.add_argument(
        "--probe-shear-stress-pa",
        type=float,
        default=375.0,
        help="Nominal probe shear stress used by the probe_only and full cases.",
    )
    parser.add_argument(
        "--element-order",
        choices=("p1", "p2"),
        default="p2",
        help="p2 also solves the intermediate-mesh P2 pair (recommended).",
    )
    args = parser.parse_args()

    cases = []
    for case in CASES:
        resolved = dict(case)
        if resolved["pretension_mn_per_m"] != 0.0:
            resolved["pretension_mn_per_m"] = args.pretension_mn_per_m
        if resolved["probe_shear_stress_pa"] != 0.0:
            resolved["probe_shear_stress_pa"] = args.probe_shear_stress_pa
        cases.append(resolved)

    print("Mechanism decomposition -- planned runs")
    print(f"  project root : {args.project_root}")
    print(f"  output root  : {args.out_root}")
    print(f"  primary mesh : {args.primary_mesh}")
    for case in cases:
        print(
            f"  {case['name']:<16} pretension={case['pretension_mn_per_m']:<6} "
            f"probe={case['probe_shear_stress_pa']:<6} -- {case['role']}"
        )
    if args.dry_run:
        print("\n--dry-run: nothing solved.")
        return 0

    runner = import_runner(args.project_root)
    args.out_root.mkdir(parents=True, exist_ok=True)

    manifest = {"cases": [], "primary_mesh": args.primary_mesh}
    for case in cases:
        out_dir = args.out_root / case["name"]
        print(f"\n=== {case['name']} -> {out_dir} ===", flush=True)
        runner.run_analysis(
            out_dir,
            primary_mesh=args.primary_mesh,
            probe_shear_stress_pa=case["probe_shear_stress_pa"],
            pretension_mn_per_m=case["pretension_mn_per_m"],
            run_intermediate_p2=(args.element_order == "p2"),
            resume=args.resume,
            sparse_backend=args.sparse_backend,
            rigid_constraint_linear_solver=args.rigid_constraint_linear_solver,
        )
        manifest["cases"].append({**case, "out_dir": str(out_dir)})

    # Summarise immediately so a long run leaves a readable artefact behind.
    discretization = f"{args.primary_mesh}_{'P2' if args.element_order == 'p2' else 'P1'}"
    results = {}
    for case in cases:
        try:
            rows = read_endpoints(args.out_root / case["name"])
            control, hyper, difference, percent = paired_change(rows, discretization)
        except (FileNotFoundError, KeyError) as exc:
            print(f"  warning: could not read {case['name']}: {exc}")
            continue
        results[case["name"]] = {
            "control_mn_per_m": control,
            "hyperosmotic_mn_per_m": hyper,
            "absolute_difference_mn_per_m": difference,
            "percent_change": percent,
        }

    if {"probe_only", "pretension_only", "full"} <= results.keys():
        full = results["full"]["absolute_difference_mn_per_m"]
        probe = results["probe_only"]["absolute_difference_mn_per_m"]
        pretension = results["pretension_only"]["absolute_difference_mn_per_m"]
        coupling = full - probe - pretension
        results["decomposition"] = {
            "discretization": discretization,
            "full_difference_mn_per_m": full,
            "probe_share_percent": 100.0 * probe / full if full else float("nan"),
            "pretension_share_percent": 100.0 * pretension / full if full else float("nan"),
            "coupling_share_percent": 100.0 * coupling / full if full else float("nan"),
            "note": (
                "Shares are of the total paired difference. A large pretension "
                "share means the headline result is dominated by an assumption "
                "with no literature provenance and must be reported as such."
            ),
        }

    manifest["results"] = results
    output = args.out_root / "mechanism_decomposition.json"
    output.write_text(json.dumps(manifest, indent=2), encoding="utf-8")
    print(f"\nWrote {output}")
    return 0


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