"""Evaluate sensitivity of GR diagnostics to Jacobian window size.

This utility replays the GR metric analysis at multiple temporal window widths
for each seed, enabling comparison of curvature, residual, and divergence
statistics relative to the canonical configuration.
"""

from __future__ import annotations

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

import numpy as np

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,
    assemble_summary,
    compute_frame_diagnostics,
    load_wormhole_series,
)

DEFAULT_WINDOWS = (5, DEFAULT_WINDOW, 13)
DEFAULT_SEEDS = (7, 11, 13, 17, 19)
DEFAULT_FEEDS = ROOT / "artifacts" / "gr_suite"
OUTPUT_JSON = ROOT / "artifacts" / "relativistic_suite" / "window_sweep.json"
OUTPUT_MD = ROOT / "artifacts" / "relativistic_suite" / "window_sweep.md"


@dataclass
class WindowSummary:
    window: int
    throat_scaled_min: float
    throat_frame: int
    condition_number: float
    curvature_median: float
    curvature_p95: float
    curvature_change_points: List[int]
    residual_median: float
    residual_p95: float
    residual_change_points: List[int]
    divergence_median: float
    divergence_max: float

    @classmethod
    def from_payload(cls, window: int, summary: Dict[str, object]) -> "WindowSummary":
        metric = summary["metric"]
        throat = metric["metric_minima_frames"]
        curvature = summary["curvature"]
        residual = summary["einstein_residual"]
        divergence = summary["divergence_proxy"]
        return cls(
            window=window,
            throat_scaled_min=float(throat["metric_min_scaled"]),
            throat_frame=int(throat["metric_argmin_frame"]),
            condition_number=float(metric["condition_number"]),
            curvature_median=float(curvature["median"]),
            curvature_p95=float(curvature["p95"]),
            curvature_change_points=[int(idx) for idx in curvature.get("change_points", [])],
            residual_median=float(residual["median"]),
            residual_p95=float(residual["p95"]),
            residual_change_points=[int(idx) for idx in residual.get("change_points", [])],
            divergence_median=float(divergence.get("median", float("nan"))),
            divergence_max=float(divergence.get("max", float("nan"))),
        )

    def to_dict(self) -> Dict[str, object]:
        return {
            "window": self.window,
            "throat_scaled_min": self.throat_scaled_min,
            "throat_frame": self.throat_frame,
            "condition_number": self.condition_number,
            "curvature": {
                "median": self.curvature_median,
                "p95": self.curvature_p95,
                "change_points": self.curvature_change_points,
            },
            "residual": {
                "median": self.residual_median,
                "p95": self.residual_p95,
                "change_points": self.residual_change_points,
            },
            "divergence": {
                "median": self.divergence_median,
                "max": self.divergence_max,
            },
        }


@dataclass
class WindowComparison:
    baseline: WindowSummary
    variants: List[Tuple[WindowSummary, Dict[str, float]]]

    def to_dict(self) -> Dict[str, object]:
        return {
            "baseline": self.baseline.to_dict(),
            "variants": [
                {
                    "window": summary.window,
                    "metrics": summary.to_dict(),
                    "delta_vs_baseline": deltas,
                }
                for summary, deltas in self.variants
            ],
        }


def _validate_windows(windows: Iterable[int]) -> List[int]:
    cleaned: List[int] = []
    for window in windows:
        if window < 3:
            raise ValueError(f"Window size must be >=3, got {window}")
        if window % 2 == 0:
            raise ValueError(f"Window size must be odd to define a centre frame (got {window})")
        cleaned.append(window)
    return sorted(set(cleaned))


def _compute_summary_for_window(
    series_path: Path,
    window: int,
    kappa: float,
) -> WindowSummary:
    series = load_wormhole_series(series_path)
    diagnostics, arrays = compute_frame_diagnostics(
        series,
        window_radius=window // 2,
        eps=DEFAULT_EPSILON,
        ridge=DEFAULT_RIDGE,
        kappa=kappa,
    )
    summary = assemble_summary(series, arrays, kappa)
    return WindowSummary.from_payload(window, summary)


def _compute_seed_windows(
    feeds_dir: Path,
    seed: int,
    windows: List[int],
    kappa: float,
) -> WindowComparison:
    feed_path = feeds_dir / f"seed_{seed}" / "wireframe.json"
    if not feed_path.exists():
        raise FileNotFoundError(f"Missing feed for seed {seed}: {feed_path}")

    window_summaries = {
        window: _compute_summary_for_window(feed_path, window, kappa) for window in windows
    }

    baseline_window = DEFAULT_WINDOW if DEFAULT_WINDOW in window_summaries else windows[0]
    baseline = window_summaries[baseline_window]

    variants: List[Tuple[WindowSummary, Dict[str, float]]] = []
    for window, summary in sorted(window_summaries.items()):
        if window == baseline_window:
            continue
        deltas = _compute_deltas(baseline, summary)
        variants.append((summary, deltas))

    return WindowComparison(baseline=baseline, variants=variants)


