"""State-dependent GR analogue diagnostics for the smooth wormhole feed.

This module replaces the older ``gr_analogue_check.py`` script by building a
frame-wise, state-dependent metric derived from the local Jacobian between the
receiver state and the packet trace.  From that metric we derive directional
curvature, stress-energy analogues, and Einstein-like residuals which sharpen
our view of where the salience wormhole agrees with (or diverges from)
relativistic expectations.

Example
-------
python scripts/gr_metric_analysis.py \
    --feed artifacts/smooth_gemini/wireframe.json \
    --kappa 418.0 \
    --window 9 \
    --out artifacts/relativistic_suite/gr_metric_seed7.json
"""

from __future__ import annotations

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

import numpy as np

DEFAULT_WINDOW = 9
DEFAULT_EPSILON = 1e-6
DEFAULT_RIDGE = 1e-3


@dataclass
class WormholeSeries:
    """Container for the geometry and HP telemetry of a wormhole regime."""

    state_a: np.ndarray  # (frames, dim)
    state_b: np.ndarray  # (frames, dim)
    packet: np.ndarray  # (frames, dim)
    hp_trace: np.ndarray  # (frames,)
    metrics: Dict[str, float]

    @property
    def frames(self) -> int:
        return self.state_a.shape[0]

    @property
    def dim(self) -> int:
        return self.state_a.shape[1]


@dataclass
class FrameDiagnostics:
    eigenvalues: np.ndarray  # (dim,)
    throat_metric: float
    throat_euclid: float
    curvature_dir: float
    stress_scalar: float
    einstein_residual: float
    divergence: float


def load_wormhole_series(feed_path: Path) -> WormholeSeries:
    data = json.loads(feed_path.read_text(encoding="utf-8"))
    try:
        wormhole = data["wormhole"][0]
        regime = next(r for r in wormhole["regimes"] if r["name"] == "wormhole")
    except (KeyError, StopIteration, IndexError) as exc:
        raise ValueError("Feed must contain wormhole.regimes[?].name == 'wormhole'") from exc

    geometry = regime.get("geometry") or {}
    state_a = np.asarray(geometry.get("state_a"), dtype=float)
    state_b = np.asarray(geometry.get("state_b"), dtype=float)
    packet = np.asarray(geometry.get("packet"), dtype=float)

    if not state_a.size or state_a.shape != state_b.shape or state_a.shape != packet.shape:
        raise ValueError("Geometry arrays state_a, state_b, packet must be aligned and non-empty")

    hp_trace = np.asarray([entry.get("match", np.nan) for entry in regime.get("hp_trace", [])], dtype=float)
    metrics = regime.get("metrics", {})

    return WormholeSeries(state_a=state_a, state_b=state_b, packet=packet, hp_trace=hp_trace, metrics=metrics)


def window_indices(center: int, upper: int, radius: int) -> np.ndarray:
    lo = max(0, center - radius)
    hi = min(upper, center + radius + 1)
    return np.arange(lo, hi, dtype=int)


def local_jacobian(
    state_window: np.ndarray,
    packet_window: np.ndarray,
    ridge: float,
    eps: float,
) -> np.ndarray:
    """Return the local Jacobian d(packet)/d(state_b) estimated via ridge regression."""

    # Center the window to avoid intercept leakage.
    state_centered = state_window - state_window.mean(axis=0, keepdims=True)
    packet_centered = packet_window - packet_window.mean(axis=0, keepdims=True)

    xtx = state_centered.T @ state_centered
    dim = xtx.shape[0]
    scale = np.trace(xtx) / max(dim, 1)
    reg = (ridge * scale + eps) * np.eye(dim)
    inv = np.linalg.pinv(xtx + reg, rcond=eps)

    jac = (packet_centered.T @ state_centered) @ inv
    return jac


