"""Generate packs for every holdout brief, then run the checks that don't need a judge.

The holdout was carved in build/score.py before the teacher ever ran, so these briefs
and their real posts have never been seen by any stage of the pipeline.

Three things get measured here:

  format validity  — does the pack obey every platform's published limits
  near-duplicates  — is the model reciting training posts back rather than writing
  baselines        — the same briefs through an un-tuned model, with and without
                     few-shot examples retrieved from the corpus

That last one is the honest question. If few-shot prompting with retrieval matches the
fine-tune, the fine-tune isn't earning its keep and the retrieval pipeline is the
thing worth shipping.

  python eval/run.py --model socialmedia --tag tuned
  python eval/run.py --model llama3.2:1b --tag base
  python eval/run.py --model llama3.2:1b --fewshot 3 --tag fewshot
"""
import argparse
import json
import pathlib
import sys
import time

sys.path.insert(0, str(pathlib.Path(__file__).resolve().parent.parent))
sys.path.insert(0, str(pathlib.Path(__file__).resolve().parent.parent / "serve"))
from common import (DATA, NICHES, PACK_SYSTEM, ROOT, read_jsonl, render_brief,  # noqa: E402
                    validate_pack, words, write_jsonl)
from generate import ask  # noqa: E402

OUT = ROOT / "out"
NEAR_DUP = 0.70


def jaccard(a, b):
    return len(a & b) / len(a | b) if a and b else 0.0


def fewshot_messages(brief, niche_label, train, n):
    """Retrieve the closest training examples and show them as prior turns. This is
    the baseline the fine-tune has to beat — it uses the same corpus, no training."""
    q = words(brief.get("topic", "") + " " + brief.get("angle", ""))
    scored = []
    for ex in train:
        user = ex["messages"][1]["content"]
        scored.append((jaccard(q, words(user)), ex))
    scored.sort(key=lambda t: t[0], reverse=True)
    msgs = [{"role": "system", "content": PACK_SYSTEM}]
    for _, ex in scored[:n]:
        msgs.append({"role": "user", "content": ex["messages"][1]["content"]})
        msgs.append({"role": "assistant", "content": ex["messages"][2]["content"]})
    msgs.append({"role": "user", "content": render_brief(brief, niche_label)})
    return msgs


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", default="socialmedia")
    ap.add_argument("--tag", required=True, help="label for the output file")
    ap.add_argument("--fewshot", type=int, default=0, help="retrieved examples to prepend")
    ap.add_argument("--limit", type=int, help="cap the number of briefs")
    ap.add_argument("--temperature", type=float, default=0.8)
    args = ap.parse_args()

    briefs = read_jsonl(DATA / "briefs_holdout.jsonl")
    if not briefs:
        sys.exit("No data/briefs_holdout.jsonl — run:\n"
                 "  python build/backtranslate.py --split holdout")
    holdout = {r["id"]: r for r in read_jsonl(DATA / "holdout.jsonl")}
    train = read_jsonl(DATA / "train.jsonl") if args.fewshot else []
    if args.fewshot and not train:
        sys.exit("--fewshot needs data/train.jsonl")

    if args.limit:
        briefs = briefs[:args.limit]

    dest = OUT / f"gen-{args.tag}.jsonl"
    rows, t0 = [], time.time()
    for i, b in enumerate(briefs, 1):
        rec = holdout.get(b["id"])
        if not rec:
            continue
        niche_label = NICHES.get(rec["niche"], "")
        try:
            if args.fewshot:
                msgs = fewshot_messages(b["brief"], niche_label, train, args.fewshot)
            else:
                msgs = [{"role": "system", "content": PACK_SYSTEM},
                        {"role": "user", "content": render_brief(b["brief"], niche_label)}]
            pack = json.loads(ask(args.model, msgs, args.temperature))
        except Exception as e:
            print(f"  {b['id']}: {type(e).__name__} {str(e)[:100]}", file=sys.stderr)
            continue
        rows.append({"id": b["id"], "niche": rec["niche"], "platform": rec["platform"],
                     "brief": b["brief"], "pack": pack, "real": rec})
        if i % 10 == 0:
            rate = (time.time() - t0) / i
            print(f"  {i}/{len(briefs)}  ({rate:.1f}s each, "
                  f"~{rate*(len(briefs)-i)/60:.0f} min left)")

    write_jsonl(dest, rows)
    print(f"\n{len(rows)} packs -> {dest}  ({(time.time()-t0)/60:.1f} min total)")

    # --- format validity -------------------------------------------------------
    bad = [(r["id"], validate_pack(r["pack"])) for r in rows]
    bad = [(i, p) for i, p in bad if p]
    print(f"\nformat validity: {len(rows)-len(bad)}/{len(rows)} "
          f"({100*(len(rows)-len(bad))/max(len(rows),1):.1f}%)")
    counts = {}
    for _, probs in bad:
        for p in probs:
            k = p.split(" is ")[0].split(" exceed")[0]
            counts[k] = counts.get(k, 0) + 1
    for k, n in sorted(counts.items(), key=lambda kv: -kv[1])[:8]:
        print(f"  {n:>4}  {k}")

    # --- near-duplicate gate ---------------------------------------------------
    train_all = train or read_jsonl(DATA / "train.jsonl")
    train_hooks = []
    for ex in train_all:
        try:
            p = json.loads(ex["messages"][2]["content"])
            train_hooks.append((words(p.get("hook", "")), words(p.get("script", ""))))
        except (json.JSONDecodeError, KeyError, IndexError):
            continue
    dups = []
    for r in rows:
        gh, gs = words(r["pack"].get("hook", "")), words(r["pack"].get("script", ""))
        worst = max((max(jaccard(gh, th), jaccard(gs, ts)) for th, ts in train_hooks),
                    default=0.0)
        if worst >= NEAR_DUP:
            dups.append((r["id"], round(worst, 2)))
    print(f"\nnear-duplicates of training data (>={NEAR_DUP}): "
          f"{len(dups)}/{len(rows)}")
    for i, s in dups[:5]:
        print(f"  {s}  {i}")

    # --- shape -----------------------------------------------------------------
    if rows:
        def real_opener(rec):
            """Each platform names its opening line differently."""
            t = rec.get("text") or {}
            return (t.get("title") or (t.get("caption") or "").split("\n")[0] or "")

        hooks = [len(r["pack"].get("hook", "")) for r in rows]
        scripts = [len(r["pack"].get("script", "")) for r in rows]
        real = [len(real_opener(r["real"])) for r in rows]
        print(f"\nmean hook {sum(hooks)/len(hooks):.0f} chars "
              f"(real winners' opening lines {sum(real)/len(real):.0f}), "
              f"mean script {sum(scripts)/len(scripts):.0f} chars")


if __name__ == "__main__":
    main()
