"""Instruction backtranslation: reconstruct the brief that would have produced each winner.

This is the only place a teacher model touches the data, and it deliberately does
the *easy* half. The winning post stays the ground truth for what good output looks
like; the teacher only reconstructs the input that would have asked for it, plus
whatever sibling-platform fields the real post doesn't cover.

Two rules carry most of the quality:

  1. The brief is written as a commission, not a description. If the teacher is
     allowed to say "write a post whose hook is X", the training task collapses to
     copying and the model learns nothing.
  2. Where a real transcript exists, the teacher reformats it but must not reword it.
     Otherwise every script in the dataset drifts to the teacher's voice, which is
     exactly the voice we're trying not to distill.

Batch API by default (half price, up to 24h). --sync for quick tests.

  python build/backtranslate.py --limit 20 --sync     # taste the output first
  python build/backtranslate.py                       # full curated run
  python build/backtranslate.py --split holdout       # briefs for eval
"""
import argparse
import json
import pathlib
import sys
import time

from openai import OpenAI

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

STATE = DATA / ".batch_state.json"

_str = {"type": "string"}
_strs = {"type": "array", "items": {"type": "string"}}

# The teacher emits the brief plus a full pack; the pack half reuses the exact schema
# the distilled model is later held to, so training targets and inference outputs
# cannot drift apart in shape.
SCHEMA = _obj({
    "brief": _obj({"topic": _str, "audience": _str, "angle": _str,
                   "research_context": _strs, "format_note": _str},
                  ["topic", "audience", "angle", "research_context", "format_note"]),
    "why_it_works": _str,
    "pack": PACK_SCHEMA,
}, ["brief", "why_it_works", "pack"])

SYSTEM = """\
You reconstruct creative briefs from finished social media posts.

You are shown one real post and the niche it belongs to. Produce three things.

1. `brief` — the brief a content planner would have written BEFORE this post existed.
   Write it as a commission to a writer, never as a description of the finished post.
   - Do NOT quote, paraphrase, or telegraph the post's hook, punchlines, or closing line.
     A writer following your brief should be able to arrive at this post, but must not
     be handed it.
   - `research_context` holds 3-5 concrete factual points a researcher would surface
     for this topic: real numbers, mechanisms, common mistakes, named things. Never
     placeholder instructions like "research the statistics".
   - `format_note` describes the intended shape (length, delivery, on-screen text).
   - Never mention performance, engagement, virality, or that the post did well.

2. `why_it_works` — one sentence on what makes the opening earn attention. This is
   analysis for the dataset's benefit and is not part of the brief.

3. `pack` — the full multi-platform set for this one piece of content.
   - If a TRANSCRIPT is supplied, it is the creator's actual spoken words. Split it
     into `hook` (the opening line that stops the scroll) and `script` (the remainder,
     broken into short beats separated by newlines). Fix speech-to-text errors and add
     punctuation. Do NOT reword, condense, sharpen, or improve it. Preserve the
     creator's own phrasing, vocabulary and rhythm exactly.
   - If no transcript is supplied, write a hook and a 30-60 second script that the
     supplied post material implies.
   - Fill every platform field. Fields the real post already provides will be
     overwritten downstream, so spend your effort on the ones it does not.
   - Match each platform's native conventions: YouTube titles are searchable and
     concrete; Instagram captions open with a hook line then breathe across short
     paragraphs; Pinterest is descriptive and keyword-led because it is a search
     engine; `long_caption` is the carousel/thread version that stands alone as text.
   - Write in the register of the niche and the creator, not in a corporate voice.
"""


def render(rec, transcript):
    t = rec.get("text") or {}
    lines = [f"PLATFORM: {rec['platform']}", f"NICHE: {rec['niche']}"]
    if rec.get("duration_s"):
        lines.append(f"DURATION: {rec['duration_s']}s")
    if rec["platform"] == "youtube":
        lines += [f"TITLE: {t.get('title','')}",
                  f"DESCRIPTION: {(t.get('description') or '')[:1200]}",
                  f"TAGS: {', '.join((t.get('tags') or [])[:20])}"]
    elif rec["platform"] == "instagram":
        lines += [f"CAPTION: {(t.get('caption') or '')[:1800]}",
                  f"HASHTAGS: {', '.join((t.get('hashtags') or [])[:25])}"]
    else:
        lines += [f"PIN TITLE: {t.get('title','')}",
                  f"PIN DESCRIPTION: {(t.get('description') or '')[:1200]}"]
    if transcript:
        lines.append(f"TRANSCRIPT (creator's actual words): {transcript[:5000]}")
    return "\n".join(lines)


def request_body(rec, transcript):
    return {
        "model": TEACHER_MODEL,
        "messages": [{"role": "system", "content": SYSTEM},
                     {"role": "user", "content": render(rec, transcript)}],
        "response_format": {"type": "json_schema",
                            "json_schema": {"name": "backtranslation",
                                            "schema": SCHEMA, "strict": True}},
        "max_completion_tokens": 2000,
    }


