#!/usr/bin/env -S uv run --quiet --script
# /// script
# requires-python = ">=3.10"
# dependencies = ["pyyaml", "piper-tts"]
# ///
"""Render one narration clip per scene, and prove it says what you wrote.

Engine selection is automatic: a cloud voice when a key is available, a
local neural voice (Piper) when there isn't one. The local path needs no
key, no network after the first voice download, and runs on macOS and
Linux alike - so a container with no secrets in it can still narrate.

Input is a scenes file: a YAML list of scenes, each with `id` and
`narration`. Output is OUTDIR/<id>.wav plus OUTDIR/manifest.json carrying
the exact text and measured duration of each clip, which is what
make-subtitles and the assembly step both read.

Usage:
  narrate SCENES.yaml OUTDIR [--engine auto|openai|openai-chat|piper]
                             [--voice NAME] [--force]
"""

import argparse
import base64
import difflib
import json
import os
import shutil
import subprocess
import sys
import tempfile
import unicodedata
import urllib.request
import uuid
import wave
from pathlib import Path

import yaml

OPENAI_TTS_MODEL = "gpt-4o-mini-tts"      # deterministic: reads what you send
OPENAI_CHAT_MODEL = "gpt-audio-1.5"       # better prosody, will ad-lib; gated
PIPER_VOICE = "en_US-lessac-medium"


def die(msg):
    print(f"narrate: {msg}", file=sys.stderr)
    sys.exit(1)


def openai_key():
    key = os.environ.get("OPENAI_API_KEY")
    if key:
        return key.strip()
    try:
        out = subprocess.run(["llm", "keys", "get", "openai"],
                             capture_output=True, text=True, timeout=15)
        if out.returncode == 0 and out.stdout.strip():
            return out.stdout.strip()
    except Exception:  # noqa: BLE001 - llm not installed is a normal outcome
        pass
    return None


SEGMENTATION_SCRIPTS = (
    (0x3040, 0x30FF),  # Hiragana and Katakana
    (0x3400, 0x4DBF),  # CJK Extension A
    (0x4E00, 0x9FFF),  # CJK Unified Ideographs
    (0xF900, 0xFAFF),  # CJK Compatibility Ideographs
    (0x20000, 0x323AF),  # CJK Unified Ideograph extensions B through I
    (0x2F800, 0x2FA1F),  # CJK Compatibility Ideographs Supplement
    (0x0E00, 0x0E7F),  # Thai
)


def norm(s):
    """Return comparable words without changing narration/cache identity."""
    words = []
    for token in unicodedata.normalize("NFKC", s).casefold().split():
        word = "".join(
            char for char in token
            if unicodedata.category(char)[0] in {"L", "N", "M"}
        )
        if word:
            words.append(word)
    return words


def requires_segmentation(s):
    return any(start <= ord(char) <= end
               for char in unicodedata.normalize("NFKC", s)
               for start, end in SEGMENTATION_SCRIPTS)


ASR_SNIPPET = """
import sys
from faster_whisper import WhisperModel
m = WhisperModel(sys.argv[2], device="cpu", compute_type="int8")
segs, _ = m.transcribe(sys.argv[1])
from pathlib import Path
import json
Path(sys.argv[3]).write_text(json.dumps({"text": " ".join(s.text.strip() for s in segs)}), encoding="utf-8")
"""


def transcribe_local(wav, model="base.en"):
    """Transcribe with a local ASR, in its own uv env so narrate stays light.
    Returns None when faster-whisper isn't available."""
    try:
        with tempfile.TemporaryDirectory(prefix="movie-asr-") as temporary:
            result_path = Path(temporary) / "transcript.json"
            out = subprocess.run(
                ["uv", "run", "--no-config", "--no-project", "--isolated",
                 "--python", sys.executable, "--with", "faster-whisper", "python",
                 "-c", ASR_SNIPPET, str(wav.resolve()), model, str(result_path)],
                cwd=temporary, capture_output=True, text=True,
                encoding="utf-8", errors="replace", timeout=900)
            for diagnostic in (out.stdout, out.stderr):
                if diagnostic.strip():
                    print(f"local ASR: {diagnostic.strip()[-600:]}", file=sys.stderr)
            if out.returncode != 0:
                print(f"local ASR exited with status {out.returncode}", file=sys.stderr)
                return None
            data = json.loads(result_path.read_text(encoding="utf-8"))
            text = data.get("text") if isinstance(data, dict) else None
            return text.strip() if isinstance(text, str) else None
    except (OSError, ValueError, subprocess.SubprocessError) as error:
        print(f"local ASR unavailable: {error}", file=sys.stderr)
        return None


