"""Generate hotspot diagnostics summary across seeds and frames.

This utility cross-references the GR metric summary hotspots with the
underlying wireframe feeds to extract detailed per-frame metrics.
"""

from __future__ import annotations

import argparse
import json
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Dict, List, Tuple
import sys

ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
    sys.path.insert(0, str(ROOT))

from scripts.gr_metric_analysis import (  # type: ignore
    DEFAULT_EPSILON,
    DEFAULT_RIDGE,
    DEFAULT_WINDOW,
    compute_frame_diagnostics,
    load_wormhole_series,
)

DEFAULT_SUMMARY = ROOT / "artifacts" / "relativistic_suite" / "gr_metric_summary.json"
DEFAULT_FEEDS = ROOT / "artifacts" / "gr_suite"
OUTPUT_JSON = ROOT / "artifacts" / "relativistic_suite" / "gr_hotspots.json"
OUTPUT_MD = ROOT / "artifacts" / "relativistic_suite" / "gr_hotspots.md"


@dataclass
class FrameSnapshot:
    frame: int
    throat_metric: float | None
    throat_euclid: float | None
    curvature: float | None
    stress: float | None
    residual: float | None
    divergence: float | None
    hp_match: float | None


@dataclass
class SeedHotspot:
    seed: int
    frames: List[int]
    divergence_median: float
    divergence_max: float
    snapshots: List[FrameSnapshot]


def load_summary(path: Path) -> Dict:
    return json.loads(path.read_text(encoding="utf-8"))


def extract_hotspots(
    feed_dir: Path,
    seed: int,
    frames: List[int],
    kappa: float,
    window_radius: int,
) -> List[FrameSnapshot]:
    feed_path = feed_dir / f"seed_{seed}" / "wireframe.json"
    if not feed_path.exists():
        raise FileNotFoundError(f"Feed for seed {seed} not found at {feed_path}")

    series = load_wormhole_series(feed_path)
    if series.frames <= max(frames, default=-1):
        raise ValueError(f"Seed {seed} feed shorter than hotspot frames {frames}")

    diagnostics, arrays = compute_frame_diagnostics(
        series,
        window_radius=window_radius,
        eps=DEFAULT_EPSILON,
        ridge=DEFAULT_RIDGE,
        kappa=kappa,
    )

    hp_trace = series.hp_trace
    snapshots: List[FrameSnapshot] = []
    for frame in sorted(frames):
        diag = diagnostics[frame]
        hp_match = None
        if hp_trace.size > frame:
            value = hp_trace[frame]
            hp_match = float(value) if value == value else None  # handle NaN

        snapshots.append(
            FrameSnapshot(
                frame=frame,
                throat_metric=_sanitize(diag.throat_metric * kappa if kappa else diag.throat_metric),
                throat_euclid=_sanitize(diag.throat_euclid),
                curvature=_sanitize(diag.curvature_dir),
                stress=_sanitize(diag.stress_scalar),
                residual=_sanitize(diag.einstein_residual),
                divergence=_sanitize(diag.divergence),
                hp_match=hp_match,
            )
        )
    return snapshots


def assemble_hotspots(summary: Dict, feeds: Path, window_radius: int) -> Tuple[float, List[SeedHotspot]]:
    hotspots_block = summary.get("hotspots") or {}
    seeds_info = summary.get("seeds") or []
    kappa = float(summary.get("kappa", 1.0))

    hotspots: List[SeedHotspot] = []
    for seed_entry in seeds_info:
        seed = int(seed_entry["seed"])
        hotspot_info = hotspots_block.get(str(seed))
        if not hotspot_info:
            continue
        frames = [int(f) for f in hotspot_info.get("frames", [])]
        if not frames:
            continue
        snapshots = extract_hotspots(feeds, seed, frames, kappa, window_radius)
        divergence_median = float(hotspot_info.get("divergence_median", float("nan")))
        divergence_max = float(hotspot_info.get("divergence_max", float("nan")))
        hotspots.append(
            SeedHotspot(
                seed=seed,
                frames=frames,
                divergence_median=divergence_median,
                divergence_max=divergence_max,
                snapshots=snapshots,
            )
        )

    return kappa, hotspots


