"""STEP 4 — 회의 녹음 → 화자가 식별된 회의 스크립트

파이프라인 3단 구성
-------------------
    ① faster-whisper (CTranslate2) : 음성 → 텍스트 + 문장 단위 타임스탬프
       · Whisper 를 4~5배 빠르게 돌리는 추론 엔진. WhisperX 가 내부에서 이걸 쓴다.
       · Whisper 원본 타임스탬프는 오차가 커서(±0.5초 이상) 화자 매칭에 그대로 쓰기 어렵다.

    ② WhisperX 강제정렬(forced alignment) : 문장 → **단어 단위** 타임스탬프
       · wav2vec2 음소 모델로 오디오와 텍스트를 다시 맞춘다(한국어: kresnik/wav2vec2-large-xlsr-korean).
       · 이 단계가 있어야 화자 경계를 단어 수준에서 정확히 자를 수 있다.

    ③ pyannote.audio 화자분리(diarization) : "언제 누가 말했는가"
       · speaker-diarization-3.1 = 발화 구간 검출 + 화자 임베딩 + 클러스터링.
       · 텍스트는 모른다. 시간축 위의 화자 구간만 낸다.

    → ②의 단어 타임스탬프와 ③의 화자 구간을 겹쳐서 단어마다 화자를 붙이면
      "화자가 식별된 회의 스크립트"가 완성된다.

출력 형식 (실무 표준)
---------------------
    .json  파이프라인 정본. 세그먼트별 start/end/speaker/text/words[]. 다음 단계(LLM) 입력.
    .vtt   WebVTT. 화자를 <v 이름> 태그로 표기. 플레이어/자막용.
    .srt   자막. 화자를 대괄호로 표기.
    .rttm  화자분리 평가 표준(NIST). 텍스트 없이 화자/시작/길이만.
    .md    사람이 읽는 회의 스크립트.
    .eval.json  정답 RTTM 이 있으면 DER 등 정확도 지표.

입출력
------
    입력  01-input/<회의명>.wav
    출력  03-transcript/<회의명>.{json,vtt,srt,rttm,md}

    정답(ground truth)이 있으면 화자 라벨을 실제 이름으로 자동 매핑하고 DER 을 계산한다.
    기본 위치는 00-audio-generation/work/reference/<회의명>.ref.json 이며,
    없으면 SPEAKER_00 같은 익명 라벨을 그대로 쓴다(실제 녹음을 다룰 때의 정상 동작).

사용법
------
    .venv\\Scripts\\python 02-STT-processing/diarize_transcribe.py --all
    .venv\\Scripts\\python 02-STT-processing/diarize_transcribe.py "01-input/xxx.wav" --num-speakers 6
    .venv\\Scripts\\python 02-STT-processing/diarize_transcribe.py --all --no-reference
"""
from __future__ import annotations

import argparse
import gc
import json
import os
import sys
import time
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from common import (DIR_INPUT, DIR_TRANSCRIPT, GEN_REFERENCE, hhmmss, mmss,
                    read_json, resolve_secret, write_json)

WHISPER_MODEL = "large-v3"
BATCH_SIZE = 8
LANG = "ko"


def patch_torch_load() -> None:
    """torch 2.6 + pyannote.audio 3.x 호환 처리.

    torch 2.6 부터 `torch.load(weights_only=...)` 기본값이 True 로 바뀌었다.
    pyannote 체크포인트에는 omegaconf 설정 객체가 pickle 로 들어 있어서
    weights_only=True 로는 로드가 실패한다(UnpicklingError).

    1순위: 필요한 클래스만 안전 목록(safe globals)에 등록한다.
    2순위: 그래도 걸리면 torch.load 기본값을 weights_only=False 로 되돌린다.
           HuggingFace 의 공식 pyannote 체크포인트만 로드하는 전제에서의 조치다.
           신뢰할 수 없는 체크포인트에는 쓰면 안 된다.
    """
    import torch
    try:
        import collections
        import typing

        from omegaconf.base import ContainerMetadata, Metadata
        from omegaconf.listconfig import ListConfig
        from omegaconf.nodes import AnyNode
        torch.serialization.add_safe_globals(
            [ListConfig, ContainerMetadata, Metadata, AnyNode,
             collections.defaultdict, dict, list, int, str, typing.Any])
    except Exception:
        pass

    _orig = torch.load

    def _load(*a, **kw):
        # lightning 이 weights_only=True 를 명시적으로 넘기므로 setdefault 로는 못 막는다.
        # 값을 강제로 덮어써야 pyannote 체크포인트가 로드된다.
        kw["weights_only"] = False
        return _orig(*a, **kw)

    torch.load = _load