def metric_from_jacobian(jac: np.ndarray, eps: float) -> np.ndarray:
    """Construct an SPD metric tensor from the Jacobian."""
    g = jac.T @ jac
    # Symmetrise to mitigate numeric drift, then stabilise.
    g = 0.5 * (g + g.T)
    dim = g.shape[0]
    g += eps * np.eye(dim)
    return g


def metric_distance(metric: np.ndarray, diff: np.ndarray) -> float:
    return float(np.sqrt(np.maximum(diff @ metric @ diff, 0.0)))


def curvature_along_path(
    g_prev: np.ndarray,
    g_curr: np.ndarray,
    g_next: np.ndarray,
    state_prev: np.ndarray,
    state_curr: np.ndarray,
    state_next: np.ndarray,
    eps: float,
) -> Tuple[float, float]:
    """Directional curvature proxy using central differences along the world-line."""
    dim = g_curr.shape[0]
    g_inv = np.linalg.pinv(g_curr, rcond=eps)

    delta_prev = state_curr - state_prev
    delta_next = state_next - state_curr

    mid_prev = 0.5 * (g_prev + g_curr)
    mid_next = 0.5 * (g_curr + g_next)

    ds_prev = np.sqrt(max(delta_prev @ mid_prev @ delta_prev, eps))
    ds_next = np.sqrt(max(delta_next @ mid_next @ delta_next, eps))
    ds = max(0.5 * (ds_prev + ds_next), eps)

    dg = 0.5 * (g_next - g_prev)
    deformation = g_inv @ dg
    curvature_mag = float(np.linalg.norm(deformation, ord="fro") ** 2 / (ds ** 2))

    return curvature_mag, ds


def stress_scalar(
    g_inv: np.ndarray,
    packet_prev: np.ndarray,
    packet_curr: np.ndarray,
    packet_next: np.ndarray,
) -> float:
    """Metric-weighted stress scalar from packet flux."""
    flux = 0.5 * (packet_next - packet_prev)
    return float(flux @ g_inv @ flux)


def compute_frame_diagnostics(
    series: WormholeSeries,
    window_radius: int,
    eps: float,
    ridge: float,
    kappa: float,
) -> Tuple[List[FrameDiagnostics], Dict[str, np.ndarray]]:
    frames, dim = series.frames, series.dim
    metrics: List[np.ndarray] = []

    # Pre-compute all metrics.
    for idx in range(frames):
        win_idx = window_indices(idx, frames, window_radius)
        if win_idx.size < 2:
            jac = np.zeros((dim, dim), dtype=float)
        else:
            jac = local_jacobian(series.state_b[win_idx], series.packet[win_idx], ridge=ridge, eps=eps)
        g = metric_from_jacobian(jac, eps)
        metrics.append(g)

    eigenvalues = np.empty((frames, dim), dtype=float)
    throat_metric = np.empty(frames, dtype=float)
    throat_euclid = np.empty(frames, dtype=float)
    curvature = np.full(frames, np.nan, dtype=float)
    stress = np.full(frames, np.nan, dtype=float)
    residual = np.full(frames, np.nan, dtype=float)
    divergence = np.full(frames, np.nan, dtype=float)

    # Throat distances & eigen spectra.
    for idx, g in enumerate(metrics):
        eigvals = np.linalg.eigvalsh(g)
        eigenvalues[idx] = eigvals
        delta = series.state_b[idx] - series.state_a[idx]
        throat_metric[idx] = metric_distance(g, delta)
        throat_euclid[idx] = float(np.linalg.norm(delta))

    # Directional curvature and stress scalars (skip boundaries).
    for idx in range(1, frames - 1):
        g_prev, g_curr, g_next = metrics[idx - 1], metrics[idx], metrics[idx + 1]
        curv, ds = curvature_along_path(
            g_prev,
            g_curr,
              g_next,
            series.state_b[idx - 1],
            series.state_b[idx],
            series.state_b[idx + 1],
            eps,
        )
        curvature[idx] = curv

        g_inv = np.linalg.pinv(g_curr, rcond=eps)
        stress_val = stress_scalar(g_inv, series.packet[idx - 1], series.packet[idx], series.packet[idx + 1])
        stress[idx] = stress_val
        residual[idx] = abs(curv - kappa * stress_val)

    # Divergence proxy (central finite difference of stress).
    for idx in range(2, frames - 2):
        prev = stress[idx - 1]
        next_ = stress[idx + 1]
        if np.isfinite(prev) and np.isfinite(next_):
            divergence[idx] = abs(next_ - prev)

    diagnostics = [
        FrameDiagnostics(
            eigenvalues=eigenvalues[idx],
            throat_metric=throat_metric[idx],
            throat_euclid=throat_euclid[idx],
            curvature_dir=curvature[idx],
            stress_scalar=stress[idx],
            einstein_residual=residual[idx],
            divergence=divergence[idx],
        )
        for idx in range(frames)
    ]

    arrays = {
        "eigenvalues": eigenvalues,
        "throat_metric": throat_metric,
        "throat_euclid": throat_euclid,
        "curvature": curvature,
        "stress": stress,
        "residual": residual,
        "divergence": divergence,
    }
    return diagnostics, arrays


