#!/usr/bin/env python3
"""Build the revised Table S2 and the decomposition/control summaries.

Two things this fixes in the current manuscript:

1. Table S2 reports endpoints but never the paired change, so the reader cannot
   see that the change runs from -37.9 % to -44.6 % across the five
   discretisations. This adds that column and splits the P1 and P2 families,
   which is what shows P1 still drifting while P2 has settled.

2. The h-refinement claim is only valid for Control. The hyperosmotic P1
   sequence is non-monotone (0.037475 -> 0.038412 -> 0.038092, the increment
   changes sign), so Richardson extrapolation and GCI do not apply and this
   script refuses to print them rather than printing a meaningless number.

Usage:
    python3 06_collect_results.py --runs RUN_DIR [RUN_DIR ...] \
        [--decomposition mechanism_decomposition.json] \
        [--homogeneous homogeneous_control.json] \
        [--markdown-out revised_tables.md]
"""

from __future__ import annotations

import argparse
import json
import math
import sys
from pathlib import Path

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

# Mesh label -> representative element size in um, used only for ordering and
# for the refinement ratio in the Richardson estimate.
MESH_SIZE_UM = {
    "coarse_h1p25": 1.25,
    "intermediate_h0p75": 0.75,
    "main_h0p45": 0.45,
}


def _mesh_and_order(discretization: str) -> tuple[str, str]:
    if discretization.endswith("_P1"):
        return discretization[:-3], "P1"
    if discretization.endswith("_P2"):
        return discretization[:-3], "P2"
    return discretization, "?"


def richardson(values: list[tuple[float, float]]) -> dict | None:
    """Richardson estimate from three (h, f) pairs, or None if non-monotone."""
    if len(values) < 3:
        return None
    values = sorted(values, key=lambda item: -item[0])
    (h1, f1), (h2, f2), (h3, f3) = values[-3:]
    first, second = f2 - f1, f3 - f2
    if second == 0.0 or first / second <= 0.0:
        return {"valid": False, "reason": "non-monotone increments; GCI undefined"}
    ratio = h1 / h2
    if not math.isclose(ratio, h2 / h3, rel_tol=0.05):
        return {"valid": False, "reason": "non-uniform refinement ratio"}
    order = math.log(first / second) / math.log(ratio)
    scale = ratio**order
    extrapolated = f3 + second / (scale - 1.0)
    gci = 1.25 * abs(second / f3) / (scale - 1.0)
    return {
        "valid": True,
        "observed_order": order,
        "extrapolated": extrapolated,
        "gci_percent": 100.0 * gci,
    }


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--runs", type=Path, nargs="+", required=True)
    parser.add_argument("--decomposition", type=Path, default=None)
    parser.add_argument("--homogeneous", type=Path, default=None)
    parser.add_argument("--markdown-out", type=Path, default=None)
    args = parser.parse_args()

    lines: list[str] = []

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

    merged: dict = {}
    for run in args.runs:
        try:
            merged.update(read_endpoints(run))
        except FileNotFoundError as exc:
            print(f"warning: {exc}", file=sys.stderr)
    if not merged:
        raise SystemExit("No summary CSVs could be read.")

    emit("## Revised Table S2 -- endpoints with the paired change")
    emit()
    emit("| discretisation | dofs | Control (mN/m) | Hyperosmotic (mN/m) | difference (mN/m) | change |")
    emit("|---|---:|---:|---:|---:|---:|")

    families: dict[str, list[tuple[float, float]]] = {}
    changes: list[float] = []
    for discretization in sorted(
        discretizations(merged),
        key=lambda name: (
            _mesh_and_order(name)[1],
            -MESH_SIZE_UM.get(_mesh_and_order(name)[0], 0.0),
        ),
    ):
        try:
            control, hyper, difference, percent = paired_change(merged, discretization)
        except KeyError:
            continue
        dofs = merged[(discretization, "Control")].get("dofs", "")
        emit(
            f"| {discretization} | {dofs} | {control:.6f} | {hyper:.6f} | "
            f"{difference:+.6f} | {percent:+.2f}% |"
        )
        changes.append(percent)
        mesh, order = _mesh_and_order(discretization)
        if mesh in MESH_SIZE_UM:
            families.setdefault(f"Control/{order}", []).append((MESH_SIZE_UM[mesh], control))
            families.setdefault(f"Hyperosmotic/{order}", []).append((MESH_SIZE_UM[mesh], hyper))

    if changes:
        emit()
        emit(
            f"Spread across tested discretisations: {min(changes):+.2f}% to "
            f"{max(changes):+.2f}% ({max(changes) - min(changes):.1f} pp). Report "
            "the endpoint to no more precision than this spread supports."
        )

    emit()
    emit("## Convergence status by family")
    emit()
    for name, series in sorted(families.items()):
        if len(series) < 3:
            ordered = ", ".join(f"h={h:g}: {f:.6f}" for h, f in sorted(series, key=lambda i: -i[0]))
            emit(f"- **{name}** -- {len(series)} mesh(es) ({ordered}); "
                 "too few for an h-study, quote the spread instead.")
            continue
        estimate = richardson(series)
        if estimate and estimate["valid"]:
            emit(
                f"- **{name}** -- observed order {estimate['observed_order']:.2f}, "
                f"Richardson limit {estimate['extrapolated']:.6f}, "
                f"GCI {estimate['gci_percent']:.2f}%."
            )
        else:
            reason = estimate["reason"] if estimate else "insufficient data"
            emit(
                f"- **{name}** -- NOT convergent in the Richardson sense "
                f"({reason}). Do not quote a GCI; state the non-monotonicity and "
                "rely on cross-order agreement instead."
            )

    if args.decomposition and args.decomposition.exists():
        data = json.loads(args.decomposition.read_text(encoding="utf-8"))
        decomposition = data.get("results", {}).get("decomposition")
        emit()
        emit("## Mechanism decomposition")
        emit()
        if decomposition:
            emit(f"At {decomposition['discretization']}, of the total paired difference "
                 f"({decomposition['full_difference_mn_per_m']:+.6f} mN/m):")
            emit()
            emit(f"- pretension relaxation: {decomposition['pretension_share_percent']:.1f}%")
            emit(f"- probe response of the shrunken body: {decomposition['probe_share_percent']:.1f}%")
            emit(f"- coupling between the two: {decomposition['coupling_share_percent']:.1f}%")
            emit()
            emit(
                "Write this into the Results. If pretension dominates, the honest "
                "sentence is that the predicted unloading is mainly relaxation of "
                "an assumed homeostatic prestress with no literature provenance."
            )
        else:
            emit("Decomposition JSON present but incomplete; check the run logs.")

    if args.homogeneous and args.homogeneous.exists():
        data = json.loads(args.homogeneous.read_text(encoding="utf-8"))
        emit()
        emit("## Homogeneous-core control")
        emit()
        emit("| variant | core E (Pa) | change |")
        emit("|---|---:|---:|")
        for label, entry in data.get("results", {}).items():
            emit(f"| {label} | {entry['core_young_pa']:.1f} | {entry['percent_change']:+.2f}% |")
        emit()
        emit(
            "Compare against the heterogeneous run at the same discretisation. If "
            "the gap is smaller than the same-mesh P1/P2 difference, the global "
            "endpoint cannot resolve core heterogeneity and the manuscript must "
            "say so rather than implying the chromatin map drives the result."
        )

    if args.markdown_out:
        args.markdown_out.write_text("\n".join(lines) + "\n", encoding="utf-8")
        print(f"\nWrote {args.markdown_out}", file=sys.stderr)
    return 0


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