#!/usr/bin/env python3
"""Transcribe audio using faster-whisper (CPU-optimized)."""
import sys
import time
from pathlib import Path
from faster_whisper import WhisperModel

BASE = Path("/opt/mia/workspace/clientes/px3lab/lancamento_photorf2/video_lancamento")
AUDIO = BASE / "audio.mp3"
TXT = BASE / "transcricao.txt"
SRT = BASE / "transcricao.srt"
JSON_OUT = BASE / "transcricao.json"

MODEL_SIZE = "small"  # medium é o sweet spot pra PT-BR em CPU
print(f"[+] Carregando modelo faster-whisper {MODEL_SIZE} (int8 CPU)...", flush=True)
t0 = time.time()
model = WhisperModel(MODEL_SIZE, device="cpu", compute_type="int8", cpu_threads=2)
print(f"[+] Modelo carregado em {time.time()-t0:.1f}s", flush=True)

print(f"[+] Transcrevendo {AUDIO}...", flush=True)
t0 = time.time()
segments, info = model.transcribe(
    str(AUDIO),
    language="pt",
    beam_size=5,
    vad_filter=True,
    vad_parameters=dict(min_silence_duration_ms=500),
    condition_on_previous_text=False,
)
print(f"[+] Detected language: {info.language} (prob={info.language_probability:.2f})", flush=True)
print(f"[+] Duration: {info.duration:.1f}s", flush=True)

# Coleta segments (streaming iterator)
segs = []
full_text_parts = []
srt_parts = []

def fmt_ts_srt(t):
    h = int(t // 3600); m = int((t % 3600) // 60); s = int(t % 60); ms = int((t - int(t)) * 1000)
    return f"{h:02d}:{m:02d}:{s:02d},{ms:03d}"

for i, seg in enumerate(segments, 1):
    segs.append({"id": i, "start": seg.start, "end": seg.end, "text": seg.text.strip()})
    full_text_parts.append(seg.text.strip())
    srt_parts.append(f"{i}\n{fmt_ts_srt(seg.start)} --> {fmt_ts_srt(seg.end)}\n{seg.text.strip()}\n")
    if i % 20 == 0:
        pct = seg.end / info.duration * 100 if info.duration else 0
        elapsed = time.time() - t0
        print(f"[{i:4d}] t={seg.end:7.1f}s ({pct:5.1f}%) elapsed={elapsed:6.1f}s", flush=True)

elapsed = time.time() - t0
print(f"[+] Transcrição concluída em {elapsed:.1f}s ({elapsed/60:.1f}min)", flush=True)
print(f"[+] Total segmentos: {len(segs)}", flush=True)

# Junta em parágrafos: quebra quando há pausa >1.2s entre segments
paragraphs = []
current = []
prev_end = 0.0
for seg in segs:
    if current and (seg["start"] - prev_end) > 1.2:
        paragraphs.append(" ".join(current))
        current = []
    current.append(seg["text"])
    prev_end = seg["end"]
if current:
    paragraphs.append(" ".join(current))

TXT.write_text("\n\n".join(paragraphs) + "\n", encoding="utf-8")
SRT.write_text("\n".join(srt_parts), encoding="utf-8")

import json
JSON_OUT.write_text(json.dumps({
    "language": info.language,
    "duration": info.duration,
    "segments": segs,
    "elapsed_seconds": elapsed,
    "model": MODEL_SIZE,
}, ensure_ascii=False, indent=2), encoding="utf-8")

print(f"[+] TXT: {TXT}", flush=True)
print(f"[+] SRT: {SRT}", flush=True)
print(f"[+] JSON: {JSON_OUT}", flush=True)

# Primeiro parágrafo pra validar
print("\n--- PRIMEIRO PARÁGRAFO ---")
print(paragraphs[0][:600] if paragraphs else "(vazio)")