def parse(rec_id, content):
    try:
        d = json.loads(content)
    except json.JSONDecodeError:
        return None
    return {"id": rec_id, "brief": d["brief"], "why_it_works": d["why_it_works"],
            "synth": d["pack"]}


def run_sync(client, todo, transcripts, out):
    buf = []
    for i, rec in enumerate(todo, 1):
        try:
            r = client.chat.completions.create(**request_body(rec, transcripts.get(rec["id"])))
            row = parse(rec["id"], r.choices[0].message.content)
            if row:
                buf.append(row)
        except Exception as e:
            print(f"  {rec['id']}: {type(e).__name__} {str(e)[:120]}", file=sys.stderr)
        if len(buf) >= 10:
            append_jsonl(out, buf)
            buf.clear()
        if i % 10 == 0 or i == len(todo):
            print(f"  {i}/{len(todo)}")
    if buf:
        append_jsonl(out, buf)


def run_batch(client, todo, transcripts, out, poll):
    """Half price, but a batch can take hours — so the batch id is written to disk
    before we start waiting, and a rerun picks the same batch back up instead of
    paying for it twice."""
    state = json.loads(STATE.read_text()) if STATE.exists() else {}
    batch_id = state.get("batch_id")

    if not batch_id:
        reqs = [{"custom_id": r["id"], "method": "POST", "url": "/v1/chat/completions",
                 "body": request_body(r, transcripts.get(r["id"]))} for r in todo]
        tmp = DATA / ".batch_input.jsonl"
        tmp.write_text("\n".join(json.dumps(r) for r in reqs) + "\n")
        print(f"uploading {len(reqs)} requests ({tmp.stat().st_size/1e6:.1f} MB)")
        up = client.files.create(file=tmp.open("rb"), purpose="batch")
        batch = client.batches.create(input_file_id=up.id,
                                      endpoint="/v1/chat/completions",
                                      completion_window="24h")
        batch_id = batch.id
        STATE.write_text(json.dumps({"batch_id": batch_id, "out": str(out)}))
        print(f"batch {batch_id} submitted — safe to Ctrl-C, rerun resumes polling")

    while True:
        b = client.batches.retrieve(batch_id)
        c = b.request_counts
        print(f"  {b.status}  {c.completed}/{c.total} done, {c.failed} failed")
        if b.status in ("completed", "failed", "expired", "cancelled"):
            break
        time.sleep(poll)

    if b.error_file_id:
        errs = client.files.content(b.error_file_id).text.strip().splitlines()
        print(f"{len(errs)} requests errored; first: {errs[0][:200]}", file=sys.stderr)
    if not b.output_file_id:
        STATE.unlink(missing_ok=True)
        sys.exit(f"batch ended {b.status} with no output")

    rows = []
    for line in client.files.content(b.output_file_id).text.splitlines():
        if not line.strip():
            continue
        d = json.loads(line)
        body = (d.get("response") or {}).get("body") or {}
        choices = body.get("choices") or []
        if choices:
            row = parse(d["custom_id"], choices[0]["message"]["content"])
            if row:
                rows.append(row)
    append_jsonl(out, rows)
    STATE.unlink(missing_ok=True)
    (DATA / ".batch_input.jsonl").unlink(missing_ok=True)
    print(f"{len(rows)} briefs written")


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--split", choices=["curated", "holdout"], default="curated")
    ap.add_argument("--sync", action="store_true", help="skip the batch API (2x cost)")
    ap.add_argument("--limit", type=int, help="cap requests — use for a taste test")
    ap.add_argument("--poll", type=int, default=60)
    args = ap.parse_args()

    load_env()
    src = DATA / f"{args.split}.jsonl"
    out = DATA / ("briefs.jsonl" if args.split == "curated" else "briefs_holdout.jsonl")
    rows = read_jsonl(src)
    if not rows:
        sys.exit(f"No {src} — run build/score.py first.")

    transcripts = {r["id"]: r["text"] for r in read_jsonl(DATA / "transcripts.jsonl")
                   if r.get("text")}
    have = seen_ids(out)
    todo = [r for r in rows if r["id"] not in have]
    if args.limit:
        todo = todo[:args.limit]
    n_tr = sum(1 for r in todo if transcripts.get(r["id"]))
    print(f"{len(todo)} to backtranslate ({len(have)} cached), "
          f"{n_tr} with real transcripts")
    if not todo and not STATE.exists():
        return

    client = OpenAI()
    if args.sync:
        run_sync(client, todo, transcripts, out)
    else:
        run_batch(client, todo, transcripts, out, args.poll)


if __name__ == "__main__":
    main()
