"""Single-pass A-vs-B comparison over BHE_991, with per-page overlays.

Renders GT (green), data-grid approach A (red), anatomy approach B (blue)
on every page and prints per-page accuracy metrics. Use this for the fast
iterate loop before scaling to repeatability runs and more documents.
"""

from __future__ import annotations

import argparse
import asyncio
from pathlib import Path

import fitz
from PIL import Image

from experiments.que245.approaches import make_locator
from experiments.que245.gt import BHE_991_BANDS, ground_truth
from experiments.que245.overlay import draw_boxes, render_page
from experiments.que245.score import page_metrics

PDF = Path("documents/BHE_991.pdf")
RENDER = Path("experiments/que245/render")


async def predict(approach: str, model: str, temperature, pages, rows=24, cols=12):
    loc = make_locator(approach, model=model, temperature=temperature, rows=rows, cols=cols)
    doc = fitz.open(str(PDF))
    out = {}
    import tempfile

    for p in pages:
        with tempfile.TemporaryDirectory() as tmp:
            img = Path(tmp) / f"p{p}.png"
            doc[p - 1].get_pixmap(dpi=200).save(str(img))
            located = await loc.locate(img, PDF, p)
        out[p] = [t.region for t in located]
    return out


async def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", default="claude-haiku-4-5-20251001")
    ap.add_argument("--temp", default="0.0")
    ap.add_argument("--pages", default="1-8")
    ap.add_argument("--rows", type=int, default=24)
    ap.add_argument("--cols", type=int, default=12)
    args = ap.parse_args()
    temp = None if args.temp == "none" else float(args.temp)
    a, b = args.pages.split("-") if "-" in args.pages else (args.pages, args.pages)
    pages = list(range(int(a), int(b) + 1))

    gt = ground_truth(PDF, BHE_991_BANDS)
    print(f"model={args.model} temp={temp}")
    pa = await predict("data-grid", args.model, temp, pages, args.rows, args.cols)
    pb = await predict("anatomy", args.model, temp, pages, args.rows, args.cols)

    print(f"\n{'pg':>2} | {'GT':>2} | A:n iou clip | B:n iou clip")
    print("-" * 56)
    sheet_imgs = []
    for p in pages:
        ma = page_metrics(gt[p], pa[p])
        mb = page_metrics(gt[p], pb[p])
        print(
            f"{p:>2} | {len(gt[p]):>2} | "
            f"{ma['pred']} {ma['mean_iou']:.2f} {ma['max_top_clip_pts']:+5.0f} | "
            f"{mb['pred']} {mb['mean_iou']:.2f} {mb['max_top_clip_pts']:+5.0f}"
        )
        im = render_page(PDF, p, dpi=120)
        im = draw_boxes(
            im,
            [("GT", gt[p], (0, 160, 0)), ("A", pa[p], (220, 0, 0)), ("B", pb[p], (0, 80, 230))],
        )
        sheet_imgs.append(im)

    w = max(i.width for i in sheet_imgs)
    h = max(i.height for i in sheet_imgs)
    cols, rows = 4, (len(sheet_imgs) + 3) // 4
    sheet = Image.new("RGB", (w * cols, h * rows), (255, 255, 255))
    for i, im in enumerate(sheet_imgs):
        sheet.paste(im, ((i % cols) * w, (i // cols) * h))
    tag = args.model.split("-")[1]
    out = RENDER / f"compare_{tag}_t{args.temp}_r{args.rows}.png"
    sheet.save(out)
    print(f"\nsaved {out}")


if __name__ == "__main__":
    asyncio.run(main())