def structural_drift(text, heard):
    """How far a transcript diverges from the script, ignoring the noise an
    ASR always makes.

    Exact word-matching is the wrong tool here: a small model mangles
    unusual names ("smevals" -> "Mevil"), and - worse - a *dropped* word
    scores as more similar than two mispronounced ones. What is detectable,
    and what actually matters, is missing or invented CONTENT: a sentence
    the voice skipped, or a preamble it invented. Returns
    (length_delta_fraction, longest_run_of_missing_added_or_changed_words).
    """
    if requires_segmentation(text) or requires_segmentation(heard):
        return None
    want, got = norm(text), norm(heard)
    if not want:
        return None
    if not got:
        return 1.0, len(want)
    delta = abs(len(got) - len(want)) / max(1, len(want))
    ops = difflib.SequenceMatcher(a=want, b=got).get_opcodes()
    worst = max((max(i2 - i1, j2 - j1) for tag, i1, i2, j1, j2 in ops if tag != "equal"),
                default=0)
    return delta, worst


def post(url, key, body, want_json=True):
    req = urllib.request.Request(
        url, data=json.dumps(body).encode(),
        headers={"Authorization": f"Bearer {key}", "Content-Type": "application/json"})
    with urllib.request.urlopen(req, timeout=180) as r:
        return json.load(r) if want_json else r.read()


def say_openai(key, text, out_wav, voice):
    data = post("https://api.openai.com/v1/audio/speech", key,
                {"model": OPENAI_TTS_MODEL, "voice": voice,
                 "input": text, "response_format": "wav"}, want_json=False)
    out_wav.write_bytes(data)
    return None                       # deterministic engine: nothing to gate


def say_openai_chat(key, text, out_wav, voice):
    doc = post("https://api.openai.com/v1/chat/completions", key, {
        "model": OPENAI_CHAT_MODEL,
        "modalities": ["text", "audio"],
        "audio": {"voice": voice, "format": "wav"},
        "messages": [{"role": "user", "content":
                      "Read this narration aloud, warm and clear, verbatim, "
                      "and say nothing else:\n\n" + text}],
    })
    audio = doc["choices"][0]["message"]["audio"]
    out_wav.write_bytes(base64.b64decode(audio["data"]))
    return audio.get("transcript", "")


def say_piper(text, out_wav, voice):
    from piper import PiperVoice
    from piper.download_voices import download_voice
    home = Path(os.environ.get("PIPER_VOICE_DIR",
                               Path.home() / ".cache" / "piper-voices"))
    home.mkdir(parents=True, exist_ok=True)
    name = voice
    onnx = home / f"{name}.onnx"
    if not onnx.exists():
        print(f"  downloading local voice {name} (one time)…")
        download_voice(name, home)
    v = PiperVoice.load(str(onnx))
    with wave.open(str(out_wav), "wb") as w:
        v.synthesize_wav(text, w)
    return None


def duration(path):
    out = subprocess.run(
        ["ffprobe", "-v", "error", "-show_entries", "format=duration",
         "-of", "csv=p=0", str(path)], capture_output=True, text=True, encoding="utf-8", errors="replace")
    return round(float(out.stdout.strip()), 3)


def accepted_wav(outdir, entry):
    wav = entry.get("wav") if isinstance(entry, dict) else None
    if not isinstance(wav, str) or not wav:
        return None
    candidate = Path(wav)
    if candidate.is_absolute():
        return None
    return outdir / candidate


def write_manifest(manifest_path, entries):
    with tempfile.NamedTemporaryFile(dir=manifest_path.parent, delete=False) as handle:
        temporary = Path(handle.name)
    try:
        temporary.write_text(json.dumps(entries, indent=2), encoding="utf-8")
        temporary.replace(manifest_path)
    finally:
        temporary.unlink(missing_ok=True)