def percentile(values: np.ndarray, q: Iterable[float]) -> List[float]:
    finite = values[np.isfinite(values)]
    if not finite.size:
        return [float("nan") for _ in q]
    return [float(np.percentile(finite, qq)) for qq in q]


def detect_change_points(values: np.ndarray, z_threshold: float = 2.5, min_gap: int = 5) -> List[int]:
    finite = values[np.isfinite(values)]
    if finite.size < 5:
        return []
    mean = float(np.nanmean(values))
    std = float(np.nanstd(values))
    if std < 1e-12:
        return []
    zscores = (values - mean) / std
    change_frames: List[int] = []
    last = -min_gap
    for idx, z in enumerate(zscores):
        if not np.isfinite(z):
            continue
        if z > z_threshold and idx - last >= min_gap:
            change_frames.append(idx)
            last = idx
    return change_frames


def assemble_summary(
    series: WormholeSeries,
    arrays: Dict[str, np.ndarray],
    kappa: float,
) -> Dict[str, object]:
    eigenvalues = arrays["eigenvalues"]
    throat_metric = arrays["throat_metric"]
    throat_euclid = arrays["throat_euclid"]
    curvature = arrays["curvature"]
    stress = arrays["stress"]
    residual = arrays["residual"]
    divergence = arrays["divergence"]

    eig_min = float(np.min(eigenvalues))
    eig_max = float(np.max(eigenvalues))
    eig_median = float(np.median(eigenvalues))

    throat_idx = int(np.nanargmin(throat_metric))
    throat_summary = {
        "euclidean_min": float(np.nanmin(throat_euclid)),
        "metric_min": float(np.nanmin(throat_metric)),
        "metric_min_scaled": float(np.nanmin(throat_metric) * kappa),
        "metric_argmin_frame": throat_idx,
    }

    summary = {
        "frames": series.frames,
        "dim": series.dim,
        "kappa": kappa,
        "metric": {
            "eigenvalue_min": eig_min,
            "eigenvalue_median": eig_median,
            "eigenvalue_max": eig_max,
            "condition_number": eig_max / max(eig_min, DEFAULT_EPSILON),
            "metric_minima_frames": throat_summary,
        },
        "curvature": {
            "p05": percentile(curvature, [5])[0],
            "median": percentile(curvature, [50])[0],
            "p95": percentile(curvature, [95])[0],
            "change_points": detect_change_points(curvature),
        },
        "stress": {
            "median": percentile(stress, [50])[0],
            "p95": percentile(stress, [95])[0],
        },
        "einstein_residual": {
            "median": percentile(residual, [50])[0],
            "p95": percentile(residual, [95])[0],
            "change_points": detect_change_points(residual),
        },
        "divergence_proxy": {
            "median": percentile(divergence, [50])[0],
            "max": float(np.nanmax(divergence)) if np.isfinite(divergence).any() else float("nan"),
        },
        "dashboard_calibration": {
            "throat_metric_mean": float(np.nanmean(throat_metric)),
            "throat_metric_std": float(np.nanstd(throat_metric)),
            "throat_metric_scaled_mean": float(np.nanmean(throat_metric) * kappa),
        },
    }
    return summary


