#!/usr/bin/env python3
"""Audit blank predictions and response-length splits in local LLM benchmark JSON.

Compatible with T20-style result files:

  {
    "model_id": "...",
    "benchmark": "...",
    "accuracy": 0.0867,
    "correct": 26,
    "total": 300,
    "questions": [
      {
        "id": "...",
        "correct": false,
        "expected": "A",
        "predicted": "",
        "raw_response": "...",
        ...
      }
    ]
  }

Usage:
  python3 blank_prediction_audit.py /path/to/result.json
  python3 blank_prediction_audit.py /path/to/dir/*.json
  python3 blank_prediction_audit.py result.json --json
"""

from __future__ import annotations

import argparse
import json
import math
import statistics
import sys
from pathlib import Path
from typing import Any


def _is_blank(pred: Any) -> bool:
    if pred is None:
        return True
    if isinstance(pred, str) and not pred.strip():
        return True
    return False


def _raw_len(item: dict) -> int:
    raw = item.get("raw_response")
    if raw is None:
        raw = item.get("response") or item.get("output") or ""
    return len(str(raw))


def _is_single_letter(item: dict) -> bool:
    raw = str(item.get("raw_response") or item.get("response") or "").strip()
    pred = str(item.get("predicted") or "").strip()
    if len(raw) == 1 and raw.upper() in "ABCDE":
        return True
    if len(pred) == 1 and pred.upper() in "ABCDE" and len(raw) <= 3:
        return True
    return False


def wilson_halfwidth(p: float, n: int, z: float = 1.96) -> float:
    """Approx half-width of Wilson / normal interval for a proportion."""
    if n <= 0:
        return float("nan")
    # normal approx is enough for a self-check checklist
    return z * math.sqrt(max(p * (1.0 - p), 1e-12) / n)


def audit_file(path: Path) -> dict[str, Any]:
    data = json.loads(path.read_text(encoding="utf-8"))
    questions = data.get("questions") or data.get("items") or data.get("results")
    if not isinstance(questions, list) or not questions:
        raise ValueError(f"{path}: no questions/items list found")

    n = len(questions)
    blank_idx = [i for i, q in enumerate(questions) if _is_blank(q.get("predicted"))]
    lengths = [_raw_len(q) for q in questions]
    single = sum(1 for q in questions if _is_single_letter(q))
    correct = sum(1 for q in questions if q.get("correct") is True)
    # accuracy from file if present, else recompute
    acc = data.get("accuracy")
    if acc is None:
        acc = correct / n if n else 0.0

    blank_n = len(blank_idx)
    blank_rate = blank_n / n if n else 0.0
    nonblank_lens = [_raw_len(q) for q in questions if not _is_blank(q.get("predicted"))]
    blank_lens = [_raw_len(questions[i]) for i in blank_idx]

    def med(xs: list[int]) -> float | None:
        return float(statistics.median(xs)) if xs else None

    report = {
        "file": str(path),
        "model_id": data.get("model_id"),
        "benchmark": data.get("benchmark"),
        "n": n,
        "accuracy": acc,
        "correct": correct,
        "blank_predicted": blank_n,
        "blank_rate": blank_rate,
        "single_letter_like": single,
        "single_letter_rate": single / n if n else 0.0,
        "raw_len_median": med(lengths),
        "raw_len_p90": float(sorted(lengths)[max(0, int(math.ceil(0.9 * n) - 1))]) if lengths else None,
        "raw_len_max": max(lengths) if lengths else None,
        "raw_len_median_blank": med(blank_lens),
        "raw_len_median_nonblank": med(nonblank_lens),
        "ci95_halfwidth_at_p": wilson_halfwidth(float(acc), n),
        "ci95_halfwidth_at_0.5": wilson_halfwidth(0.5, n),
        "veto": [],
    }

    if blank_rate > 0.10:
        report["veto"].append(f"blank_rate={blank_rate:.1%} > 10%: do not rank on this benchmark mean")
    if float(acc) < 0.20 and blank_rate > 0.20:
        report["veto"].append(
            f"accuracy={float(acc):.1%} below 5-way chance (~20%) with high blank rate: treat as broken meter first"
        )
    if n < 100:
        report["veto"].append(f"n={n} < 100: ranking is a working hypothesis only")

    return report


def format_text(r: dict[str, Any]) -> str:
    lines = [
        f"file:        {r['file']}",
        f"model:       {r.get('model_id')}",
        f"benchmark:   {r.get('benchmark')}",
        f"n:           {r['n']}",
        f"accuracy:    {float(r['accuracy']):.4f}  ({r['correct']}/{r['n']})",
        f"blank:       {r['blank_predicted']}  ({r['blank_rate']:.1%})",
        f"single-ish:  {r['single_letter_like']}  ({r['single_letter_rate']:.1%})",
        f"raw median:  {r['raw_len_median']}",
        f"raw p90/max: {r['raw_len_p90']} / {r['raw_len_max']}",
        f"raw med blank/nonblank: {r['raw_len_median_blank']} / {r['raw_len_median_nonblank']}",
        f"~95% half-width @acc:  ±{r['ci95_halfwidth_at_p']*100:.1f} pp",
        f"~95% half-width @0.5:  ±{r['ci95_halfwidth_at_0.5']*100:.1f} pp",
    ]
    if r["veto"]:
        lines.append("veto flags:")
        for v in r["veto"]:
            lines.append(f"  - {v}")
    else:
        lines.append("veto flags: (none)")
    return "\n".join(lines)


def main(argv: list[str] | None = None) -> int:
    ap = argparse.ArgumentParser(description="Audit blank predictions in local LLM benchmark JSON")
    ap.add_argument("paths", nargs="+", help="JSON file(s) or globs")
    ap.add_argument("--json", action="store_true", help="print machine-readable JSON")
    args = ap.parse_args(argv)

    files: list[Path] = []
    for p in args.paths:
        path = Path(p)
        if any(ch in p for ch in "*?[]"):
            files.extend(sorted(Path().glob(p)))
        elif path.is_dir():
            files.extend(sorted(path.glob("*.json")))
        else:
            files.append(path)

    if not files:
        print("no files matched", file=sys.stderr)
        return 2

    reports = []
    for f in files:
        try:
            reports.append(audit_file(f))
        except Exception as e:
            print(f"ERROR {f}: {e}", file=sys.stderr)
            return 1

    if args.json:
        print(json.dumps(reports if len(reports) > 1 else reports[0], ensure_ascii=False, indent=2))
    else:
        for i, r in enumerate(reports):
            if i:
                print("-" * 60)
            print(format_text(r))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
