#!/usr/bin/env python3
"""Transcribe audio y video a .srt con faster-whisper.

Saca .srt (no .txt) a proposito: el timestamp es lo que permite preguntar
"minuto 12:34 no entendi" y que la skill ingles-vocab lo resuelva.

Decisiones medidas sobre el material de YouTalkTV (ver libros.md):

- BatchedInferencePipeline con batch_size 16, unas 3 veces mas rapido que
  el modo secuencial (71x contra 21x tiempo real).
- El idioma se fuerza por carpeta, no por segmento. El nivel ya lo predice:
  Intermedio y Avanzado son ingles, el resto mezcla espanol e ingles.
  Forzarlo es mas rapido y mas preciso que multilingual=True.
- initial_prompt SIEMPRE. Sin el, el ingles sale sin puntuacion ni
  mayusculas, un muro de texto en minusculas. Con el, la puntuacion sube
  de 89 a 146 marcas en la misma clase.
- beam_size 5. Con 1 va marginalmente mas rapido pero pierde calidad.

    python transcribir.py ENTRADA SALIDA [--solo 03-Avanzado 02-Intermedio]
"""
import argparse
import os
import sys
import time

MEDIA = {".mp3", ".wav", ".m4a", ".ogg", ".wma", ".aac", ".flac",
         ".mp4", ".avi", ".mkv", ".mov", ".wmv", ".flv"}

# PyAV abre los mp4 directamente, asi que no hay que convertirlos a mp3.

PROMPT_EN = ("Hello everybody, and welcome back to the class. Today, we are going to "
             "study grammar, sentences and vocabulary. Let's begin, shall we?")
PROMPT_MIX = ("Hola a todos, bienvenidos de nuevo a la clase. Hoy vamos a ver gramatica, "
              "frases y vocabulario. Let's begin, shall we? Empezamos.")

# Carpetas cuyo contenido es ingles casi puro. El resto va en modo mixto.
SOLO_INGLES = ("02-Intermedio", "03-Avanzado")


def config_idioma(rel):
    """Devuelve (kwargs de idioma, etiqueta) segun la carpeta de nivel."""
    raiz = rel.split(os.sep)[0]
    if raiz in SOLO_INGLES:
        return {"language": "en", "initial_prompt": PROMPT_EN}, "en"
    return {"multilingual": True, "initial_prompt": PROMPT_MIX}, "mixto"


def ts(seconds):
    """Segundos a HH:MM:SS,mmm del formato SRT."""
    ms = int(round(seconds * 1000))
    h, ms = divmod(ms, 3_600_000)
    m, ms = divmod(ms, 60_000)
    s, ms = divmod(ms, 1000)
    return f"{h:02d}:{m:02d}:{s:02d},{ms:03d}"


def encontrar(raiz, solo):
    """Lista media bajo raiz, saltando .origen y ocultos. Ordena por prioridad."""
    hallados = []
    for dirpath, dirnames, files in os.walk(raiz):
        dirnames[:] = [d for d in dirnames if not d.startswith(".")]
        for fn in files:
            if os.path.splitext(fn)[1].lower() in MEDIA:
                hallados.append(os.path.relpath(os.path.join(dirpath, fn), raiz))

    if solo:
        hallados = [r for r in hallados if r.split(os.sep)[0] in solo]
        orden = {n: i for i, n in enumerate(solo)}
        hallados.sort(key=lambda r: (orden[r.split(os.sep)[0]], r))
    else:
        hallados.sort()
    return hallados


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("entrada")
    ap.add_argument("salida")
    ap.add_argument("--model", default="large-v3")
    ap.add_argument("--gpu", type=int, default=0)
    ap.add_argument("--batch", type=int, default=16)
    ap.add_argument("--beam", type=int, default=5)
    ap.add_argument("--solo", nargs="*", default=None,
                    help="carpetas de primer nivel, en orden de prioridad")
    ap.add_argument("--cpu", action="store_true")
    args = ap.parse_args()

    from faster_whisper import WhisperModel, BatchedInferencePipeline

    device = "cpu" if args.cpu else "cuda"
    kw = {"compute_type": "int8" if args.cpu else "float16"}
    if device == "cuda":
        kw["device_index"] = args.gpu

    print(f"Cargando {args.model} en {device}:{args.gpu} ...", flush=True)
    pipe = BatchedInferencePipeline(
        model=WhisperModel(args.model, device=device, **kw))

    archivos = encontrar(args.entrada, args.solo)
    print(f"{len(archivos)} archivos de media\n", flush=True)

    hechos = saltados = fallidos = 0
    audio_total = 0.0
    t0 = time.time()

    for i, rel in enumerate(archivos, 1):
        destino = os.path.join(args.salida, os.path.splitext(rel)[0] + ".srt")
        if os.path.exists(destino):
            saltados += 1
            continue

        os.makedirs(os.path.dirname(destino), exist_ok=True)
        idioma_kw, etiqueta = config_idioma(rel)
        print(f"[{i}/{len(archivos)}] [{etiqueta}] {rel}", flush=True)

        try:
            t1 = time.time()
            segmentos, info = pipe.transcribe(
                os.path.join(args.entrada, rel),
                batch_size=args.batch,
                beam_size=args.beam,
                vad_filter=True,
                **idioma_kw,
            )
            # .tmp y renombrar, para que un corte no deje un srt a medias
            tmp = destino + ".tmp"
            with open(tmp, "w", encoding="utf-8") as f:
                for n, seg in enumerate(segmentos, 1):
                    f.write(f"{n}\n{ts(seg.start)} --> {ts(seg.end)}\n"
                            f"{seg.text.strip()}\n\n")
            os.replace(tmp, destino)

            dt = time.time() - t1
            audio_total += info.duration
            print(f"    {info.duration/60:5.1f} min de audio en {dt:5.1f}s "
                  f"({info.duration/dt:.0f}x)", flush=True)
            hechos += 1
        except Exception as exc:
            print(f"    FALLO: {type(exc).__name__}: {exc}",
                  file=sys.stderr, flush=True)
            fallidos += 1

    mins = (time.time() - t0) / 60
    print(f"\nHechos {hechos}, saltados {saltados}, fallidos {fallidos}")
    print(f"{audio_total/3600:.1f} h de audio en {mins:.1f} min")


if __name__ == "__main__":
    main()