def print_summary(summary: Dict[str, object]) -> None:
    metric = summary["metric"]
    throat = metric["metric_minima_frames"]
    curvature = summary["curvature"]
    stress = summary["stress"]
    residual = summary["einstein_residual"]
    divergence = summary["divergence_proxy"]
    dashboard = summary["dashboard_calibration"]

    print("Frames:", summary["frames"], "| Dimension:", summary["dim"])
    print(f"Metric eigenvalues: min={metric['eigenvalue_min']:.3e}, median={metric['eigenvalue_median']:.3e}, max={metric['eigenvalue_max']:.3e}")
    print(f"Metric condition number: {metric['condition_number']:.3e}")
    print(
        "Throat minima: frame={metric_min_frame}, Euclid={euclid:.4f}, Metric={metric_min:.4e}, Metric×κ={scaled:.4f}".format(
            metric_min_frame=throat["metric_argmin_frame"],
            euclid=throat["euclidean_min"],
            metric_min=throat["metric_min"],
            scaled=throat["metric_min_scaled"],
        )
    )

    print(
        "Curvature dir stats: p05={p05:.3e}, median={med:.3e}, p95={p95:.3e}".format(
            p05=curvature["p05"], med=curvature["median"], p95=curvature["p95"]
        )
    )
    if curvature["change_points"]:
        print("  Curvature surges at frames:", ", ".join(map(str, curvature["change_points"])))

    print(
        "Stress scalar stats: median={med:.3e}, p95={p95:.3e}".format(
            med=stress["median"], p95=stress["p95"]
        )
    )

    print(
        "Einstein residual stats: median={med:.3e}, p95={p95:.3e}".format(
            med=residual["median"], p95=residual["p95"]
        )
    )
    if residual["change_points"]:
        print("  Residual change-points:", ", ".join(map(str, residual["change_points"])))

    print(
        "Divergence proxy: median={med:.3e}, max={mx:.3e}".format(
            med=divergence["median"], mx=divergence["max"]
        )
    )

    print(
        "Dashboard throat calibration: mean={mean:.4e}, std={std:.4e}, scaled mean={scaled:.4f}".format(
            mean=dashboard["throat_metric_mean"],
            std=dashboard["throat_metric_std"],
            scaled=dashboard["throat_metric_scaled_mean"],
        )
    )


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="State-dependent GR analogue analysis for wormhole feeds")
    parser.add_argument("--feed", type=Path, required=True, help="Path to wireframe JSON feed")
    parser.add_argument("--window", type=int, default=DEFAULT_WINDOW, help="Radius (in frames) for local Jacobian fits")
    parser.add_argument("--epsilon", type=float, default=DEFAULT_EPSILON, help="Diagonal stabiliser for metric inverses")
    parser.add_argument("--ridge", type=float, default=DEFAULT_RIDGE, help="Relative ridge strength for Jacobian fits")
    parser.add_argument("--kappa", type=float, default=1.0, help="Scale factor that maps metric units to dashboard units")
    parser.add_argument("--out", type=Path, default=None, help="Optional JSON file to dump diagnostic summary")
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    series = load_wormhole_series(args.feed)

    diagnostics, arrays = compute_frame_diagnostics(
        series,
        window_radius=args.window,
        eps=args.epsilon,
        ridge=args.ridge,
        kappa=args.kappa,
    )

    summary = assemble_summary(series, arrays, kappa=args.kappa)
    print_summary(summary)

    if args.out is not None:
        args.out.parent.mkdir(parents=True, exist_ok=True)
        args.out.write_text(json.dumps(summary, indent=2), encoding="utf-8")


if __name__ == "__main__":
    main()
