#!/usr/bin/env python3
"""Baja el audio de un video corto (YouTube, Shorts, TikTok) y lo transcribe.

Pensado para shadowing: un clip de 30-90 s de habla real, con su texto exacto.

    python clip.py URL nombre-del-clip
    python clip.py URL nombre --inicio 00:01:20 --fin 00:02:10

Deja en 02-Areas/English/shadowing/:
    nombre.mp3   el audio
    nombre.srt   la transcripcion con timestamps
"""
import argparse
import os
import subprocess
import sys

YTDLP = "/data/users/julio/.conda/envs/tts/bin/yt-dlp"
FFMPEG = os.path.expanduser("~/bin/ffmpeg")
DESTINO = "/data/users/julio/Notes/02-Areas/English/shadowing"


def run(cmd, **kw):
    r = subprocess.run(cmd, capture_output=True, text=True, **kw)
    if r.returncode != 0:
        sys.exit((r.stderr or r.stdout).strip()[-1500:])
    return r


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("url")
    ap.add_argument("nombre", help="nombre base, sin extension")
    ap.add_argument("--inicio", help="recortar desde, HH:MM:SS")
    ap.add_argument("--fin", help="recortar hasta, HH:MM:SS")
    ap.add_argument("--destino", default=DESTINO)
    args = ap.parse_args()

    os.makedirs(args.destino, exist_ok=True)
    base = os.path.join(args.destino, args.nombre)
    crudo = base + ".crudo.mp3"
    final = base + ".mp3"

    print("bajando audio...")
    run([YTDLP, "-x", "--audio-format", "mp3", "--audio-quality", "5",
         "--ffmpeg-location", os.path.dirname(FFMPEG),
         "-o", crudo.replace(".mp3", ".%(ext)s"), args.url])

    if args.inicio or args.fin:
        cmd = [FFMPEG, "-y", "-loglevel", "error", "-i", crudo]
        if args.inicio:
            cmd += ["-ss", args.inicio]
        if args.fin:
            cmd += ["-to", args.fin]
        cmd += [final]
        run(cmd)
        os.remove(crudo)
    else:
        os.replace(crudo, final)

    print("transcribiendo...")
    from faster_whisper import WhisperModel
    m = WhisperModel("large-v3", device="cuda", device_index=0,
                     compute_type="float16")
    segs, info = m.transcribe(final, language="en", vad_filter=True, beam_size=5)

    def ts(s):
        ms = int(round(s * 1000))
        h, ms = divmod(ms, 3_600_000)
        mi, ms = divmod(ms, 60_000)
        se, ms = divmod(ms, 1000)
        return f"{h:02d}:{mi:02d}:{se:02d},{ms:03d}"

    with open(base + ".srt", "w", encoding="utf-8") as f:
        for n, s in enumerate(segs, 1):
            f.write(f"{n}\n{ts(s.start)} --> {ts(s.end)}\n{s.text.strip()}\n\n")

    print(f"\n{final}\n{base}.srt\n{info.duration:.0f}s de audio")


if __name__ == "__main__":
    main()
