"""YouTube Data API v3 collector.

The free quota is 10,000 units/day and search.list costs 100 units per call while
playlistItems.list costs 1. So we never enumerate videos with search — we resolve
each channel once, then page its uploads playlist. That single choice is the
difference between ~90 videos/day and ~15,000/day.

  python collect/youtube.py --discover --niche finance   # find candidate channels
  python collect/youtube.py --niche finance              # collect from creators.yaml
  python collect/youtube.py                              # all niches
"""
import argparse
import re
import sys
import time

import requests
import yaml

sys.path.insert(0, str(__import__("pathlib").Path(__file__).resolve().parent.parent))
from common import (MAX_DURATION_S, NICHES, RAW, ROOT, append_jsonl, need,  # noqa: E402
                    seen_ids)

API = "https://www.googleapis.com/youtube/v3"
OUT = RAW / "youtube.jsonl"
_DUR = re.compile(r"P(?:(\d+)D)?T?(?:(\d+)H)?(?:(\d+)M)?(?:(\d+)S)?")


class QuotaExceeded(Exception):
    pass


class YT:
    def __init__(self, key, ceiling=9500):
        self.key = key
        self.units = 0
        self.ceiling = ceiling
        self.session = requests.Session()

    def get(self, endpoint, cost, **params):
        if self.units + cost > self.ceiling:
            raise QuotaExceeded(f"local ceiling {self.ceiling} reached")
        params["key"] = self.key
        for attempt in range(4):
            r = self.session.get(f"{API}/{endpoint}", params=params, timeout=30)
            if r.status_code == 200:
                self.units += cost
                return r.json()
            body = r.json() if r.headers.get("content-type", "").startswith("application/json") else {}
            reason = ""
            errs = body.get("error", {}).get("errors") or []
            if errs:
                reason = errs[0].get("reason", "")
            if reason in ("quotaExceeded", "dailyLimitExceeded"):
                raise QuotaExceeded(reason)
            if r.status_code in (403, 429, 500, 503) and attempt < 3:
                time.sleep(2 ** attempt)
                continue
            raise RuntimeError(f"{endpoint} {r.status_code}: {r.text[:300]}")
        raise RuntimeError(f"{endpoint}: retries exhausted")


def parse_duration(iso):
    """PT1M30S -> 90. Returns None when YouTube gives us something unparseable."""
    if not iso:
        return None
    m = _DUR.fullmatch(iso)
    if not m:
        return None
    d, h, mi, s = (int(x) if x else 0 for x in m.groups())
    return d * 86400 + h * 3600 + mi * 60 + s


def resolve_channel(yt, ref):
    """@handle or UC... -> channel metadata. 1 quota unit."""
    params = {"part": "snippet,statistics,contentDetails"}
    if ref.startswith("UC") and len(ref) == 24:
        params["id"] = ref
    else:
        params["forHandle"] = ref if ref.startswith("@") else "@" + ref
    data = yt.get("channels", 1, **params)
    items = data.get("items") or []
    if not items:
        return None
    it = items[0]
    return {
        "id": it["id"],
        "name": it["snippet"]["title"],
        "subs": int(it["statistics"].get("subscriberCount", 0) or 0),
        "uploads": it["contentDetails"]["relatedPlaylists"].get("uploads"),
    }


def uploads_video_ids(yt, playlist_id, limit):
    """Page the uploads playlist. 1 unit per 50 videos."""
    ids, token = [], None
    while len(ids) < limit:
        data = yt.get("playlistItems", 1, part="contentDetails",
                      playlistId=playlist_id, maxResults=50,
                      **({"pageToken": token} if token else {}))
        for it in data.get("items", []):
            vid = it["contentDetails"].get("videoId")
            if vid:
                ids.append(vid)
        token = data.get("nextPageToken")
        if not token:
            break
    return ids[:limit]


def fetch_videos(yt, video_ids):
    """1 unit per 50 videos."""
    out = []
    for i in range(0, len(video_ids), 50):
        chunk = video_ids[i:i + 50]
        data = yt.get("videos", 1, part="snippet,statistics,contentDetails",
                      id=",".join(chunk), maxResults=50)
        out.extend(data.get("items", []))
    return out