def render_payload(hotspots: List[SeedHotspot], kappa: float) -> Dict[str, object]:
    return {
        "kappa": kappa,
        "hotspots": [
            {
                "seed": entry.seed,
                "frames": entry.frames,
                "divergence_median": entry.divergence_median,
                "divergence_max": entry.divergence_max,
                "snapshots": [asdict(snapshot) for snapshot in entry.snapshots],
            }
            for entry in hotspots
        ],
    }


def render_markdown(hotspots: List[SeedHotspot], kappa: float) -> str:
    lines: List[str] = ["# GR Hotspot Report", "", f"κ scale: {kappa:.5f}", ""]
    for entry in hotspots:
        lines.append(f"## Seed {entry.seed}")
        lines.append(
            f"Hotspot frames: {', '.join(map(str, entry.frames)) or '—'} (divergence median {entry.divergence_median:.3e}, max {entry.divergence_max:.3e})"
        )
        lines.append("")
        lines.append(
            "| frame | throat_metric | throat_euclid | curvature | stress | residual | divergence | hp_match |"
        )
        lines.append("|-------|---------------|---------------|-----------|--------|----------|-----------|----------|")
        for snapshot in entry.snapshots:
            lines.append(
                "| {frame} | {throat_metric} | {throat_euclid} | {curvature} | {stress} | {residual} | {divergence} | {hp} |".format(
                    frame=snapshot.frame,
                    throat_metric=_fmt(snapshot.throat_metric),
                    throat_euclid=_fmt(snapshot.throat_euclid),
                    curvature=_fmt(snapshot.curvature),
                    stress=_fmt(snapshot.stress),
                    residual=_fmt(snapshot.residual),
                    divergence=_fmt(snapshot.divergence),
                    hp="{:.3f}".format(snapshot.hp_match) if snapshot.hp_match is not None else "—",
                )
            )
        lines.append("")
    return "\n".join(lines).strip() + "\n"


def emit_reports(hotspots: List[SeedHotspot], kappa: float, json_path: Path, md_path: Path) -> Dict[str, object]:
    payload = render_payload(hotspots, kappa)
    json_path.write_text(json.dumps(payload, indent=2, allow_nan=False), encoding="utf-8")
    md_path.write_text(render_markdown(hotspots, kappa), encoding="utf-8")
    return payload


def generate_hotspot_reports(
    summary_path: Path = DEFAULT_SUMMARY,
    feeds_path: Path = DEFAULT_FEEDS,
    window: int = DEFAULT_WINDOW,
    json_path: Path = OUTPUT_JSON,
    md_path: Path = OUTPUT_MD,
) -> Dict[str, object]:
    summary = load_summary(summary_path)
    kappa, hotspots = assemble_hotspots(summary, feeds_path, window // 2)
    return emit_reports(hotspots, kappa, json_path, md_path)


def _sanitize(value: float | None) -> float | None:
    if value is None:
        return None
    if isinstance(value, (float, int)):
        if value != value or value in (float("inf"), float("-inf")):
            return None
        return float(value)
    return None


def _fmt(value: float | None) -> str:
    return f"{value:.3e}" if value is not None else "—"


def main() -> None:
    parser = argparse.ArgumentParser(description="Assemble hotspot diagnostics across seeds")
    parser.add_argument("--summary", type=Path, default=DEFAULT_SUMMARY)
    parser.add_argument("--feeds", type=Path, default=DEFAULT_FEEDS)
    parser.add_argument("--window", type=int, default=DEFAULT_WINDOW)
    args = parser.parse_args()

    generate_hotspot_reports(
        summary_path=args.summary,
        feeds_path=args.feeds,
        window=args.window,
        json_path=OUTPUT_JSON,
        md_path=OUTPUT_MD,
    )
    print(f"Hotspot report written to {OUTPUT_JSON} and {OUTPUT_MD}")


if __name__ == "__main__":
    main()
