import os
import pickle
import numpy as np
import matplotlib.pyplot as plt
from matplotlib.backends.backend_pdf import PdfPages

# ============================================================
# Config
# ============================================================
PKL_PATH = "/nfs/fanae/user/jprado/LLP_Showers/shower-studies/wire_profile_run1.pkl"
OUTDIR = "figures"
OUTPDF = os.path.join(OUTDIR, "wire_profiles_by_category.pdf")

N_EVENTS = 100          # number of events (histograms) per category
GRID_ROWS, GRID_COLS = 10, 10   # 10x10 = 100 subplots per page

CATEGORIES = ("TP", "FN", "FP", "TN")

COLORS = {
    "TP": "tab:green",
    "FN": "tab:orange",
    "FP": "tab:red",
    "TN": "tab:blue",
}

RNG_SEED = 42

os.makedirs(OUTDIR, exist_ok=True)

with open(PKL_PATH, "rb") as f:
    data = pickle.load(f)

wireProfile = data["wireProfile"]

rng = np.random.default_rng(RNG_SEED)

with PdfPages(OUTPDF) as pdf:

    for cat in CATEGORIES:

        events = wireProfile.get(cat, [])
        n_available = len(events)

        if n_available == 0:
            print(f"[{cat}] no events found, skipping.")
            continue

        n_to_plot = min(N_EVENTS, n_available)

        # random sample without replacement of up to N_EVENTS events
        idx = rng.choice(n_available, size=n_to_plot, replace=False)
        idx.sort()

        fig, axes = plt.subplots(
            GRID_ROWS, GRID_COLS,
            figsize=(GRID_COLS * 2, GRID_ROWS * 2),
        )
        axes = axes.flatten()

        for ax_i, ev_i in enumerate(idx):
            evt = np.asarray(events[ev_i])
            wires = np.arange(len(evt))

            ax = axes[ax_i]
            ax.bar(wires, evt, width=1.0, color=COLORS[cat])
            ax.set_title(f"evt {ev_i}", fontsize=6)
            ax.tick_params(labelsize=5)
            ax.set_xlim(0, len(evt))

        # hide any unused axes (if fewer events than grid size)
        for ax_j in range(n_to_plot, len(axes)):
            axes[ax_j].axis("off")

        fig.suptitle(
            f"{cat} wire profiles "
            f"({n_to_plot} of {n_available} events shown)",
            fontsize=16,
        )
        fig.supxlabel("Wire index", fontsize=10)
        fig.supylabel("Hits", fontsize=10)

        plt.tight_layout(rect=[0.02, 0.02, 1, 0.97])
        pdf.savefig(fig, dpi=150)
        plt.close(fig)

        print(f"[{cat}] plotted {n_to_plot}/{n_available} events.")

print(f"\nPDF saved to: {os.path.abspath(OUTPDF)}")