"""Aggregate GR analogue diagnostics across seeds and emit summary artifacts."""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Iterable, List

ROOT = Path(__file__).resolve().parents[1]
DEFAULT_BASE = ROOT / "artifacts" / "relativistic_suite"
DEFAULT_FEEDS = ROOT / "artifacts" / "gr_suite"
SEEDS_DEFAULT = (7, 11, 13, 17, 19)

if str(ROOT) not in sys.path:
    sys.path.insert(0, str(ROOT))

from scripts.gr_hotspot_report import generate_hotspot_reports  # type: ignore
from scripts.gr_metric_analysis import DEFAULT_WINDOW  # type: ignore


@dataclass
class SeedMetrics:
    seed: int
    eigenvalue_max: float
    condition_number: float
    throat_frame: int
    throat_metric: float
    throat_metric_scaled: float
    throat_euclidean: float
    curvature_median: float
    curvature_change_points: List[int]
    residual_median: float
    residual_change_points: List[int]
    divergence_median: float
    divergence_max: float

    @classmethod
    def from_payload(cls, seed: int, payload: dict) -> "SeedMetrics":
        metric = payload["metric"]
        throat = metric["metric_minima_frames"]
        curvature = payload["curvature"]
        residual = payload["einstein_residual"]
        divergence = payload["divergence_proxy"]
        return cls(
            seed=seed,
            eigenvalue_max=float(metric["eigenvalue_max"]),
            condition_number=float(metric["condition_number"]),
            throat_frame=int(throat["metric_argmin_frame"]),
            throat_metric=float(throat["metric_min"]),
            throat_metric_scaled=float(throat["metric_min_scaled"]),
            throat_euclidean=float(throat["euclidean_min"]),
            curvature_median=float(curvature["median"]),
            curvature_change_points=[int(x) for x in curvature.get("change_points", [])],
            residual_median=float(residual["median"]),
            residual_change_points=[int(x) for x in residual.get("change_points", [])],
            divergence_median=float(divergence.get("median", 0.0)),
            divergence_max=float(divergence.get("max", 0.0)),
        )


def load_seed_metrics(base: Path, seeds: Iterable[int]) -> tuple[float, List[SeedMetrics]]:
    kappa: float | None = None
    metrics: List[SeedMetrics] = []
    for seed in seeds:
        path = base / f"gr_metric_seed{seed}.json"
        payload = json.loads(path.read_text(encoding="utf-8"))
        if kappa is None:
            kappa = float(payload["kappa"])
        metrics.append(SeedMetrics.from_payload(seed, payload))
    if kappa is None:
        raise RuntimeError("No diagnostics found; cannot infer κ scale")
    return kappa, metrics


def emit_summary(
    base: Path,
    kappa: float,
    metrics: List[SeedMetrics],
    feed_dir: Path | None,
    hotspot_window: int,
) -> None:
    summary_payload = {
        "kappa": kappa,
        "seeds": [metric.__dict__ for metric in metrics],
        "hotspots": {
            str(metric.seed): {
                "frames": sorted(set(metric.curvature_change_points + metric.residual_change_points)),
                "divergence_max": metric.divergence_max,
                "divergence_median": metric.divergence_median,
            }
            for metric in metrics
        },
    }
    summary_path = base / "gr_metric_summary.json"
    summary_path.write_text(json.dumps(summary_payload, indent=2), encoding="utf-8")

    header = [
        "# GR Analogue Diagnostics",
        "",
        f"κ scale: {kappa:.5f}",
        "",
        "| seed | eig_max | cond | throat_frame | throat_metric | throat_scaled | euclid_min | curvature_median | curvature_spikes | residual_median | residual_spikes | divergence_med | divergence_max |",
        "|------|---------|------|--------------|---------------|---------------|------------|------------------|------------------|-----------------|-----------------|----------------|----------------|",
    ]
    rows = [
        "| {seed} | {eig:.3e} | {cond:.3e} | {frame} | {throat:.3e} | {scaled:.3e} | {euclid:.3f} | {curv:.3e} | {curv_spikes} | {res:.3e} | {res_spikes} | {div_med:.3e} | {div:.3e} |".format(
            seed=metric.seed,
            eig=metric.eigenvalue_max,
            cond=metric.condition_number,
            frame=metric.throat_frame,
            throat=metric.throat_metric,
            scaled=metric.throat_metric_scaled,
            euclid=metric.throat_euclidean,
            curv=metric.curvature_median,
            curv_spikes=",".join(map(str, metric.curvature_change_points)) or "-",
            res=metric.residual_median,
            res_spikes=",".join(map(str, metric.residual_change_points)) or "-",
            div_med=metric.divergence_median,
            div=metric.divergence_max,
        )
        for metric in metrics
    ]
    (base / "gr_metric_report.md").write_text(
        "\n".join(header + rows) + "\n",
        encoding="utf-8",
    )

    if feed_dir is not None:
        hotspot_json = base / "gr_hotspots.json"
        hotspot_md = base / "gr_hotspots.md"
        generate_hotspot_reports(
            summary_path=summary_path,
            feeds_path=feed_dir,
            window=hotspot_window,
            json_path=hotspot_json,
            md_path=hotspot_md,
        )


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Aggregate GR diagnostics and emit summaries")
    parser.add_argument(
        "--base",
        type=Path,
        default=DEFAULT_BASE,
        help="Directory containing gr_metric_seed*.json files",
    )
    parser.add_argument(
        "--feeds",
        type=Path,
        default=DEFAULT_FEEDS,
        help="Directory containing per-seed wireframe feeds for hotspot extraction",
    )
    parser.add_argument(
        "--seeds",
        type=int,
        nargs="*",
        default=list(SEEDS_DEFAULT),
        help="Seeds to aggregate",
    )
    parser.add_argument(
        "--hotspot-window",
        type=int,
        default=DEFAULT_WINDOW,
        help="Window size used for diagnostics (controls hotspot extraction smoothing)",
    )
    parser.add_argument(
        "--skip-hotspots",
        action="store_true",
        help="Do not regenerate hotspot payloads",
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    base = args.base.resolve()
    base.mkdir(parents=True, exist_ok=True)

    kappa, metrics = load_seed_metrics(base, args.seeds)

    feed_dir: Path | None = None if args.skip_hotspots else args.feeds.resolve()
    if feed_dir is not None and not feed_dir.exists():
        raise FileNotFoundError(f"Feed directory {feed_dir} not found for hotspot extraction")

    emit_summary(
        base=base,
        kappa=kappa,
        metrics=metrics,
        feed_dir=feed_dir,
        hotspot_window=args.hotspot_window,
    )


if __name__ == "__main__":
    main()
