"""Render the entire saved Flask shortlist curve from the adjacent verified CSV.

Requires Matplotlib only for rendering. Run: python plot.py
Optional regeneration/verification from the frozen research result:
  python plot.py --source-results /path/to/cutoff_curve/results.json
No inference. The default path has no dependency on the research repository.
"""

import argparse
import csv
import hashlib
import json
from pathlib import Path

HERE = Path(__file__).resolve().parent
SOURCE_SHA256 = "def5bd7cb7b506faf5769cb2d310adc709ffc55e6fbc4001edcd28aeb57ecd7a"
FIELDS = [
    "cutoff",
    "gold_size_ceiling",
    "keep_score",
    "binary_keep_then_bm25",
    "bm25",
    "answerable_denominator",
    "unavailable_excluded",
]


def from_results(path):
    raw = path.read_bytes()
    if hashlib.sha256(raw).hexdigest() != SOURCE_SHA256:
        raise ValueError("Frozen source result hash mismatch")
    data = json.loads(raw)
    if data["answerable"] != 32 or data["unavailable"] != 8 or data["failures"]:
        raise ValueError("Unexpected result cohort")
    return [
        {
            "cutoff": r["cutoff"],
            "gold_size_ceiling": r["gold_size_ceiling"],
            **r["complete"],
            "answerable_denominator": 32,
            "unavailable_excluded": 8,
        }
        for r in data["curve"]
    ]


def read_csv(path):
    with path.open(newline="", encoding="utf-8") as handle:
        reader = csv.DictReader(handle)
        if reader.fieldnames != FIELDS:
            raise ValueError("Unexpected CSV schema")
        rows = [{key: int(value) for key, value in row.items()} for row in reader]
    if [r["cutoff"] for r in rows] != list(range(1, 17)):
        raise ValueError("The entire integer grid 1..16 is required")
    for row in rows:
        if row["answerable_denominator"] != 32 or row["unavailable_excluded"] != 8:
            raise ValueError("Unexpected denominator")
        if not all(
            0 <= row[key] <= row["gold_size_ceiling"] <= 32 for key in FIELDS[2:5]
        ):
            raise ValueError("Invalid count or ceiling")
    return rows


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--source-results", type=Path)
    args = parser.parse_args()
    csv_path = HERE / "shortlist-cutoff.csv"
    if args.source_results:
        expected = from_results(args.source_results)
        if csv_path.exists():
            if read_csv(csv_path) != expected:
                raise ValueError("CSV differs from frozen results")
        else:
            with csv_path.open("w", newline="", encoding="utf-8") as handle:
                writer = csv.DictWriter(handle, fieldnames=FIELDS)
                writer.writeheader()
                writer.writerows(expected)
    rows = read_csv(csv_path)
    import matplotlib

    matplotlib.use("Agg")
    import matplotlib.pyplot as plt
    from matplotlib.ticker import MaxNLocator

    plt.rcParams.update(
        {
            "font.family": "DejaVu Sans",
            "font.size": 11,
            "svg.fonttype": "none",
            "svg.hashsalt": "lecter-shortlist-cutoff-v1",
            "axes.spines.top": False,
            "axes.spines.right": False,
        }
    )
    fig, ax = plt.subplots(figsize=(11, 6.6))
    fig.subplots_adjust(left=0.10, right=0.97, bottom=0.31, top=0.79)
    x = [r["cutoff"] for r in rows]
    ax.plot(
        x,
        [r["gold_size_ceiling"] for r in rows],
        color="#666666",
        linestyle="--",
        linewidth=2,
        label="Gold-size / candidate ceiling",
        zorder=1,
    )
    curves = [
        ("keep_score", "Numerical keep score", "#2166AC", "o"),
        ("binary_keep_then_bm25", "Binary keep-first + BM25", "#B35C00", "s"),
        ("bm25", "BM25", "#22836D", "^"),
    ]
    for key, label, color, marker in curves:
        ax.plot(
            x,
            [r[key] for r in rows],
            label=label,
            color=color,
            marker=marker,
            markersize=5,
            markerfacecolor="white",
            markeredgewidth=1.5,
            linewidth=2,
            zorder=3,
        )
    ax.set_xlim(0.6, 16.4)
    ax.set_ylim(0, 34)
    ax.set_xticks(range(1, 17))
    ax.set_yticks(range(0, 33, 4))
    ax.yaxis.set_major_locator(MaxNLocator(nbins=9, integer=True))
    ax.set_xlabel("Shortlist cutoff (spans)", labelpad=9)
    ax.set_ylabel("Complete required-evidence sets (of 32)", labelpad=10)
    ax.grid(axis="y", color="#E5E7EB", linewidth=0.8)
    ax.set_axisbelow(True)
    fig.text(
        0.10,
        0.93,
        "Required-evidence coverage across every cutoff",
        fontsize=18,
        weight="bold",
        ha="left",
    )
    fig.text(
        0.10,
        0.875,
        "Retrospective saved Flask cohort · shortlist only · same 64 candidates",
        fontsize=11.5,
        color="#444444",
        ha="left",
    )
    handles, labels = ax.get_legend_handles_labels()
    order = [1, 2, 3, 0]
    fig.legend(
        [handles[i] for i in order],
        [labels[i] for i in order],
        loc="lower center",
        bbox_to_anchor=(0.54, 0.115),
        ncol=2,
        frameon=False,
        columnspacing=2.8,
        handlelength=2.8,
    )
    fig.text(
        0.10,
        0.065,
        "32 answerable cases; 8 unavailable cases excluded from coverage. Ceiling uses annotated set size within the candidate pool.",
        fontsize=9,
        color="#555555",
    )
    fig.text(
        0.10,
        0.035,
        "Exploratory full grid after the cutoff-16 tie. No optimal cutoff selected; no new inference or downstream outcome claim.",
        fontsize=9,
        color="#555555",
    )
    for suffix in ("svg", "png"):
        metadata = (
            {"Date": None}
            if suffix == "svg"
            else {"Software": "Matplotlib; offline frozen-result plot"}
        )
        fig.savefig(HERE / ("shortlist-cutoff." + suffix), dpi=180, metadata=metadata)
    plt.close(fig)
    print("Rendered all 16 cutoffs from verified CSV; no inference.")


if __name__ == "__main__":
    main()
