from __future__ import annotations

import json
from pathlib import Path
from typing import Dict, List

import matplotlib.pyplot as plt

ROOT = Path(__file__).resolve().parents[1]
ARTIFACTS = ROOT / "artifacts" / "relativistic_suite"
DATA_PATH = ARTIFACTS / "relativistic_suite.json"


def _load_data() -> Dict:
    if not DATA_PATH.exists():
        raise FileNotFoundError(f"Missing experiment bundle: {DATA_PATH}")
    with DATA_PATH.open("r", encoding="utf-8") as f:
        return json.load(f)


def plot_lensing(experiments: Dict[str, List[Dict]]) -> Path:
    lensing = experiments.get("lensing", [])
    if not lensing:
        raise ValueError("No lensing data in experiment bundle")

    strengths = sorted({entry["strength"] for entry in lensing})
    plt.figure(figsize=(8, 5))
    for strength in strengths:
        subset = [entry for entry in lensing if entry["strength"] == strength]
        subset.sort(key=lambda e: e["offset"])
        offsets = [entry["offset"] for entry in subset]
        angles = [entry["deflection_angle"] for entry in subset]
        plt.plot(offsets, angles, marker="o", label=f"λh={strength:.0f}")

    plt.title("Information Lensing: Deflection vs Offset")
    plt.xlabel("Offset (impact parameter proxy)")
    plt.ylabel("Deflection angle (rad)")
    plt.legend()
    plt.grid(alpha=0.3)
    output = ARTIFACTS / "lensing_deflection.png"
    plt.tight_layout()
    plt.savefig(output, dpi=160)
    plt.close()
    return output


def plot_frame_dragging(experiments: Dict[str, List[Dict]]) -> Path:
    frame_drag = experiments.get("frame_dragging", [])
    if not frame_drag:
        raise ValueError("No frame dragging data in experiment bundle")

    angular_momenta = sorted({entry["angular_momentum"] for entry in frame_drag})
    plt.figure(figsize=(8, 5))
    for ang in angular_momenta:
        subset = [entry for entry in frame_drag if entry["angular_momentum"] == ang]
        subset.sort(key=lambda e: e["radius"])
        radii = [entry["radius"] for entry in subset]
        rates = [entry["mean_rotation_rate"] for entry in subset]
        plt.plot(radii, rates, marker="s", label=f"ω={ang:.2f}")

    plt.title("Frame Dragging: Mean Rotation vs Radius")
    plt.xlabel("Radius (arbitrary units)")
    plt.ylabel("Mean rotation rate")
    plt.legend()
    plt.grid(alpha=0.3)
    output = ARTIFACTS / "frame_dragging.png"
    plt.tight_layout()
    plt.savefig(output, dpi=160)
    plt.close()
    return output


def plot_tunneling(experiments: Dict[str, List[Dict]]) -> Path:
    tunneling = experiments.get("quantum_tunneling", [])
    if not tunneling:
        raise ValueError("No tunneling data in experiment bundle")

    tunneling.sort(key=lambda e: e["barrier_height"])
    barriers = [entry["barrier_height"] for entry in tunneling]
    throughput = [entry["throughput"] for entry in tunneling]

    plt.figure(figsize=(8, 5))
    plt.plot(barriers, throughput, marker="^", color="#ff7f0e")
    plt.title("Quantum Tunneling: Throughput vs Barrier Height")
    plt.xlabel("Barrier height (verify tolerance)")
    plt.ylabel("Throughput")
    plt.grid(alpha=0.3)
    output = ARTIFACTS / "tunneling_throughput.png"
    plt.tight_layout()
    plt.savefig(output, dpi=160)
    plt.close()
    return output


def main() -> None:
    data = _load_data()
    experiments = data.get("experiments", {})
    ARTIFACTS.mkdir(parents=True, exist_ok=True)

    outputs = [
        plot_lensing(experiments),
        plot_frame_dragging(experiments),
        plot_tunneling(experiments),
    ]

    for path in outputs:
        print(f"Wrote plot: {path}")


if __name__ == "__main__":
    main()
