"""Blind pairwise judging against the real winners.

The bar is deliberately harsh: every comparison is against a post that actually
landed in the top ~18% of its creator's output. Beating that half the time would be
remarkable; 45% is a genuinely good result and 30% is still a usable model.

Comparisons are made on the real post's own platform, using only the fields that
platform actually published — generating a Pinterest description and scoring it
against a YouTube title would measure nothing. Sides are shuffled per item, by a
hash of the id so reruns are reproducible, and the judge never learns which side
is which.

  python eval/judge.py --a out/gen-tuned.jsonl                        # vs real winners
  python eval/judge.py --a out/gen-tuned.jsonl --b out/gen-fewshot.jsonl
"""
import argparse
import hashlib
import json
import pathlib
import sys
from collections import Counter, defaultdict
from concurrent.futures import ThreadPoolExecutor

from openai import OpenAI

sys.path.insert(0, str(pathlib.Path(__file__).resolve().parent.parent))
from common import DATA, TEACHER_MODEL, _obj, load_env, read_jsonl  # noqa: E402

VERDICT = _obj({"winner": {"type": "string", "enum": ["A", "B"]},
                "reason": {"type": "string"}}, ["winner", "reason"])

SYSTEM = """\
You are judging short-form social content. You will see a brief and two candidate
pieces written for the same platform. Pick the one that would perform better with the
brief's audience.

Weigh, in this order:
  1. Does the opening earn the next three seconds — is it specific, surprising, or
     does it name a real problem?
  2. Is the substance concrete? Real numbers, mechanisms and named things beat
     generalities and encouragement.
  3. Does it read as native to the platform rather than repurposed filler?

Ignore length, formatting polish, emoji and hashtag counts. Do not reward whichever
sounds more like an advertisement. If both are weak, still pick the less weak one.
"""


def side_from_pack(pack, platform):
    """Render a generated pack the way the given platform would publish it."""
    if platform == "youtube":
        yt = pack.get("youtube") or {}
        return (f"TITLE: {yt.get('title','')}\n"
                f"OPENING LINE: {pack.get('hook','')}\n"
                f"SCRIPT: {pack.get('script','')}")
    if platform == "instagram":
        ig = pack.get("instagram") or {}
        return (f"OPENING LINE: {pack.get('hook','')}\n"
                f"CAPTION: {ig.get('caption','')}")
    pin = pack.get("pinterest") or {}
    return f"TITLE: {pin.get('title','')}\nDESCRIPTION: {pin.get('description','')}"


def side_from_real(rec, transcripts):
    """The real post, using only what that platform actually published."""
    t = rec.get("text") or {}
    if rec["platform"] == "youtube":
        tr = transcripts.get(rec["id"], "")
        return (f"TITLE: {t.get('title','')}\n"
                + (f"SCRIPT: {tr[:2500]}" if tr
                   else f"DESCRIPTION: {(t.get('description') or '')[:800]}"))
    if rec["platform"] == "instagram":
        return f"CAPTION: {(t.get('caption') or '')[:1800]}"
    return f"TITLE: {t.get('title','')}\nDESCRIPTION: {(t.get('description') or '')[:800]}"


def brief_text(brief):
    return (f"TOPIC: {brief.get('topic','')}\n"
            f"AUDIENCE: {brief.get('audience','')}\n"
            f"ANGLE: {brief.get('angle','')}")


def judge_one(client, item):
    """Returns which label won, or None if the call failed."""
    left, right = (item["a"], item["b"]) if item["flip"] else (item["b"], item["a"])
    user = (f"{brief_text(item['brief'])}\n\nPLATFORM: {item['platform']}\n\n"
            f"--- CANDIDATE A ---\n{left}\n\n--- CANDIDATE B ---\n{right}")
    try:
        r = client.chat.completions.create(
            model=TEACHER_MODEL,
            messages=[{"role": "system", "content": SYSTEM},
                      {"role": "user", "content": user}],
            response_format={"type": "json_schema",
                             "json_schema": {"name": "verdict", "schema": VERDICT,
                                             "strict": True}},
            max_completion_tokens=300)
        d = json.loads(r.choices[0].message.content)
    except Exception as e:
        print(f"  {item['id']}: {type(e).__name__} {str(e)[:90]}", file=sys.stderr)
        return None
    chose_left = d["winner"] == "A"
    winner = item["a_label"] if (chose_left == item["flip"]) else item["b_label"]
    return {"id": item["id"], "niche": item["niche"], "platform": item["platform"],
            "winner": winner, "reason": d["reason"]}


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--a", required=True, help="generated packs jsonl")
    ap.add_argument("--b", help="second generated file; omit to judge against the real posts")
    ap.add_argument("--limit", type=int)
    ap.add_argument("--workers", type=int, default=6)
    ap.add_argument("--out", default="out/judge.jsonl")
    args = ap.parse_args()

    load_env()
    rows_a = {r["id"]: r for r in read_jsonl(args.a)}
    if not rows_a:
        sys.exit(f"No packs in {args.a} — run eval/run.py first.")
    a_label = pathlib.Path(args.a).stem.replace("gen-", "")

    transcripts = {r["id"]: r["text"] for r in read_jsonl(DATA / "transcripts.jsonl")
                   if r.get("text")}

    if args.b:
        rows_b = {r["id"]: r for r in read_jsonl(args.b)}
        b_label = pathlib.Path(args.b).stem.replace("gen-", "")
        ids = [i for i in rows_a if i in rows_b]
    else:
        rows_b, b_label = None, "real"
        ids = list(rows_a)
    if args.limit:
        ids = ids[:args.limit]
    print(f"judging {a_label} vs {b_label} on {len(ids)} items")

    items = []
    for i in ids:
        ra = rows_a[i]
        plat = ra["platform"]
        b_side = (side_from_pack(rows_b[i]["pack"], plat) if rows_b
                  else side_from_real(ra["real"], transcripts))
        items.append({
            "id": i, "niche": ra["niche"], "platform": plat, "brief": ra["brief"],
            "a": side_from_pack(ra["pack"], plat), "b": b_side,
            "a_label": a_label, "b_label": b_label,
            # deterministic coin flip so reruns judge the same layout
            "flip": int(hashlib.sha1(i.encode()).hexdigest(), 16) % 2 == 0,
        })

    client = OpenAI()
    with ThreadPoolExecutor(max_workers=args.workers) as pool:
        results = [r for r in pool.map(lambda it: judge_one(client, it), items) if r]

    if not results:
        sys.exit("every judging call failed")

    pathlib.Path(args.out).parent.mkdir(parents=True, exist_ok=True)
    with open(args.out, "w") as fh:
        for r in results:
            fh.write(json.dumps(r, ensure_ascii=False) + "\n")

    tally = Counter(r["winner"] for r in results)
    wins = tally[a_label]
    print(f"\n{a_label} beat {b_label} on {wins}/{len(results)} "
          f"({100*wins/len(results):.1f}%)")

    for dim in ("niche", "platform"):
        print(f"\nby {dim}:")
        buckets = defaultdict(lambda: [0, 0])
        for r in results:
            buckets[r[dim]][0] += r["winner"] == a_label
            buckets[r[dim]][1] += 1
        for k, (w, n) in sorted(buckets.items()):
            print(f"  {k:<14} {w:>4}/{n:<4} {100*w/n:>5.1f}%")

    print(f"\nsample reasons where {b_label} won:")
    for r in [x for x in results if x["winner"] != a_label][:3]:
        print(f"  - {r['reason'][:150]}")
    print(f"\nfull verdicts -> {args.out}")


if __name__ == "__main__":
    main()