def hf_token() -> str | None:
    for name in ("HF_TOKEN", "HUGGINGFACE_TOKEN", "HUGGING_FACE_HUB_TOKEN"):
        if v := resolve_secret(name):
            return v
    cached = Path.home() / ".cache" / "huggingface" / "token"
    if cached.exists():
        t = cached.read_text(encoding="utf-8").strip()
        if t:
            return t
    return None


def diarize_ecapa(audio, segments: list[dict], num_speakers: int | None, device: str):
    """pyannote 없이 돌리는 대체 화자분리기 (게이트 없는 모델만 사용).

    pyannote/speaker-diarization-3.1 은 HuggingFace 라이선스 동의가 필요한 gated 모델이라
    접근 권한이 없으면 아예 받을 수 없다. 그 경우를 위한 폴백이다.

        ① 강제정렬된 단어 타임스탬프로 실제 발화 구간을 구한다 (VAD 대용).
        ② 구간 위에 1.5초 창을 0.75초씩 밀며 SpeechBrain ECAPA-TDNN 화자 임베딩을 뽑는다.
        ③ 코사인 거리 + average linkage 계층 클러스터링으로 화자 수만큼 묶는다.
        ④ 인접한 같은 화자 창을 합쳐 화자 구간(turn)을 만든다.

    pyannote 3.1 보다 정확도는 낮다. 특히 겹쳐 말하는 구간을 처리하지 못한다.
    구조를 이해하는 용도이자, 권한이 없을 때의 우회로다.
    """
    import numpy as np
    import pandas as pd
    import torch
    from scipy.cluster.hierarchy import fcluster, linkage
    from scipy.spatial.distance import pdist
    from speechbrain.inference.speaker import EncoderClassifier

    SR, WIN, HOP, MIN_WIN = 16000, 1.5, 0.75, 0.45

    # ① 단어 타임스탬프 → 발화 구간
    words = [w for s in segments for w in (s.get("words") or [])
             if w.get("start") is not None and w.get("end") is not None]
    words.sort(key=lambda w: w["start"])
    regions: list[list[float]] = []
    for w in words:
        if regions and w["start"] - regions[-1][1] < 0.30:
            regions[-1][1] = max(regions[-1][1], w["end"])
        else:
            regions.append([w["start"], w["end"]])

    # ② 창 분할
    spans = []
    for a, b in regions:
        if b - a < MIN_WIN:
            continue
        if b - a <= WIN:
            spans.append((a, b))
            continue
        t = a
        while t + MIN_WIN < b:
            spans.append((t, min(t + WIN, b)))
            t += HOP
    if not spans:
        raise RuntimeError("발화 구간을 찾지 못했습니다.")

    enc = EncoderClassifier.from_hparams(
        source="speechbrain/spkrec-ecapa-voxceleb",
        savedir=str(Path.home() / ".cache" / "speechbrain" / "spkrec-ecapa-voxceleb"),
        run_opts={"device": device})

    embs = []
    B = 32
    for i in range(0, len(spans), B):
        chunk = spans[i:i + B]
        n = max(int((b - a) * SR) for a, b in chunk)
        batch = np.zeros((len(chunk), n), dtype=np.float32)
        lens = np.zeros(len(chunk), dtype=np.float32)
        for j, (a, b) in enumerate(chunk):
            seg = audio[int(a * SR):int(b * SR)]
            batch[j, :len(seg)] = seg
            lens[j] = len(seg) / n
        with torch.no_grad():
            e = enc.encode_batch(torch.from_numpy(batch).to(device),
                                 torch.from_numpy(lens).to(device))
        embs.append(e.squeeze(1).cpu().numpy())
    X = np.vstack(embs)
    X = X / (np.linalg.norm(X, axis=1, keepdims=True) + 1e-9)

    # ③ 클러스터링
    if num_speakers and num_speakers > 1 and len(X) > num_speakers:
        Z = linkage(pdist(X, metric="cosine"), method="average")
        labels = fcluster(Z, t=num_speakers, criterion="maxclust")
    else:
        labels = np.ones(len(X), dtype=int)

    # ④ 인접 창 병합 → 화자 구간
    turns = []
    for (a, b), lab in zip(spans, labels):
        spk = f"SPEAKER_{int(lab)-1:02d}"
        if turns and turns[-1]["speaker"] == spk and a - turns[-1]["end"] <= HOP + 0.05:
            turns[-1]["end"] = max(turns[-1]["end"], b)
        else:
            turns.append({"start": a, "end": b, "speaker": spk})
    return pd.DataFrame(turns), len(set(labels))