def _compute_deltas(baseline: WindowSummary, candidate: WindowSummary) -> Dict[str, float]:
    def pct_delta(new: float, base: float) -> float:
        if not np.isfinite(base) or abs(base) < 1e-12:
            return float("nan")
        return (new - base) / base

    return {
        "curvature_median_delta": candidate.curvature_median - baseline.curvature_median,
        "curvature_median_pct": pct_delta(candidate.curvature_median, baseline.curvature_median),
        "divergence_median_delta": candidate.divergence_median - baseline.divergence_median,
        "divergence_median_pct": pct_delta(candidate.divergence_median, baseline.divergence_median),
        "divergence_max_delta": candidate.divergence_max - baseline.divergence_max,
        "divergence_max_pct": pct_delta(candidate.divergence_max, baseline.divergence_max),
    }


def _emit_reports(
    results: Dict[int, WindowComparison],
    kappa: float,
    windows: List[int],
    json_path: Path,
    markdown_path: Path,
) -> None:
    payload = {
        "kappa": kappa,
        "windows": windows,
        "seeds": {
            str(seed): comparison.to_dict() for seed, comparison in results.items()
        },
    }
    json_path.write_text(json.dumps(payload, indent=2), encoding="utf-8")

    lines: List[str] = ["# Jacobian Window Sensitivity", "", f"κ scale: {kappa:.5f}", ""]
    for seed, comparison in results.items():
        lines.append(f"## Seed {seed}")
        lines.append(
            f"Baseline window {comparison.baseline.window}: curvature median {comparison.baseline.curvature_median:.3e}, divergence median {comparison.baseline.divergence_median:.3e}"
        )
        lines.append("")
        lines.append(
            "| window | curvature_median | divergence_median | divergence_max | Δ curvature (abs / %) | Δ divergence_med (abs / %) | Δ divergence_max (abs / %) | change_points (curv / residual) |"
        )
        lines.append(
            "|--------|------------------|-------------------|----------------|------------------------|-----------------------------|-----------------------------|---------------------------------|"
        )
        for summary, deltas in comparison.variants:
            lines.append(
                "| {window} | {curv_med:.3e} | {div_med:.3e} | {div_max:.3e} | {dc:.3e} / {dcpct} | {dd:.3e} / {ddpct} | {dx:.3e} / {dxpct} | {cpts} / {rpts} |".format(
                    window=summary.window,
                    curv_med=summary.curvature_median,
                    div_med=summary.divergence_median,
                    div_max=summary.divergence_max,
                    dc=deltas["curvature_median_delta"],
                    dcpct=_fmt_pct(deltas["curvature_median_pct"]),
                    dd=deltas["divergence_median_delta"],
                    ddpct=_fmt_pct(deltas["divergence_median_pct"]),
                    dx=deltas["divergence_max_delta"],
                    dxpct=_fmt_pct(deltas["divergence_max_pct"]),
                    cpts=",".join(map(str, summary.curvature_change_points)) or "-",
                    rpts=",".join(map(str, summary.residual_change_points)) or "-",
                )
            )
        lines.append("")
    markdown_path.write_text("\n".join(lines).strip() + "\n", encoding="utf-8")


def _fmt_pct(value: float) -> str:
    if not np.isfinite(value):
        return "—"
    return f"{value * 100:.2f}%"


def main() -> None:
    parser = argparse.ArgumentParser(description="Sweep GR diagnostics across Jacobian windows")
    parser.add_argument("--feeds", type=Path, default=DEFAULT_FEEDS, help="Directory holding seed feeds")
    parser.add_argument("--windows", type=int, nargs="*", default=list(DEFAULT_WINDOWS), help="Odd window sizes to evaluate")
    parser.add_argument("--seeds", type=int, nargs="*", default=list(DEFAULT_SEEDS), help="Seeds to evaluate")
    parser.add_argument("--kappa", type=float, default=5.20208, help="κ calibration factor")
    parser.add_argument("--json", type=Path, default=OUTPUT_JSON, help="Output JSON path")
    parser.add_argument("--markdown", type=Path, default=OUTPUT_MD, help="Output Markdown path")
    args = parser.parse_args()

    windows = _validate_windows(args.windows)
    feeds_dir = args.feeds
    results: Dict[int, WindowComparison] = {}
    for seed in args.seeds:
        comparison = _compute_seed_windows(feeds_dir, seed, windows, args.kappa)
        results[seed] = comparison
        print(f"Seed {seed}: baseline window {comparison.baseline.window} evaluated with {len(comparison.variants)} variants")

    _emit_reports(results, args.kappa, windows, args.json, args.markdown)
    print(f"Window sweep written to {args.json} and {args.markdown}")


if __name__ == "__main__":
    main()
