"""Assemble training examples: brief in, multi-platform pack out.

Real data always beats synthesised data, so wherever the corpus actually contains a
field it overwrites whatever the teacher wrote. Where the same creator posted the
same content to several platforms, all of those real fields land in one example —
that is the ideal training row, and provenance is recorded so eval can report on
real and synthesised fields separately.

The split is by *creator*, not by row. Splitting by row puts the same creator on both
sides, the model recognises the voice it already memorised, and val loss lies to you.

  python build/pack.py
"""
import argparse
import datetime as dt
import hashlib
import json
import re
import sys
from collections import defaultdict

import yaml

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

PAIR_WINDOW_DAYS = 7
PAIR_MIN_JACCARD = 0.30


def norm(s):
    return re.sub(r"[^a-z0-9]", "", str(s or "").lower())


def creator_index():
    """Canonical creator id per (platform, handle), from creators.yaml.

    Needed because the same person is a channel *title* on YouTube and a lowercase
    *handle* on Instagram. Keying on either alone would both miss real cross-platform
    pairs and, worse, let one creator land in train and val at the same time.
    """
    path = ROOT / "creators.yaml"
    if not path.exists():
        return {}
    cfg = yaml.safe_load(path.read_text()) or {}
    index = {}
    for niche, block in cfg.items():
        for i, c in enumerate(((block or {}).get("creators") or [])):
            key = f"{niche}:{i}"
            for platform in ("youtube", "instagram", "pinterest"):
                if c.get(platform):
                    index[norm(c[platform])] = key
    return index


def canonical(rec, index):
    """Prefer the explicit creators.yaml mapping; fall back to a normalised name."""
    for candidate in (rec.get("creator"), rec.get("creator_name")):
        hit = index.get(norm(candidate))
        if hit:
            return hit
    return norm(rec.get("creator_name")) or norm(rec.get("creator"))


def post_text(rec):
    t = rec.get("text") or {}
    return " ".join(str(v) for v in t.values() if isinstance(v, str))


def when(rec):
    try:
        s = str(rec.get("published_at", "")).replace("Z", "+00:00")
        d = dt.datetime.fromisoformat(s)
        return d if d.tzinfo else d.replace(tzinfo=dt.timezone.utc)
    except (ValueError, TypeError):
        return None


def jaccard(a, b):
    if not a or not b:
        return 0.0
    return len(a & b) / len(a | b)


def find_siblings(rows, index):
    """Map post id -> [posts on other platforms that are the same content].

    Same creator, within a week, and overlapping vocabulary. Creators who repurpose
    one video across platforms are common; this finds them cheaply with token overlap
    instead of paying for embeddings.
    """
    by_creator = defaultdict(list)
    for r in rows:
        by_creator[(r["niche"], canonical(r, index))].append(r)

    sibs = defaultdict(list)
    for group in by_creator.values():
        if len(group) < 2:
            continue
        toks = {r["id"]: words(post_text(r)) for r in group}
        times = {r["id"]: when(r) for r in group}
        for i, a in enumerate(group):
            for b in group[i + 1:]:
                if a["platform"] == b["platform"]:
                    continue
                ta, tb = times[a["id"]], times[b["id"]]
                if ta and tb and abs((ta - tb).days) > PAIR_WINDOW_DAYS:
                    continue
                if jaccard(toks[a["id"]], toks[b["id"]]) >= PAIR_MIN_JACCARD:
                    sibs[a["id"]].append(b)
                    sibs[b["id"]].append(a)
    return sibs