def free_gpu(*objs):
    for o in objs:
        del o
    gc.collect()
    try:
        import torch
        torch.cuda.empty_cache()
    except Exception:
        pass


# ── 화자 라벨 매핑 ────────────────────────────────────────────────────────────
def map_speakers(segments: list[dict], ref: dict | None) -> tuple[dict, dict]:
    """pyannote 의 SPEAKER_00/01/... 을 실제 참석자 이름에 대응시킨다.

    실무에서는 사람이 앞부분을 듣고 손으로 매핑하거나 화자 등록(enrollment)을 쓴다.
    여기서는 STEP 3 이 남긴 정답(ref)과 **시간 겹침이 최대인 화자**끼리 짝지어 자동 매핑한다.
    """
    if not ref:
        return {}, {}
    overlap: dict[tuple[str, str], float] = {}
    for s in segments:
        spk = s.get("speaker")
        if not spk:
            continue
        for r in ref["segments"]:
            ov = min(s["end"], r["end"]) - max(s["start"], r["start"])
            if ov > 0:
                overlap[(spk, r["speaker_id"])] = overlap.get((spk, r["speaker_id"]), 0.0) + ov

    mapping, used = {}, set()
    for (hyp, real), _ in sorted(overlap.items(), key=lambda kv: -kv[1]):
        if hyp in mapping or real in used:
            continue
        mapping[hyp] = real
        used.add(real)
    labels = {sp["speaker_id"]: sp["label"] for sp in ref["speakers"]}
    return mapping, labels


def evaluate(hyp_segments: list[dict], ref: dict) -> dict:
    """DER(Diarization Error Rate)과 화자 귀속 정확도."""
    from pyannote.core import Annotation, Segment
    from pyannote.metrics.diarization import DiarizationErrorRate

    reference, hypothesis = Annotation(), Annotation()
    for r in ref["segments"]:
        reference[Segment(r["start"], r["end"])] = r["speaker_id"]
    for i, s in enumerate(hyp_segments):
        if s.get("speaker"):
            hypothesis[Segment(s["start"], s["end"])] = s["speaker"]

    metric = DiarizationErrorRate()
    der = metric(reference, hypothesis)
    detail = metric(reference, hypothesis, detailed=True)
    return {
        "DER": round(float(der), 4),
        "missed_detection_sec": round(float(detail.get("missed detection", 0)), 2),
        "false_alarm_sec": round(float(detail.get("false alarm", 0)), 2),
        "confusion_sec": round(float(detail.get("confusion", 0)), 2),
        "total_ref_speech_sec": round(float(detail.get("total", 0)), 2),
        "ref_speakers": len(ref["speakers"]),
    }


# ── 출력 포맷 ────────────────────────────────────────────────────────────────
def write_outputs(stem: str, payload: dict) -> list[Path]:
    outs = []
    segs = payload["segments"]

    outs.append(write_json(DIR_TRANSCRIPT / f"{stem}.json", payload))

    vtt = ["WEBVTT", ""]
    for i, s in enumerate(segs, 1):
        vtt += [str(i),
                f"{hhmmss(s['start'], '.')} --> {hhmmss(s['end'], '.')}",
                f"<v {s['speaker_label']}>{s['text']}", ""]
    p = DIR_TRANSCRIPT / f"{stem}.vtt"; p.write_text("\n".join(vtt), encoding="utf-8"); outs.append(p)

    srt = []
    for i, s in enumerate(segs, 1):
        srt += [str(i),
                f"{hhmmss(s['start'], ',')} --> {hhmmss(s['end'], ',')}",
                f"[{s['speaker_label']}] {s['text']}", ""]
    p = DIR_TRANSCRIPT / f"{stem}.srt"; p.write_text("\n".join(srt), encoding="utf-8"); outs.append(p)

    uri = stem.replace(" ", "_")
    rttm = [f"SPEAKER {uri} 1 {s['start']:.3f} {s['end']-s['start']:.3f} <NA> <NA> "
            f"{s['speaker'].replace(' ', '')} <NA> <NA>" for s in segs if s.get("speaker")]
    p = DIR_TRANSCRIPT / f"{stem}.rttm"; p.write_text("\n".join(rttm) + "\n", encoding="utf-8"); outs.append(p)

    m = payload["meeting"]
    md = [f"# {m.get('title') or stem}", ""]
    if m.get("date"):
        md.append(f"- 일시: {m['date']}")
    if m.get("place"):
        md.append(f"- 장소: {m['place']}")
    md += [f"- 녹음 길이: {mmss(payload['duration_sec'])}",
           f"- 인식 화자 수: {payload['speaker_count']}명",
           f"- STT: faster-whisper {WHISPER_MODEL} + WhisperX 정렬 + pyannote 3.1 화자분리", ""]
    if payload.get("evaluation"):
        md.append(f"- 화자분리 정확도: DER {payload['evaluation']['DER']*100:.2f}%")
        md.append("")
    md.append("---")
    md.append("")
    last = None
    for s in segs:
        if s["speaker_label"] != last:
            md.append("")
            md.append(f"**[{mmss(s['start'])}] {s['speaker_label']}**")
            md.append("")
            last = s["speaker_label"]
        md.append(s["text"])
    p = DIR_TRANSCRIPT / f"{stem}.md"; p.write_text("\n".join(md) + "\n", encoding="utf-8"); outs.append(p)
    return outs


