"""Turn the raw corpus into 'winners' — the posts worth learning from.

The naive move is to rank by likes, which just surfaces the biggest accounts and
teaches the model to imitate whoever has the most followers. Instead every post is
scored against *its own creator's* baseline, so a 20k-subscriber channel's breakout
outranks a 2M-subscriber channel's routine upload.

  reach       = per-day traction, divided by that creator's median per-day traction
  engagement  = interaction rate, divided by that creator's median interaction rate
  score       = 0.6 * z(log reach) + 0.4 * z(log engagement), z-scored within niche+platform

Logs first, because both ratios are heavy-tailed and a raw z-score would let one
viral outlier dominate the whole distribution.

  python build/score.py
"""
import argparse
import datetime as dt
import hashlib
import statistics
import sys
from collections import defaultdict

import numpy as np

sys.path.insert(0, str(__import__("pathlib").Path(__file__).resolve().parent.parent))
from common import (DATA, HOLDOUT_SIZE, MIN_AGE_DAYS, MIN_POSTS_PER_CREATOR,  # noqa: E402
                    NICHES, RAW, WINNER_PERCENTILE, read_jsonl, write_jsonl)

NOW = dt.datetime.now(dt.timezone.utc)


def age_days(published):
    """Returns None when the timestamp is missing or unparseable — the caller drops
    those rather than guessing, and the run report shows how many were lost."""
    if not published:
        return None
    try:
        if isinstance(published, (int, float)):
            d = dt.datetime.fromtimestamp(float(published), dt.timezone.utc)
        else:
            s = str(published).strip().replace("Z", "+00:00")
            d = dt.datetime.fromisoformat(s)
            if d.tzinfo is None:
                d = d.replace(tzinfo=dt.timezone.utc)
    except (ValueError, OSError, OverflowError):
        return None
    return max((NOW - d).total_seconds() / 86400.0, 0.0)


def signals(rec, age):
    """(reach_per_day, engagement_rate) for one post, or None if unscoreable.

    A comment is weighted 5x a like throughout: it costs far more effort and
    correlates better with the kind of post that actually lands.
    """
    m = rec.get("metrics") or {}
    size = rec.get("creator_size") or 0
    days = max(age, 1.0)
    p = rec["platform"]

    if p == "youtube":
        views = m.get("views", 0)
        if views < 100:
            return None
        inter = m.get("likes", 0) + 5 * m.get("comments", 0)
        return views / days, inter / views

    if p == "instagram":
        inter = m.get("likes", 0) + 5 * m.get("comments", 0)
        if inter <= 0:
            return None
        views = m.get("views") or 0
        return (views or inter) / days, inter / max(size, 1)

    if p == "pinterest":
        saves = m.get("saves", 0) + m.get("reactions", 0)
        if saves <= 0:
            return None
        return saves / days, saves / max(size, 1)

    return None


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--percentile", type=float, default=WINNER_PERCENTILE)
    ap.add_argument("--holdout", type=int, default=HOLDOUT_SIZE)
    args = ap.parse_args()

    rows = []
    for f in ("youtube.jsonl", "instagram.jsonl", "pinterest.jsonl"):
        rows.extend(read_jsonl(RAW / f))
    if not rows:
        sys.exit(f"No raw data in {RAW}. Run the collectors first.")
    print(f"loaded {len(rows)} raw posts")

    dropped = defaultdict(int)
    scored = []
    for r in rows:
        age = age_days(r.get("published_at"))
        if age is None:
            dropped["no usable timestamp"] += 1
            continue
        if age < MIN_AGE_DAYS:
            dropped[f"younger than {MIN_AGE_DAYS}d"] += 1
            continue
        sig = signals(r, age)
        if sig is None:
            dropped["no usable metrics"] += 1
            continue
        r["_age"], r["_reach"], r["_eng"] = age, sig[0], sig[1]
        scored.append(r)

    # Creator baselines. A creator with too few posts has no stable median, and
    # dividing by a noisy median is worse than not normalising at all.
    by_creator = defaultdict(list)
    for r in scored:
        by_creator[(r["platform"], r["creator"])].append(r)

    relative = []
    for key, posts in by_creator.items():
        if len(posts) < MIN_POSTS_PER_CREATOR:
            dropped[f"creator under {MIN_POSTS_PER_CREATOR} posts"] += len(posts)
            continue
        med_reach = statistics.median(p["_reach"] for p in posts)
        med_eng = statistics.median(p["_eng"] for p in posts)
        if med_reach <= 0 or med_eng <= 0:
            dropped["creator median is zero"] += len(posts)
            continue
        for p in posts:
            p["_reach_rel"] = p["_reach"] / med_reach
            p["_eng_rel"] = p["_eng"] / med_eng
            relative.append(p)

    if not relative:
        sys.exit("Nothing survived filtering — check the drop report above and "
                 "collect more posts per creator.")

    # z-score within niche+platform: comparing a Pinterest save rate to a YouTube
    # view rate is meaningless, so each cohort is standardised on its own.
    cohorts = defaultdict(list)
    for r in relative:
        cohorts[(r["niche"], r["platform"])].append(r)

    winners = []
    print("\ncohort                        posts   kept   cutoff")
    for (niche, platform), posts in sorted(cohorts.items()):
        lr = np.log(np.array([p["_reach_rel"] for p in posts]) + 1e-9)
        le = np.log(np.array([p["_eng_rel"] for p in posts]) + 1e-9)
        zr = (lr - lr.mean()) / (lr.std() or 1.0)
        ze = (le - le.mean()) / (le.std() or 1.0)
        score = 0.6 * zr + 0.4 * ze
        cutoff = float(np.quantile(score, args.percentile)) if len(posts) > 4 else -np.inf
        kept = 0
        for p, s in zip(posts, score):
            p["score"] = round(float(s), 4)
            for k in ("_age", "_reach", "_eng", "_reach_rel", "_eng_rel"):
                p.pop(k, None)
            if s >= cutoff:
                winners.append(p)
                kept += 1
        print(f"{niche:<14} {platform:<12} {len(posts):>6} {kept:>6}   {cutoff:>6.2f}")

    # Holdout is carved here, before the teacher model touches anything, so eval
    # is against posts the training pipeline has never seen. Deterministic by id
    # hash so reruns produce the same split.
    winners.sort(key=lambda r: r["score"], reverse=True)
    by_cohort = defaultdict(list)
    for w in winners:
        by_cohort[(w["niche"], w["platform"])].append(w)

    holdout, per = [], max(1, args.holdout // max(len(by_cohort), 1))
    for cohort, posts in sorted(by_cohort.items()):
        posts.sort(key=lambda r: hashlib.sha1(r["id"].encode()).hexdigest())
        holdout.extend(posts[:per])
    holdout = holdout[:args.holdout]
    hold_ids = {h["id"] for h in holdout}
    curated = [w for w in winners if w["id"] not in hold_ids]

    write_jsonl(DATA / "curated.jsonl", curated)
    write_jsonl(DATA / "holdout.jsonl", holdout)

    print("\ndropped:")
    for reason, n in sorted(dropped.items(), key=lambda kv: -kv[1]):
        print(f"  {n:>7}  {reason}")
    print(f"\n{len(curated)} winners -> data/curated.jsonl")
    print(f"{len(holdout)} held out -> data/holdout.jsonl  (never shown to the teacher)")
    missing = set(NICHES) - {w["niche"] for w in winners}
    if missing:
        print(f"\nno winners at all for: {', '.join(sorted(missing))}")


if __name__ == "__main__":
    main()