def main():
    for stream in (sys.stdout, sys.stderr):
        if hasattr(stream, "reconfigure"):
            stream.reconfigure(errors="backslashreplace")
    if len(sys.argv) == 4 and sys.argv[1] == "--drift-check":
        script = Path(sys.argv[2]).read_text(encoding="utf-8-sig")
        heard = Path(sys.argv[3]).read_text(encoding="utf-8-sig")
        drift = structural_drift(script, heard)
        if drift is None:
            print("comparison unavailable -> MISMATCH")
            return 1
        delta, worst = drift
        bad = delta > 0.15 or worst >= 4
        print(f"length change {delta:.0%}, worst run {worst} -> "
              f"{'MISMATCH' if bad else 'ok'}")
        return 1 if bad else 0

    ap = argparse.ArgumentParser()
    ap.add_argument("scenes", type=Path)
    ap.add_argument("outdir", type=Path)
    ap.add_argument("--engine", default="auto",
                    choices=["auto", "openai", "openai-chat", "piper"])
    ap.add_argument("--voice", default=None)
    ap.add_argument("--force", action="store_true")
    ap.add_argument("--verify", default="auto", choices=["auto", "on", "off"],
                    help="local ASR: on requires verification; auto tries it "
                         "for piper/openai-chat but allows unavailable ASR; "
                         "off skips ASR (default: auto)")
    ap.add_argument("--asr-model", default="base.en")
    args = ap.parse_args()

    doc = yaml.safe_load(args.scenes.read_text(encoding="utf-8-sig"))
    scenes = [s for s in doc.get("scenes", []) if (s.get("narration") or "").strip()]
    if not scenes:
        die("no scenes with narration")
    if not shutil.which("ffprobe"):
        die("ffprobe not on PATH (narration durations cannot be measured)")

    key = openai_key()
    engine = args.engine
    if engine == "auto":
        engine = "openai" if key else "piper"
    if engine.startswith("openai") and not key:
        die("no OPENAI_API_KEY (and `llm keys get openai` found nothing). "
            "Use --engine piper for a local voice.")
    print(f"engine: {engine}" + ("" if key or engine == "piper" else ""))
    voice = args.voice or (PIPER_VOICE if engine == "piper" else "nova")
    synthesis = {"engine": engine, "voice": voice, "model": {
        "openai": OPENAI_TTS_MODEL,
        "openai-chat": OPENAI_CHAT_MODEL,
        "piper": voice,
    }[engine]}

    # a deterministic cloud endpoint reads exactly what you send it, so the
    # ear-check is optional there; anything else gets listened to by default
    verify = args.verify == "on" or (args.verify == "auto" and engine != "openai")

    args.outdir.mkdir(parents=True, exist_ok=True)
    # Cache only clips rendered from the requested text and synthesis settings.
    prior = {}
    prior_path = args.outdir / "manifest.json"
    if prior_path.exists():
        try:
            prior = {e["id"]: e for e in
                     json.loads(prior_path.read_text(encoding="utf-8-sig"))}
        except Exception:  # noqa: BLE001 - a corrupt manifest just means no cache
            prior = {}
    current_ids = {sc["id"] for sc in scenes}
    manifest = {sid: entry for sid, entry in prior.items() if sid in current_ids}
    write_manifest(prior_path, list(manifest.values()))
    failures = []

    def withdraw(sid):
        if sid in manifest:
            del manifest[sid]
        write_manifest(prior_path, list(manifest.values()))

    def accept(entry):
        manifest[entry["id"]] = entry
        write_manifest(prior_path, list(manifest.values()))

    for sc in scenes:
        sid = sc["id"]
        text = " ".join((sc["narration"] or "").split())
        previous = prior.get(sid, {})
        previous_wav = accepted_wav(args.outdir, previous)
        cached = (previous_wav is not None and previous_wav.exists() and not args.force
                  and previous.get("text") == text
                  and previous.get("synthesis") == synthesis)
        if cached:
            print(f"{sid}: cached")
        elif previous_wav is not None and previous_wav.exists() and not args.force and sid in prior:
            print(f"{sid}: text or synthesis settings changed - redoing")
        if engine == "openai-chat" and structural_drift(text, text) is None:
            print(f"{sid}: chat transcript comparison unavailable", file=sys.stderr)
            withdraw(sid)
            failures.append(sid)
            continue
        withdraw(sid)
        accepted = False
        for attempt in ((1,) if cached else (1, 2)):
            claimed = None
            candidate = previous_wav if cached else args.outdir / (
                f".{sid}.attempt-{uuid.uuid4().hex}.wav"
            )
            if not cached:
                try:
                    if engine == "openai":
                        claimed = say_openai(key, text, candidate, voice)
                    elif engine == "openai-chat":
                        claimed = say_openai_chat(key, text, candidate, voice)
                    else:
                        claimed = say_piper(text, candidate, voice)
                except Exception as error:  # rejected bytes remain as evidence
                    print(f"{sid}: synthesis failed (attempt {attempt}: {error})", file=sys.stderr)
                    continue

            # Preserve the engine transcript gate and the ASR drift thresholds.
            if engine == "openai-chat" and not cached:
                if not isinstance(claimed, str) or not norm(claimed):
                    print(f"{sid}: chat transcript contains no speech", file=sys.stderr)
                    continue
                drift_result = structural_drift(text, claimed)
                if drift_result is None:
                    print(f"{sid}: chat transcript comparison unavailable", file=sys.stderr)
                    break
                want, got = norm(text), norm(claimed)
                drift = abs(len(want) - len(got)) + sum(
                    1 for a, b in zip(want, got) if a != b)
                if drift > max(2, len(want) // 25):
                    print(f"{sid}: engine ad-libbed (attempt {attempt}, drift {drift})")
                    continue
            if verify:
                heard = transcribe_local(candidate, args.asr_model)
                if heard is None:
                    if args.verify == "on":
                        print(f"{sid}: required verification unavailable", file=sys.stderr)
                        break
                    else:
                        print(f"{sid}: verification unavailable (no local ASR)")
                else:
                    drift_result = structural_drift(text, heard)
                    if drift_result is None:
                        if args.verify == "on":
                            print(f"{sid}: required verification unavailable", file=sys.stderr)
                            break
                        print(f"{sid}: verification unavailable (unsupported script)")
                    else:
                        delta, worst = drift_result
                        if delta > 0.15 or worst >= 4:
                            print(f"{sid}: what came out does not match the script "
                                  f"(attempt {attempt}: {delta:.0%} length change, "
                                  f"{worst} words in a row wrong)")
                            print(f"       heard: {heard[:120]}")
                            continue
                        print(f"{sid}: ok (verified by ear: {delta:.0%} length "
                              f"change, worst run {worst})")
            if not verify:
                print(f"{sid}: ok")
            try:
                clip_duration = duration(candidate)
            except Exception as error:  # unaccepted bytes remain as evidence
                print(f"{sid}: duration failed ({error})", file=sys.stderr)
                break
            if not cached:
                candidate.replace(args.outdir / f"{sid}.wav")
                candidate = args.outdir / f"{sid}.wav"
            wav_name = previous["wav"] if cached else candidate.name
            accept({"id": sid, "text": text, "wav": wav_name,
                    "duration": clip_duration, "synthesis": synthesis})
            accepted = True
            break
        if not accepted:
            withdraw(sid)
            failures.append(sid)

    entries = list(manifest.values())
    total = sum(m["duration"] for m in entries)
    print(f"\n{len(entries)} clips, {total:.1f}s total -> {args.outdir}/manifest.json")
    if engine == "piper":
        print("local voice: it mispronounces unusual names rather than dropping "
              "them - listen to one clip before you commit to a voice.")
    if verify:
        print("the ear-check catches missing or invented sentences, not "
              "pronunciation: an ASR mangles jargon too.")
    if failures:
        print(f"FAILED verbatim delivery: {failures}", file=sys.stderr)
        return 1
    return 0


if __name__ == "__main__":
    sys.exit(main())