def to_record(v, channel, niche):
    dur = parse_duration(v["contentDetails"].get("duration"))
    if dur is None or dur > MAX_DURATION_S or dur < 3:
        return None
    sn, st = v["snippet"], v.get("statistics", {})
    views = int(st.get("viewCount", 0) or 0)
    if views < 100:
        return None  # too little signal to score
    return {
        "id": f"yt:{v['id']}",
        "platform": "youtube",
        "niche": niche,
        "creator": channel["id"],
        "creator_name": channel["name"],
        "creator_size": channel["subs"],
        "published_at": sn["publishedAt"],
        "duration_s": dur,
        "url": f"https://www.youtube.com/watch?v={v['id']}",
        "text": {
            "title": sn.get("title", ""),
            "description": sn.get("description", ""),
            "tags": sn.get("tags", []) or [],
        },
        "metrics": {
            "views": views,
            "likes": int(st.get("likeCount", 0) or 0),
            "comments": int(st.get("commentCount", 0) or 0),
        },
    }


def discover(yt, niche, cfg):
    """search.list at 100 units/query. Writes candidates for you to prune by hand
    rather than editing creators.yaml in place — that file has comments worth keeping."""
    queries = cfg.get("discover_queries") or []
    known = {c.get("youtube") for c in (cfg.get("creators") or []) if c.get("youtube")}
    found = {}
    for q in queries:
        try:
            data = yt.get("search", 100, part="snippet", q=q, type="channel",
                          maxResults=25, relevanceLanguage="en")
        except QuotaExceeded:
            print(f"  quota reached during discovery at query {q!r}", file=sys.stderr)
            break
        for it in data.get("items", []):
            cid = it["snippet"].get("channelId") or it.get("id", {}).get("channelId")
            title = it["snippet"].get("title", "")
            if cid and cid not in known:
                found[cid] = title
        print(f"  {q!r}: {len(found)} unique candidates so far ({yt.units} units)")

    dest = RAW.parent / f"discovered_{niche}.yaml"
    lines = [f"# Candidates for {niche}. Prune hard, then paste under",
             f"# creators.yaml -> {niche}: creators:", ""]
    for cid, title in found.items():
        lines.append(f"  - youtube: {cid}    # {title}")
        lines.append("    instagram:")
        lines.append("    pinterest:")
    dest.parent.mkdir(parents=True, exist_ok=True)
    dest.write_text("\n".join(lines) + "\n")
    print(f"  wrote {len(found)} candidates -> {dest}")


def collect(yt, niche, cfg, per_channel):
    creators = [c for c in (cfg.get("creators") or []) if c.get("youtube")]
    if not creators:
        print(f"  {niche}: no youtube creators in creators.yaml — run --discover first")
        return 0
    have = seen_ids(OUT)
    total = 0
    for c in creators:
        ref = str(c["youtube"])
        try:
            ch = resolve_channel(yt, ref)
        except QuotaExceeded:
            raise
        except Exception as e:
            print(f"  {ref}: resolve failed ({e}) — skipping")
            continue
        if not ch or not ch["uploads"]:
            print(f"  {ref}: no uploads playlist — skipping")
            continue

        vids = uploads_video_ids(yt, ch["uploads"], per_channel)
        fresh = [v for v in vids if f"yt:{v}" not in have]
        if not fresh:
            print(f"  {ch['name']}: nothing new")
            continue

        rows = []
        for v in fetch_videos(yt, fresh):
            rec = to_record(v, ch, niche)
            if rec:
                rows.append(rec)
                have.add(rec["id"])
        append_jsonl(OUT, rows)   # checkpoint per channel
        total += len(rows)
        print(f"  {ch['name']}: +{len(rows)} short-form of {len(fresh)} new "
              f"({yt.units} units)")
    return total


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--niche", choices=list(NICHES), help="default: all")
    ap.add_argument("--discover", action="store_true",
                    help="find candidate channels (100 units per query)")
    ap.add_argument("--limit", type=int, default=250,
                    help="max videos per channel (default 250)")
    ap.add_argument("--quota", type=int, default=9500,
                    help="stop before spending this many units (default 9500)")
    args = ap.parse_args()

    cfg = yaml.safe_load((ROOT / "creators.yaml").read_text()) or {}
    yt = YT(need("YOUTUBE_API_KEY"), ceiling=args.quota)
    niches = [args.niche] if args.niche else list(NICHES)

    total = 0
    try:
        for n in niches:
            print(f"{n}:")
            if args.discover:
                discover(yt, n, cfg.get(n) or {})
            else:
                total += collect(yt, n, cfg.get(n) or {}, args.limit)
    except QuotaExceeded as e:
        print(f"\nquota exhausted ({e}). Progress is saved — rerun tomorrow "
              f"and it resumes where it stopped.", file=sys.stderr)
    print(f"\n{total} new records -> {OUT}   ({yt.units} quota units used)")


if __name__ == "__main__":
    main()