def apply_real(pack, rec, prov):
    """Overwrite synthesised fields with what the platform actually published."""
    t = rec.get("text") or {}
    p = rec["platform"]
    if p == "youtube":
        if t.get("title"):
            pack["youtube"]["title"] = t["title"]
            prov.add("youtube.title")
        if (t.get("description") or "").strip():
            pack["youtube"]["description"] = t["description"]
            prov.add("youtube.description")
        if t.get("tags"):
            pack["youtube"]["tags"] = t["tags"]
            prov.add("youtube.tags")
    elif p == "instagram":
        if (t.get("caption") or "").strip():
            pack["instagram"]["caption"] = t["caption"]
            prov.add("instagram.caption")
        if t.get("hashtags"):
            pack["instagram"]["hashtags"] = t["hashtags"]
            prov.add("instagram.hashtags")
    elif p == "pinterest":
        if t.get("title"):
            pack["pinterest"]["title"] = t["title"]
            prov.add("pinterest.title")
        if (t.get("description") or "").strip():
            pack["pinterest"]["description"] = t["description"]
            prov.add("pinterest.description")


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--val-frac", type=float, default=0.10)
    ap.add_argument("--out-prefix", default="")
    args = ap.parse_args()

    curated = read_jsonl(DATA / "curated.jsonl")
    briefs = {b["id"]: b for b in read_jsonl(DATA / "briefs.jsonl")}
    has_transcript = {r["id"] for r in read_jsonl(DATA / "transcripts.jsonl") if r.get("text")}
    if not curated:
        sys.exit("No data/curated.jsonl — run build/score.py first.")
    if not briefs:
        sys.exit("No data/briefs.jsonl — run build/backtranslate.py first.")

    index = creator_index()
    sibs = find_siblings(curated, index)
    print(f"{len(curated)} winners, {len(briefs)} briefs, "
          f"{len(sibs)} posts with a cross-platform sibling "
          f"({len(index)} handles mapped in creators.yaml)")

    examples, seen_hooks = [], set()
    stats = defaultdict(int)
    for rec in curated:
        b = briefs.get(rec["id"])
        if not b:
            stats["no brief"] += 1
            continue
        pack = json.loads(json.dumps(b["synth"]))   # deep copy, teacher fills the gaps
        prov = set()

        apply_real(pack, rec, prov)
        for sib in sibs.get(rec["id"], []):
            apply_real(pack, sib, prov)             # real fields from the sibling posts

        if rec["id"] in has_transcript:
            prov.add("hook")
            prov.add("script")

        if not pack.get("hook", "").strip() or not pack.get("script", "").strip():
            stats["empty hook/script"] += 1
            continue

        # Near-duplicate hooks teach the model one phrase very well and nothing else.
        key = " ".join(sorted(words(pack["hook"])))
        if key in seen_hooks:
            stats["duplicate hook"] += 1
            continue
        seen_hooks.add(key)

        examples.append({
            "messages": [
                {"role": "system", "content": PACK_SYSTEM},
                {"role": "user", "content": render_brief(b["brief"], NICHES.get(rec["niche"], ""))},
                {"role": "assistant", "content": json.dumps(pack, ensure_ascii=False)},
            ],
            "meta": {"id": rec["id"], "niche": rec["niche"], "platform": rec["platform"],
                     "creator": canonical(rec, index),
                     "score": rec.get("score"), "real_fields": sorted(prov),
                     "siblings": [s["id"] for s in sibs.get(rec["id"], [])]},
        })

    if not examples:
        sys.exit("No examples built — check the stats above.")

    # Creator-level split, deterministic and stratified by niche.
    per_niche = defaultdict(set)
    for e in examples:
        per_niche[e["meta"]["niche"]].add(e["meta"]["creator"])
    val_creators = set()
    for niche, creators in per_niche.items():
        ordered = sorted(creators, key=lambda c: hashlib.sha1(c.encode()).hexdigest())
        n = max(1, round(len(ordered) * args.val_frac))
        val_creators.update(ordered[:n])

    train = [e for e in examples if e["meta"]["creator"] not in val_creators]
    val = [e for e in examples if e["meta"]["creator"] in val_creators]

    write_jsonl(DATA / f"{args.out_prefix}train.jsonl", train)
    write_jsonl(DATA / f"{args.out_prefix}val.jsonl", val)

    print("\ndropped:")
    for k, v in sorted(stats.items(), key=lambda kv: -kv[1]):
        print(f"  {v:>6}  {k}")

    real_counts = defaultdict(int)
    for e in examples:
        for f in e["meta"]["real_fields"]:
            real_counts[f] += 1
    print("\nfields backed by real data (rest are teacher-written):")
    for f, n in sorted(real_counts.items(), key=lambda kv: -kv[1]):
        print(f"  {n:>6}/{len(examples)}  {f}")

    print(f"\nper niche: " + ", ".join(
        f"{n}={sum(1 for e in examples if e['meta']['niche']==n)}" for n in NICHES))
    print(f"{len(train)} train / {len(val)} val "
          f"({len(val_creators)} creators held out for val)")


if __name__ == "__main__":
    main()
