#!/usr/bin/env python3
"""
녹음 파일을 받아서
  1) faster-whisper로 한국어 음성 인식
  2) pyannote.audio로 화자 분리
  3) 두 결과를 시간 기준으로 합쳐서
JSON으로 표준출력에 뽑아내는 스크립트.

전부 로컬에서 돌아가고 무료입니다 (모델 첫 다운로드에만 인터넷 필요).

사용법:
  python transcribe.py <오디오파일경로> [--min-speakers N] [--max-speakers N]

표준출력(stdout)에는 JSON만 출력하고, 진행 로그는 전부 stderr로 보냅니다.
Node 서버는 stdout만 읽으면 됩니다.
"""

import argparse
import json
import os
import subprocess
import sys
import tempfile

def log(msg):
    print(msg, file=sys.stderr, flush=True)


def to_wav16k_mono(src_path: str) -> str:
    """pyannote/whisper 둘 다 16kHz mono wav를 제일 안정적으로 처리하므로 ffmpeg로 미리 변환."""
    fd, out_path = tempfile.mkstemp(suffix=".wav")
    os.close(fd)
    cmd = [
        "ffmpeg", "-y", "-i", src_path,
        "-ac", "1", "-ar", "16000",
        out_path,
    ]
    result = subprocess.run(cmd, capture_output=True, text=True)
    if result.returncode != 0:
        raise RuntimeError(f"ffmpeg 변환 실패:\n{result.stderr}")
    return out_path


def run_whisper(wav_path: str, model_size: str, device: str, compute_type: str):
    from faster_whisper import WhisperModel

    log(f"[whisper] 모델 로딩 중: {model_size} ({device}/{compute_type})")
    model = WhisperModel(model_size, device=device, compute_type=compute_type)

    log("[whisper] 음성 인식 중...")
    segments, info = model.transcribe(
        wav_path,
        language="ko",
        beam_size=5,
        vad_filter=True,
        vad_parameters={"min_silence_duration_ms": 500},
    )

    whisper_segments = []
    for seg in segments:
        whisper_segments.append({
            "start": seg.start,
            "end": seg.end,
            "text": seg.text.strip(),
        })
        log(f"  [{seg.start:6.1f}-{seg.end:6.1f}] {seg.text.strip()}")

    return whisper_segments


def run_diarization(wav_path: str, hf_token: str, min_speakers, max_speakers):
    import torch
    from pyannote.audio import Pipeline

    log("[diarization] 화자분리 모델 로딩 중 (pyannote/speaker-diarization-3.1)...")
    pipeline = Pipeline.from_pretrained(
        "pyannote/speaker-diarization-3.1",
        use_auth_token=hf_token,
    )
    if torch.cuda.is_available():
        pipeline.to(torch.device("cuda"))
        log("[diarization] GPU 사용")
    else:
        log("[diarization] CPU 사용 (녹음이 길면 시간이 꽤 걸릴 수 있어요)")

    kwargs = {}
    if min_speakers:
        kwargs["min_speakers"] = min_speakers
    if max_speakers:
        kwargs["max_speakers"] = max_speakers

    log("[diarization] 화자분리 실행 중...")
    diarization = pipeline(wav_path, **kwargs)

    turns = []
    for turn, _, speaker in diarization.itertracks(yield_label=True):
        turns.append({"start": turn.start, "end": turn.end, "speaker": speaker})
    return turns


def assign_speaker(seg, turns):
    """whisper 세그먼트 구간과 가장 많이 겹치는 화자 턴을 찾는다."""
    best_speaker = None
    best_overlap = 0.0
    for t in turns:
        overlap = min(seg["end"], t["end"]) - max(seg["start"], t["start"])
        if overlap > best_overlap:
            best_overlap = overlap
            best_speaker = t["speaker"]
    return best_speaker or "화자?"


def speaker_display_name(label: str) -> str:
    # pyannote 라벨은 SPEAKER_00, SPEAKER_01 ... 형태 → A, B, C ...로 보기 좋게 변환
    if label.startswith("SPEAKER_"):
        try:
            idx = int(label.split("_")[-1])
            return chr(ord("A") + idx) if idx < 26 else label
        except ValueError:
            return label
    return label


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("audio_path")
    parser.add_argument("--min-speakers", type=int, default=None)
    parser.add_argument("--max-speakers", type=int, default=None)
    parser.add_argument("--model", default=os.environ.get("WHISPER_MODEL", "large-v3"))
    parser.add_argument("--device", default=os.environ.get("WHISPER_DEVICE", "cpu"))
    parser.add_argument("--compute-type", default=os.environ.get("WHISPER_COMPUTE_TYPE", "int8"))
    args = parser.parse_args()

    hf_token = os.environ.get("HF_TOKEN")
    if not hf_token:
        log("오류: HF_TOKEN 환경변수가 없습니다. .env에 HuggingFace 액세스 토큰을 넣어주세요.")
        sys.exit(1)

    wav_path = to_wav16k_mono(args.audio_path)
    try:
        whisper_segments = run_whisper(wav_path, args.model, args.device, args.compute_type)
        turns = run_diarization(wav_path, hf_token, args.min_speakers, args.max_speakers)

        merged = []
        for seg in whisper_segments:
            if not seg["text"]:
                continue
            raw_speaker = assign_speaker(seg, turns)
            merged.append({
                "start": round(seg["start"] * 1000),
                "end": round(seg["end"] * 1000),
                "speaker": speaker_display_name(raw_speaker),
                "text": seg["text"],
            })

        output = {
            "fullText": " ".join(s["text"] for s in merged),
            "speakers": sorted({s["speaker"] for s in merged}),
            "segments": merged,
        }
        print(json.dumps(output, ensure_ascii=False))
    finally:
        try:
            os.remove(wav_path)
        except OSError:
            pass


if __name__ == "__main__":
    main()
