"""Fetch real transcripts for the YouTube winners.

Without this the `script` field in every training example would be invented by the
teacher model, and the distilled model would learn GPT's idea of a short-form script
instead of what actually performed. With it, the hook and script targets are the real
creator's words.

Runs only over curated winners plus holdout (a few thousand), not the whole raw
corpus, so it stays quick and keeps YouTube's patience.

  python collect/transcripts.py
"""
import argparse
import random
import sys
import threading
import time
from concurrent.futures import ThreadPoolExecutor

from youtube_transcript_api import YouTubeTranscriptApi

sys.path.insert(0, str(__import__("pathlib").Path(__file__).resolve().parent.parent))
from common import DATA, append_jsonl, read_jsonl, seen_ids  # noqa: E402

OUT = DATA / "transcripts.jsonl"
_lock = threading.Lock()


def fetch_one(api, rec, langs):
    vid = rec["id"].split(":", 1)[1]
    try:
        tr = api.fetch(vid, languages=langs)
        text = " ".join(s.text for s in tr).strip()
        text = " ".join(text.split())
        if len(text) < 40:
            return {"id": rec["id"], "text": "", "error": "too short"}
        return {"id": rec["id"], "text": text, "error": ""}
    except Exception as e:
        return {"id": rec["id"], "text": "", "error": type(e).__name__}


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--workers", type=int, default=4,
                    help="keep this low; YouTube throttles aggressive callers")
    ap.add_argument("--langs", default="en,en-US,en-GB")
    args = ap.parse_args()

    rows = read_jsonl(DATA / "curated.jsonl") + read_jsonl(DATA / "holdout.jsonl")
    if not rows:
        sys.exit("No curated.jsonl — run build/score.py first.")

    have = seen_ids(OUT)
    todo = [r for r in rows if r["platform"] == "youtube" and r["id"] not in have]
    print(f"{len(todo)} youtube winners need transcripts "
          f"({len(have)} already cached)")
    if not todo:
        return

    api = YouTubeTranscriptApi()
    langs = args.langs.split(",")
    done = {"n": 0, "ok": 0}
    buf = []

    def work(rec):
        time.sleep(random.uniform(0.1, 0.6))   # spread the load
        res = fetch_one(api, rec, langs)
        with _lock:
            buf.append(res)
            done["n"] += 1
            done["ok"] += bool(res["text"])
            if len(buf) >= 25:                  # checkpoint
                append_jsonl(OUT, buf)
                buf.clear()
            if done["n"] % 50 == 0:
                print(f"  {done['n']}/{len(todo)}  ({done['ok']} with text)")

    with ThreadPoolExecutor(max_workers=args.workers) as pool:
        list(pool.map(work, todo))
    if buf:
        append_jsonl(OUT, buf)

    all_rows = read_jsonl(OUT)
    ok = sum(1 for r in all_rows if r.get("text"))
    print(f"\n{ok}/{len(all_rows)} transcripts retrieved -> {OUT}")
    if ok < len(all_rows) * 0.5:
        errs = {}
        for r in all_rows:
            if not r.get("text"):
                errs[r.get("error", "?")] = errs.get(r.get("error", "?"), 0) + 1
        print("failures:", errs)
        print("If most are IpBlocked/TooManyRequests, rerun later or from another IP; "
              "the teacher falls back to synthesising scripts for whatever is missing.")


if __name__ == "__main__":
    main()
