"""Visualization utilities for the standalone wormhole proof experiment."""

from __future__ import annotations

from pathlib import Path
from typing import List

import matplotlib.pyplot as plt
import numpy as np

from .core.sweep import SweepResult
from .core.wormhole import RegimeResult, SimulationResult

plt.style.use("seaborn-v0_8-darkgrid")


def _plot_wormhole_throughput(output_dir: Path, simulations: List[SimulationResult]) -> Path:
    seeds = []
    throughputs = []
    hp_match = []
    for sim in simulations:
        wormhole_regime = next(r for r in sim.regimes if r.name == "wormhole")
        seeds.append(sim.seed)
        throughputs.append(wormhole_regime.throughput)
        hp_match.append(wormhole_regime.metrics.hp_match_fraction)

    fig, ax1 = plt.subplots(figsize=(8, 4.5))
    indices = np.arange(len(seeds))
    bars = ax1.bar(indices - 0.15, throughputs, width=0.3, color="#4c78a8", label="throughput")
    ax1.set_xticks(indices)
    ax1.set_xticklabels([f"seed {s}" for s in seeds])
    ax1.set_ylabel("Throughput norm")
    ax1.set_title("Wormhole throughput and HP match fraction")

    ax2 = ax1.twinx()
    ax2.plot(indices + 0.15, hp_match, color="#f58518", marker="o", label="HP match")
    ax2.set_ylabel("HP match fraction")
    ax2.set_ylim(0, 1.05)

    handles, labels = [], []
    for ax in (ax1, ax2):
        h, l = ax.get_legend_handles_labels()
        handles.extend(h)
        labels.extend(l)
    ax1.legend(handles, labels, loc="upper left")

    path = output_dir / "figure_wormhole_throughput.png"
    fig.tight_layout()
    fig.savefig(path, dpi=200)
    plt.close(fig)
    return path


def _plot_sweep(output_dir: Path, sweep: SweepResult) -> Path:
    lam = np.array([p.lambda_h for p in sweep.points])
    throughput = np.array([p.throughput for p in sweep.points])
    soft_frac = np.array([p.soft_rank_fraction for p in sweep.points])

    fig, ax1 = plt.subplots(figsize=(8, 4.5))
    ax1.plot(lam, throughput, color="#4c78a8", marker="o", label="throughput")
    ax1.set_xlabel("lambda_h")
    ax1.set_ylabel("Throughput norm")
    ax1.set_title("Critical sweep: throughput vs soft-mode fraction")

    ax2 = ax1.twinx()
    ax2.plot(lam, soft_frac, color="#54a24b", marker="s", label="soft rank fraction")
    ax2.set_ylabel("Soft rank fraction")
    ax2.set_ylim(0, 1.05)

    handles, labels = [], []
    for ax in (ax1, ax2):
        h, l = ax.get_legend_handles_labels()
        handles.extend(h)
        labels.extend(l)
    ax1.legend(handles, labels, loc="upper right")

    if sweep.analysis.lambda_crit_est is not None:
        ax1.axvline(sweep.analysis.lambda_crit_est, color="#e45756", linestyle="--", label="lambda_crit")

    path = output_dir / "figure_sweep.png"
    fig.tight_layout()
    fig.savefig(path, dpi=200)
    plt.close(fig)
    return path


def _plot_hp_trace(output_dir: Path, simulations: List[SimulationResult]) -> Path | None:
    traces = []
    for sim in simulations:
        for regime in sim.regimes:
            if regime.trace is None:
                continue
            if regime.trace.hp_trace:
                traces.append((sim.seed, regime.name, regime.trace.hp_trace))

    if not traces:
        return None

    fig, ax = plt.subplots(figsize=(8, 4.5))
    for seed, regime_name, entries in traces:
        packets = [e.packet_index for e in entries]
        match = [e.match_fraction for e in entries]
        ax.plot(packets, match, marker="o", label=f"seed {seed} ({regime_name})")

    ax.set_xlabel("Packet index")
    ax.set_ylabel("Match fraction")
    ax.set_ylim(0, 1.05)
    ax.set_title("Hayden–Preskill decoding trace")
    ax.legend()

    path = output_dir / "figure_hp_trace.png"
    fig.tight_layout()
    fig.savefig(path, dpi=200)
    plt.close(fig)
    return path


def generate_visualizations(
    io: ExperimentIO,
    simulations: List[SimulationResult],
    sweep: SweepResult,
    benchmark_metrics: dict,
) -> List[Path]:
    paths: List[Path] = []
    output_dir = io.output_dir

    paths.append(_plot_wormhole_throughput(output_dir, simulations))
    paths.append(_plot_sweep(output_dir, sweep))
    hp_path = _plot_hp_trace(output_dir, simulations)
    if hp_path:
        paths.append(hp_path)

    # Bridge benchmark bar chart for loss gains
    fig, ax = plt.subplots(figsize=(6, 4))
    metrics = {
        "loss_gain_mean": benchmark_metrics.get("loss_gain_mean", 0.0),
        "loss_gain_std": benchmark_metrics.get("loss_gain_std", 0.0),
    }
    ax.bar(["mean"], [metrics["loss_gain_mean"]], color="#f58518", yerr=[metrics["loss_gain_std"]])
    ax.set_ylabel("Objective reduction")
    ax.set_title("CCA bridge objective gain")
    bridge_path = output_dir / "figure_bridge_gain.png"
    fig.tight_layout()
    fig.savefig(bridge_path, dpi=200)
    plt.close(fig)
    paths.append(bridge_path)

    return paths