def merge_same_speaker(segs: list[dict], max_gap: float = 1.6) -> list[dict]:
    """연속된 같은 화자의 세그먼트를 하나의 발화로 합친다."""
    out: list[dict] = []
    for s in segs:
        if out and out[-1]["speaker"] == s["speaker"] and s["start"] - out[-1]["end"] <= max_gap:
            prev = out[-1]
            prev["end"] = s["end"]
            prev["text"] = (prev["text"].rstrip() + " " + s["text"].lstrip()).strip()
            prev["words"].extend(s.get("words", []))
        else:
            out.append({**s, "words": list(s.get("words", []))})
    return out


def run(wav: Path, num_speakers: int | None, device: str, diarizer: str = "pyannote",
        reference_dir: Path | None = None) -> dict:
    import whisperx
    from whisperx.diarize import DiarizationPipeline, assign_word_speakers

    stem = wav.stem
    ref_path = (reference_dir / f"{stem}.ref.json") if reference_dir else None
    ref = read_json(ref_path) if (ref_path and ref_path.exists()) else None
    if num_speakers is None and ref:
        num_speakers = len(ref["speakers"])

    print(f"\n{'='*78}\n[{stem}]\n{'='*78}")
    audio = whisperx.load_audio(str(wav))
    dur = len(audio) / 16000
    print(f"  녹음 길이 {mmss(dur)} / 예상 화자 {num_speakers or '자동'}명 / device={device}")
    timing = {}

    # ① faster-whisper 전사
    t0 = time.time()
    compute = "float16" if device == "cuda" else "int8"
    model = whisperx.load_model(WHISPER_MODEL, device, compute_type=compute, language=LANG)
    result = model.transcribe(audio, batch_size=BATCH_SIZE, language=LANG)
    timing["transcribe"] = round(time.time() - t0, 1)
    print(f"  ① 전사        {len(result['segments'])}개 세그먼트  ({timing['transcribe']}초)")
    free_gpu(model)

    # ② WhisperX 강제정렬 → 단어 타임스탬프
    t0 = time.time()
    model_a, metadata = whisperx.load_align_model(language_code=LANG, device=device)
    result = whisperx.align(result["segments"], model_a, metadata, audio, device,
                            return_char_alignments=False)
    timing["align"] = round(time.time() - t0, 1)
    nwords = sum(len(s.get("words", [])) for s in result["segments"])
    print(f"  ② 강제정렬    단어 {nwords}개에 타임스탬프  ({timing['align']}초)")
    free_gpu(model_a)

    # ③ 화자분리
    t0 = time.time()
    engine = diarizer
    if diarizer == "pyannote":
        tok = hf_token()
        if not tok:
            raise SystemExit("HuggingFace 토큰이 필요합니다. HF_TOKEN 환경변수 또는 `huggingface-cli login`.")
        try:
            diar = DiarizationPipeline(use_auth_token=tok, device=device)
            if diar.model is None:
                raise RuntimeError("pyannote 파이프라인을 받지 못했습니다(gated repo 접근 거부).")
            diar_segments = diar(audio, num_speakers=num_speakers) if num_speakers else diar(audio)
            free_gpu(diar)
        except Exception as e:
            print(f"  [!] pyannote 사용 불가 → ECAPA 폴백으로 전환합니다. ({type(e).__name__})")
            print("      pyannote/speaker-diarization-3.1 과 segmentation-3.0 의 라이선스에 동의하면")
            print("      https://huggingface.co/settings/gated-repos 에서 ACCEPTED 로 바뀝니다.")
            engine = "ecapa"
    if engine == "ecapa":
        diar_segments, _ = diarize_ecapa(audio, result["segments"], num_speakers, device)

    timing["diarize"] = round(time.time() - t0, 1)
    found = sorted(diar_segments["speaker"].unique())
    print(f"  ③ 화자분리    {len(found)}명 검출 {found}  [{engine}]  ({timing['diarize']}초)")

    # ④ 단어에 화자 할당
    result = assign_word_speakers(diar_segments, result)

    segs = []
    for s in result["segments"]:
        text = (s.get("text") or "").strip()
        if not text:
            continue
        segs.append({
            "start": round(float(s["start"]), 3),
            "end": round(float(s["end"]), 3),
            "speaker": s.get("speaker") or "UNKNOWN",
            "text": text,
            "words": [{"word": w.get("word"), "start": w.get("start"), "end": w.get("end"),
                       "score": w.get("score"), "speaker": w.get("speaker")}
                      for w in (s.get("words") or []) if w.get("start") is not None],
        })
    segs = merge_same_speaker(segs)

    # ⑤ 화자 라벨을 실제 참석자 이름으로
    mapping, labels = map_speakers(segs, ref)
    for i, s in enumerate(segs, 1):
        s["id"] = i
        real = mapping.get(s["speaker"])
        s["speaker_id"] = real
        s["speaker_label"] = labels.get(real, s["speaker"]) if real else s["speaker"]

    payload = {
        "audio": wav.name,
        "duration_sec": round(dur, 2),
        "language": LANG,
        "pipeline": {
            "asr": f"faster-whisper {WHISPER_MODEL} (ctranslate2)",
            "alignment": "whisperx forced alignment / kresnik/wav2vec2-large-xlsr-korean",
            "diarization": ("pyannote/speaker-diarization-3.1" if engine == "pyannote"
                            else "speechbrain/spkrec-ecapa-voxceleb + agglomerative clustering (폴백)"),
            "device": device, "batch_size": BATCH_SIZE,
        },
        "timing_sec": timing,
        "meeting": (ref or {}).get("meeting", {"title": stem}),
        "speaker_count": len({s["speaker"] for s in segs}),
        "speaker_mapping": mapping,
        "segments": segs,
    }
    if ref:
        payload["evaluation"] = evaluate(segs, ref)
        ev = payload["evaluation"]
        print(f"  ④ 정확도      DER {ev['DER']*100:.2f}%  "
              f"(누락 {ev['missed_detection_sec']}s / 오검출 {ev['false_alarm_sec']}s / 화자혼동 {ev['confusion_sec']}s)")
        print(f"     화자 매핑   " + ", ".join(f"{h}→{labels.get(r, r)}" for h, r in sorted(mapping.items())))

    DIR_TRANSCRIPT.mkdir(parents=True, exist_ok=True)
    for p in write_outputs(stem, payload):
        print(f"  → {p.name}")
    print(f"  총 처리시간 {sum(timing.values()):.0f}초 (실시간 대비 {dur/max(sum(timing.values()),1):.1f}배속)")
    return payload


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("audio", nargs="?")
    ap.add_argument("--all", action="store_true")
    ap.add_argument("--num-speakers", type=int, default=None)
    ap.add_argument("--device", default=None, help="cuda | cpu (기본: 가능하면 cuda)")
    ap.add_argument("--diarizer", choices=["pyannote", "ecapa"], default="pyannote",
                    help="pyannote=권장(gated 모델). ecapa=게이트 없는 폴백.")
    ap.add_argument("--reference-dir", default=str(GEN_REFERENCE),
                    help="화자분리 정답(.ref.json)이 있는 폴더. 있으면 실제 이름 매핑 + DER 계산.")
    ap.add_argument("--no-reference", action="store_true",
                    help="정답을 무시하고 SPEAKER_00 같은 익명 라벨을 그대로 쓴다(실제 녹음 시나리오).")
    args = ap.parse_args()

    patch_torch_load()

    if args.device:
        device = args.device
    else:
        import torch
        device = "cuda" if torch.cuda.is_available() else "cpu"

    ref_dir = None if args.no_reference else Path(args.reference_dir)

    targets = ([Path(args.audio)] if args.audio and not args.all
               else sorted(p for p in DIR_INPUT.glob("*.wav")))
    if not targets:
        raise SystemExit(f"{DIR_INPUT} 에 wav 가 없습니다.")
    for t in targets:
        run(t, args.num_speakers, device, args.diarizer, ref_dir)


if __name__ == "__main__":
    main()
