#!/usr/bin/env -S uv run --quiet --script
# /// script
# requires-python = ">=3.10"
# ///
"""Build an SRT from narrate's manifest, timed to the measured clips.

Subtitles are not decoration. A movie gets watched muted - in a PR, on a
phone, in an open-plan office, by someone who is deaf - and an unsubtitled
narrated movie simply doesn't communicate to those viewers. They also make
the movie searchable and let a reviewer check what was said without
listening.

Cue timing is proportional to character count within each scene's measured
audio, which tracks speech closely enough for reading. If you need
word-exact timing, transcribe the rendered audio with a word-timestamp API
and use those offsets instead.

Usage:
  make-subtitles MANIFEST.json OUT.srt [--offsets SCENE=SECONDS ...]
                                       [--max-chars N] [--max-secs S]
"""

import argparse
import json
import math
import sys
from pathlib import Path

MAX_CHARS = 84          # two comfortable lines
MAX_SECS = 5.5


def cue_chunks(text, max_chars):
    """Split into cue-sized pieces on sentence, then clause, then word."""
    words, chunks, cur = text.split(), [], ""
    for w in words:
        candidate = f"{cur} {w}".strip()
        if len(candidate) > max_chars and cur:
            chunks.append(cur)
            cur = w
        else:
            cur = candidate
            if cur.endswith((".", "!", "?")) and len(cur) > max_chars * 0.45:
                chunks.append(cur)
                cur = ""
    if cur:
        chunks.append(cur)
    return chunks or [text]


def scene_cues(text, start, duration, max_chars, max_secs):
    """Allocate positive millisecond cues across the complete measured scene."""
    end = start + duration
    if (not all(math.isfinite(value) for value in (start, duration, end))
            or start < 0 or duration <= 0):
        raise ValueError("start must be finite and nonnegative; "
                         "duration must be finite and positive")
    start_ms, end_ms = round(start * 1000), round(end * 1000)
    available = end_ms - start_ms
    if available <= 0:
        raise ValueError("scene has no representable millisecond subtitle interval")

    # Readability guides word-boundary splitting, never cuts off the scene tail.
    char_limit = max(1, min(max_chars,
                            int(len(text) * min(1.0, max_secs / duration))))
    chunks = cue_chunks(text, char_limit)
    if len(chunks) > available:
        chunks = [" ".join(chunks[i * len(chunks) // available:
                                  (i + 1) * len(chunks) // available])
                  for i in range(available)]
    while True:
        total_chars = sum(len(chunk) for chunk in chunks) or 1
        elapsed_chars, previous = 0, start_ms
        cues, split = [], None
        for i, chunk in enumerate(chunks):
            elapsed_chars += len(chunk)
            remaining = len(chunks) - i - 1
            boundary = start_ms + round(available * elapsed_chars / total_chars)
            # Reserve one millisecond per remaining cue, even for very uneven text.
            boundary = min(end_ms - remaining, max(previous + 1, boundary))
            if not remaining:
                boundary = end_ms
            if (split is None and boundary - previous > max_secs * 1000
                    and len(chunk.split()) > 1):
                split = i
            cues.append((previous / 1000, boundary / 1000, wrap(chunk)))
            previous = boundary
        if split is None or len(chunks) == available:
            return cues
        # Splitting removes a space from the allocation weights, so remeasure all
        # cues until every splittable chunk fits or milliseconds limit the count.
        words = chunks[split].split()
        midpoint = len(words) // 2
        chunks[split:split + 1] = [" ".join(words[:midpoint]),
                                   " ".join(words[midpoint:])]


def wrap(line, width=42):
    words, out, cur = line.split(), [], ""
    for w in words:
        if len(f"{cur} {w}".strip()) > width and cur:
            out.append(cur)
            cur = w
        else:
            cur = f"{cur} {w}".strip()
    if cur:
        out.append(cur)
    return "\n".join(out[:2]) if len(out) <= 2 else "\n".join(
        [" ".join(out[:len(out) // 2]), " ".join(out[len(out) // 2:])])


def ts(seconds):
    ms = int(round(seconds * 1000))
    h, ms = divmod(ms, 3600000)
    m, ms = divmod(ms, 60000)
    s, ms = divmod(ms, 1000)
    return f"{h:02d}:{m:02d}:{s:02d},{ms:03d}"


def main():
    for stream in (sys.stdout, sys.stderr):
        if hasattr(stream, "reconfigure"):
            stream.reconfigure(errors="backslashreplace")
    ap = argparse.ArgumentParser()
    ap.add_argument("manifest", type=Path)
    ap.add_argument("out", type=Path)
    ap.add_argument("--offsets", nargs="*", default=[],
                    help="SCENE=SECONDS start overrides; other scenes run "
                         "back to back in manifest order")
    ap.add_argument("--offsets-json", type=Path, default=None,
                    help="segments/offsets.json from assemble - selects scenes "
                         "in the finished cut and sets their start times")
    ap.add_argument("--max-chars", type=int, default=MAX_CHARS)
    ap.add_argument("--max-secs", type=float, default=MAX_SECS)
    args = ap.parse_args()
    if args.max_chars <= 0 or not math.isfinite(args.max_secs) or args.max_secs <= 0:
        ap.error("--max-chars and --max-secs must be finite and positive")

    manifest = json.loads(args.manifest.read_text(encoding="utf-8-sig"))
    overrides = {}
    if args.offsets_json:
        overrides.update({k: float(v) for k, v in
                          json.loads(args.offsets_json.read_text(encoding="utf-8-sig")).items()})
        # Assembly offsets identify the cut's scenes; manual offsets only retime them.
        manifest = [e for e in manifest if e["id"] in overrides]
    for spec in args.offsets:
        k, _, v = spec.partition("=")
        overrides[k] = float(v)

    cues, clock = [], 0.0
    for entry in manifest:
        start = overrides.get(entry["id"], clock)
        try:
            dur = float(entry["duration"])
            cues.extend(scene_cues(entry["text"], start, dur,
                                   args.max_chars, args.max_secs))
        except (ValueError, TypeError, OverflowError) as error:
            ap.error(f"scene {entry['id']}: {error}")
        clock = start + dur

    lines = []
    for i, (a, b, text) in enumerate(cues, 1):
        lines += [str(i), f"{ts(a)} --> {ts(b)}", text, ""]
    args.out.write_text("\n".join(lines), encoding="utf-8")
    end = cues[-1][1] if cues else 0.0
    print(f"{len(cues)} cues, ends at {ts(end)} -> {args.out}")
    return 0


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