#!/usr/bin/env python3
"""Mac-native voice extraction — Single idol, MLX + MPS 스택.

관측 가능한 로컬 파이프라인. 각 단계 실시간 로그 출력.

Pipeline:
  1 fetch      yt-dlp → wav 16k mono (per URL)
  2 vocals     demucs-mlx htdemucs CLI (batch N 트랙)
  3 diarize    pyannote/speaker-diarization-3.1 on MPS (per URL)
  4 select     auto_dominant (default) 또는 --speaker SPEAKER_XX
  5 match      pyannote WeSpeaker-ResNet34 on MPS → cosine ≥ THRESHOLD
  6 filter     silero-vad voice_ratio + SNR + < 30s
  7 concat     cosine 상위 PROMPT_SECONDS prompt wav
  8 transcribe docksvc whisper-cuda-spark (HTTP)

Usage:
  extract-voice-mac.py --idol jennie --urls URL1 URL2 URL3 [--speaker SPEAKER_XX]
  extract-voice-mac.py --idol jennie --list-speakers --urls URL1

Output:
  ~/ai/experiments/voice-extract-mac/{idol}/data/filtered/*.wav
  ~/ai/experiments/voice-extract-mac/{idol}/data/prompt.wav
  ~/ai/experiments/voice-extract-mac/{idol}/data/prompt.txt
  ~/ai/experiments/voice-extract-mac/{idol}/data/manifest.json
"""
from __future__ import annotations

import argparse
import json
import os
import shutil
import subprocess
import sys
import time
from pathlib import Path
from urllib.parse import urlparse, parse_qs

SR = 16000
# CosyVoice prompt 길이 규약 (CEO 2026-04-21): 10-15초. 10초 미만 = 화자 정보 부족, 15초 초과 = 의미 없음.
PROMPT_SECONDS = 12.0
SEG_MIN_S = 2.0   # 개별 wav 하한 (너무 짧으면 prompt 후보로 부적합)
SEG_MAX_S = 15.0  # 개별 wav 상한 (CosyVoice best_segment가 이 범위 내에서 뽑혀야 함)
COSINE_THRESHOLD = 0.55
WHISPER_URL = "http://edgexpert-e1a2.tailc38d4e.ts.net:8720/v1/audio/transcriptions"


def ensure_hf_token() -> None:
    tok_path = Path.home() / ".cache" / "huggingface" / "token"
    if tok_path.exists() and "HF_TOKEN" not in os.environ:
        os.environ["HF_TOKEN"] = tok_path.read_text().strip()


def data_root(idol: str) -> Path:
    return Path.home() / "ai" / "experiments" / "voice-extract-mac" / idol / "data"


def log(msg: str) -> None:
    ts = time.strftime("%H:%M:%S")
    print(f"[{ts}] {msg}", flush=True)


def video_id(url: str) -> str:
    parsed = urlparse(url)
    if parsed.hostname in ("youtu.be",):
        return parsed.path.lstrip("/")
    qs = parse_qs(parsed.query)
    if "v" in qs:
        return qs["v"][0]
    return parsed.path.split("/")[-1]


def run(cmd, **kw) -> None:
    subprocess.run(cmd, check=True, **kw)


def parse_rttm(path: Path):
    for line in path.read_text().splitlines():
        parts = line.split()
        if not parts or parts[0] != "SPEAKER":
            continue
        yield {"start": float(parts[3]), "dur": float(parts[4]), "spk": parts[7]}


# ─── Phase 1: fetch ───────────────────────────────────────────────────────────
def phase_fetch(root: Path, urls: list[str]) -> None:
    raw_dir = root / "raw"
    raw_dir.mkdir(parents=True, exist_ok=True)
    for url in urls:
        vid = video_id(url)
        out = raw_dir / f"{vid}.wav"
        if out.exists():
            log(f"  fetch [skip] {vid}")
            continue
        log(f"  fetch [yt-dlp] {vid} ...")
        raw_tmpl = raw_dir / f"{vid}_raw.%(ext)s"
        t0 = time.time()
        run(["yt-dlp", "-x", "--audio-format", "wav", "-o", str(raw_tmpl), url],
            stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
        raw_wav = raw_dir / f"{vid}_raw.wav"
        run(["ffmpeg", "-y", "-i", str(raw_wav), "-ar", str(SR), "-ac", "1", str(out)],
            stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
        raw_wav.unlink()
        dur = out.stat().st_size / 2 / SR
        log(f"         done ({dur:.0f}s, {time.time()-t0:.1f}s elapsed)")


# ─── Phase 2: vocals (demucs-mlx CLI, batch) ──────────────────────────────────
def phase_vocals(root: Path, urls: list[str]) -> None:
    voc_dir = root / "vocals"
    voc_dir.mkdir(parents=True, exist_ok=True)

    todo = []
    for url in urls:
        vid = video_id(url)
        src = root / "raw" / f"{vid}.wav"
        out = voc_dir / f"{vid}.wav"
        if out.exists():
            log(f"  vocals [skip] {vid}")
            continue
        todo.append((vid, src, out))

    if not todo:
        return

    import shutil as _sh
    demucs_bin = _sh.which("demucs-mlx") or str(Path(sys.executable).parent / "demucs-mlx")
    tmp_out = voc_dir / "_demucs_tmp"
    if tmp_out.exists():
        shutil.rmtree(tmp_out)
    tmp_out.mkdir(parents=True)

    # 한 번에 N 트랙 처리 (모델 1회 로드)
    track_paths = [str(src) for _, src, _ in todo]
    log(f"  vocals [demucs-mlx] batch separating {len(track_paths)} track(s) ...")
    t0 = time.time()
    try:
        r = subprocess.run(
            [demucs_bin, "-n", "htdemucs", "-o", str(tmp_out), *track_paths],
            capture_output=True, text=True,
        )
    except Exception as e:
        sys.exit(f"ERROR: demucs-mlx 실행 실패: {e}")
    if r.returncode != 0:
        sys.exit(f"ERROR: demucs-mlx exit={r.returncode}\nstderr: {r.stderr[-1000:]}")
    log(f"             batch done ({time.time()-t0:.1f}s elapsed)")

    for vid, src, out in todo:
        # demucs-mlx 출력: {tmp_out}/{src_stem}/vocals.wav (htdemucs 중간 dir 없음)
        produced = tmp_out / src.stem / "vocals.wav"
        if not produced.exists():
            cand = list((tmp_out / src.stem).glob("vocals.*"))
            if not cand:
                sys.exit(f"ERROR: demucs-mlx vocals 출력 못 찾음 ({vid}) — checked {produced}")
            produced = cand[0]
        run(["ffmpeg", "-y", "-i", str(produced), "-ar", str(SR), "-ac", "1", str(out)],
            stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
        log(f"             [{vid}] -> {out.name}")
    shutil.rmtree(tmp_out, ignore_errors=True)


# ─── Phase 3: diarize ─────────────────────────────────────────────────────────
def phase_diarize(root: Path, urls: list[str], device: str = "mps") -> None:
    from pyannote.audio import Pipeline
    import torch

    dia_dir = root / "diarized"
    dia_dir.mkdir(parents=True, exist_ok=True)

    todo = []
    for url in urls:
        vid = video_id(url)
        rttm = dia_dir / f"{vid}.rttm"
        if rttm.exists() and rttm.stat().st_size > 0:
            log(f"  diarize [skip] {vid}")
            continue
        todo.append((vid, rttm))

    if not todo:
        return

    log(f"  diarize [pyannote] loading speaker-diarization-3.1 on {device} ...")
    t0 = time.time()
    pipeline = Pipeline.from_pretrained("pyannote/speaker-diarization-3.1", token=True)
    pipeline.to(torch.device(device))
    log(f"          loaded ({time.time()-t0:.1f}s)")

    for vid, rttm in todo:
        src = root / "vocals" / f"{vid}.wav"
        log(f"  diarize [run] {vid} ...")
        t0 = time.time()
        diarization = pipeline(str(src))
        annot = getattr(diarization, "speaker_diarization", diarization)
        with rttm.open("w") as f:
            annot.write_rttm(f)
        segs = list(parse_rttm(rttm))
        spks = set(s["spk"] for s in segs)
        log(f"          done ({time.time()-t0:.1f}s) — {len(segs)} segments, {len(spks)} speakers: {sorted(spks)}")


# ─── embedder (WeSpeaker ResNet34 MLX-wrapped via pyannote) ───────────────────
def _embedder(device: str = "mps"):
    from pyannote.audio import Model
    import torch
    m = Model.from_pretrained("pyannote/wespeaker-voxceleb-resnet34-LM")
    m = m.to(torch.device(device))
    return m


def _embed(model, wav_np, device="mps"):
    import torch
    t = torch.from_numpy(wav_np).float().unsqueeze(0).unsqueeze(0).to(device)
    with torch.no_grad():
        emb = model(t).squeeze().detach().cpu().numpy()
    n = float((emb ** 2).sum() ** 0.5)
    return emb / max(n, 1e-8)


def select_speaker_auto(root: Path, bootstrap_vid: str) -> str:
    rttm = root / "diarized" / f"{bootstrap_vid}.rttm"
    totals: dict[str, float] = {}
    for seg in parse_rttm(rttm):
        totals[seg["spk"]] = totals.get(seg["spk"], 0) + seg["dur"]
    if not totals:
        sys.exit("ERROR: diarization RTTM 비어있음")
    dominant = max(totals.items(), key=lambda kv: kv[1])
    log(f"  select auto_dominant: {dominant[0]} (total {dominant[1]:.1f}s)")
    for spk, total in sorted(totals.items(), key=lambda kv: -kv[1]):
        log(f"           {spk}: {total:.1f}s")
    return dominant[0]


def phase_match(root: Path, urls: list[str], target_spk: str, device: str = "mps") -> dict:
    import numpy as np
    import soundfile as sf

    log(f"  match [load WeSpeaker] on {device} ...")
    t0 = time.time()
    model = _embedder(device)
    log(f"        loaded ({time.time()-t0:.1f}s)")

    bootstrap_vid = video_id(urls[0])
    src_wav, sr = sf.read(root / "vocals" / f"{bootstrap_vid}.wav")
    if src_wav.ndim > 1:
        src_wav = src_wav.mean(axis=1)
    log(f"  match [centroid] from {bootstrap_vid} {target_spk} ...")
    embeddings = []
    for seg in parse_rttm(root / "diarized" / f"{bootstrap_vid}.rttm"):
        if seg["spk"] != target_spk or seg["dur"] < 1.0:
            continue
        s, e = int(seg["start"] * sr), int((seg["start"] + seg["dur"]) * sr)
        embeddings.append(_embed(model, src_wav[s:e], device))
    if not embeddings:
        sys.exit(f"ERROR: '{target_spk}' segments 0건")
    centroid = np.mean(embeddings, axis=0)
    centroid /= np.linalg.norm(centroid)
    np.save(root / "centroid.npy", centroid)
    log(f"        centroid: {len(embeddings)} segments")

    matched_count = {}
    matched_root = root / "matched"
    for url in urls:
        vid = video_id(url)
        out_dir = matched_root / vid
        out_dir.mkdir(parents=True, exist_ok=True)
        wav, sr = sf.read(root / "vocals" / f"{vid}.wav")
        if wav.ndim > 1:
            wav = wav.mean(axis=1)
        n = 0
        t0 = time.time()
        segs = list(parse_rttm(root / "diarized" / f"{vid}.rttm"))
        for seg in segs:
            if seg["dur"] < 1.0:
                continue
            s, e = int(seg["start"] * sr), int((seg["start"] + seg["dur"]) * sr)
            clip = wav[s:e]
            emb = _embed(model, clip, device)
            cos = float(np.dot(centroid, emb))
            if cos >= COSINE_THRESHOLD:
                out = out_dir / f"{int(seg['start']*1000):08d}_{int(seg['dur']*1000):05d}_{cos:.3f}.wav"
                sf.write(out, clip, sr)
                n += 1
        matched_count[vid] = n
        log(f"  match [{vid}] {n}/{len(segs)} segments ≥ {COSINE_THRESHOLD} ({time.time()-t0:.1f}s)")

    return matched_count


def _snr(audio, vad_ts) -> float:
    import numpy as np
    voiced = np.zeros(len(audio), dtype=bool)
    for t in vad_ts:
        voiced[t["start"]:t["end"]] = True
    sig = audio[voiced]
    noise = audio[~voiced]
    if len(noise) == 0 or len(sig) == 0:
        return 60.0
    sig_rms = float(np.sqrt((sig.astype("float64") ** 2).mean()))
    noise_rms = float(np.sqrt((noise.astype("float64") ** 2).mean()))
    if noise_rms < 1e-8:
        return 60.0
    return 20.0 * float(np.log10(sig_rms / noise_rms))


# ─── Phase 5.5: stitch ─────────────────────────────────────────────────────────
# 연속된 matched segments 를 이어 10-15초 clip 생성.
# 짧은 conversational segments (Weverse live 등) 는 개별로 10초 못 채우는 경우가 대부분이므로 stitch 필수.
# matched/{vid}/*.wav 의 timestamps (파일명 {start_ms}_{dur_ms}_{cos}) 로 인접성 판단,
# vocals/{vid}.wav 에서 실제 구간 추출 (prosody·pause 자연스럽게 유지).
def phase_stitch(root: Path, urls: list[str],
                 min_dur: float = 10.0, max_dur: float = 15.0,
                 gap_tolerance: float = 0.5) -> int:
    import soundfile as sf

    stitched = root / "stitched"
    if stitched.exists():
        for f in stitched.glob("*.wav"):
            f.unlink()
    stitched.mkdir(parents=True, exist_ok=True)

    total = 0
    for url in urls:
        vid = video_id(url)
        matched_dir = root / "matched" / vid
        vocal = root / "vocals" / f"{vid}.wav"
        if not matched_dir.exists() or not vocal.exists():
            continue
        audio, sr = sf.read(vocal)
        if audio.ndim > 1:
            audio = audio.mean(axis=1)

        events = []
        for p in matched_dir.glob("*.wav"):
            try:
                parts = p.stem.split("_")
                events.append({
                    "start": int(parts[0]) / 1000.0,
                    "dur": int(parts[1]) / 1000.0,
                    "cos": float(parts[2]),
                })
            except (ValueError, IndexError):
                continue
        events.sort(key=lambda e: e["start"])

        i = 0
        count = 0
        while i < len(events):
            start = events[i]["start"]
            end = events[i]["start"] + events[i]["dur"]
            cos_sum = events[i]["cos"] * events[i]["dur"]
            cos_dur = events[i]["dur"]
            j = i + 1
            while j < len(events):
                gap = events[j]["start"] - end
                new_end = events[j]["start"] + events[j]["dur"]
                if gap > gap_tolerance or (new_end - start) > max_dur:
                    break
                end = new_end
                cos_sum += events[j]["cos"] * events[j]["dur"]
                cos_dur += events[j]["dur"]
                j += 1

            dur = end - start
            if min_dur <= dur <= max_dur:
                avg_cos = cos_sum / cos_dur
                clip = audio[int(start * sr):int(end * sr)]
                fname = f"{vid}_{int(start*1000):08d}_{int(dur*1000):05d}_{avg_cos:.3f}.wav"
                sf.write(stitched / fname, clip, sr)
                count += 1
            i = j if j > i else i + 1

        log(f"  stitch [{vid}]: {count} clips ({min_dur}-{max_dur}s)")
        total += count
    log(f"  stitched total: {total}")
    return total


def phase_filter(root: Path, urls: list[str]) -> int:
    """stitched/*.wav 대상 VAD + SNR 품질 필터 → filtered/"""
    import torch
    import soundfile as sf
    from silero_vad import load_silero_vad, get_speech_timestamps

    stitched = root / "stitched"
    filt = root / "filtered"
    if filt.exists():
        for old in filt.glob("*.wav"):
            old.unlink()
    filt.mkdir(parents=True, exist_ok=True)

    log("  filter [silero-vad] loading ...")
    vad_model = load_silero_vad()

    kept = 0
    rejected = {"no_voice": 0, "voice_dur": 0, "voice_ratio": 0, "snr": 0}
    for wav_path in sorted(stitched.glob("*.wav")):
        audio, sr = sf.read(wav_path)
        if audio.ndim > 1:
            audio = audio.mean(axis=1)
        dur = len(audio) / sr
        ts_list = get_speech_timestamps(torch.from_numpy(audio).float(),
                                        vad_model, sampling_rate=sr)
        ts = [{"start": t["start"], "end": t["end"]} for t in ts_list]
        if not ts:
            rejected["no_voice"] += 1; continue
        voice_dur = sum(t["end"] - t["start"] for t in ts) / sr
        if voice_dur < 1.5:
            rejected["voice_dur"] += 1; continue
        voice_ratio = voice_dur / dur
        if voice_ratio < 0.6:
            rejected["voice_ratio"] += 1; continue
        if _snr(audio, ts) < 10.0:
            rejected["snr"] += 1; continue
        sf.write(filt / wav_path.name, audio, sr)
        kept += 1
    log(f"  filter kept: {kept} / {len(list(stitched.glob('*.wav')))} stitched clips")
    log(f"         rejected: {rejected}")
    return kept


def phase_concat(root: Path) -> tuple[float, int]:
    """filtered 중 cosine 최고 단일 clip 을 prompt.wav 로 사용 (이미 10-15s 로 stitched)."""
    import soundfile as sf

    filt = root / "filtered"
    candidates = []
    for wav_path in filt.glob("*.wav"):
        try:
            cos = float(wav_path.stem.split("_")[-1])
        except ValueError:
            continue
        candidates.append((cos, wav_path))
    candidates.sort(reverse=True)
    if not candidates:
        sys.exit("ERROR: filter 통과 후보 0건")

    cos, top = candidates[0]
    audio, sr = sf.read(top)
    if audio.ndim > 1:
        audio = audio.mean(axis=1)
    pwav = root / "prompt.wav"
    sf.write(pwav, audio, sr)
    dur = len(audio) / sr
    log(f"  prompt {pwav.name}: {dur:.2f}s ({top.name}, cos={cos:.3f})")
    return dur, 1


def phase_transcribe(root: Path) -> str:
    import requests
    pwav = root / "prompt.wav"
    log(f"  transcribe [docksvc whisper] {pwav.name} ...")
    t0 = time.time()
    with pwav.open("rb") as f:
        r = requests.post(WHISPER_URL,
                          files={"file": f},
                          data={"model": "whisper-speech", "language": "ko"},
                          timeout=300)
    r.raise_for_status()
    text = r.json().get("text", "").strip()
    ptxt = root / "prompt.txt"
    ptxt.write_text(text + "\n")
    log(f"             done ({time.time()-t0:.1f}s) — {len(text)} chars")
    log(f"             \"{text}\"")
    return text


def main() -> None:
    p = argparse.ArgumentParser()
    p.add_argument("--idol", required=True)
    p.add_argument("--urls", nargs="+", required=False)
    p.add_argument("--speaker", help="target SPEAKER_XX (미지정 시 auto_dominant)")
    p.add_argument("--list-speakers", action="store_true")
    p.add_argument("--device", default="mps", choices=["mps", "cpu"])
    p.add_argument("--skip-transcribe", action="store_true")
    args = p.parse_args()

    ensure_hf_token()
    root = data_root(args.idol)
    root.mkdir(parents=True, exist_ok=True)

    if not args.urls:
        sys.exit("ERROR: --urls 필수")

    urls = args.urls
    log(f"=== Mac voice-extract: idol={args.idol}, urls={len(urls)} ===")
    log(f"  root: {root}")
    log(f"  device: {args.device}")

    log("[1/7] fetch")
    phase_fetch(root, urls)
    log("[2/7] vocals (demucs-mlx)")
    phase_vocals(root, urls)
    log("[3/7] diarize (pyannote 3.1 / MPS)")
    phase_diarize(root, urls, args.device)

    if args.list_speakers:
        bootstrap_vid = video_id(urls[0])
        totals = {}
        for seg in parse_rttm(root / "diarized" / f"{bootstrap_vid}.rttm"):
            totals[seg["spk"]] = totals.get(seg["spk"], 0) + seg["dur"]
        log("  speakers (total seconds):")
        for spk, total in sorted(totals.items(), key=lambda kv: -kv[1]):
            log(f"    {spk}: {total:.1f}s")
        return

    log("[4/7] select")
    if args.speaker:
        target = args.speaker
        log(f"  select manual: {target}")
    else:
        target = select_speaker_auto(root, video_id(urls[0]))

    log("[5/8] match (WeSpeaker ResNet34 / MPS)")
    matched = phase_match(root, urls, target, args.device)
    log(f"  matched total: {sum(matched.values())}")

    log(f"[6/8] stitch (인접 matched → {int(PROMPT_SECONDS-2)}-{int(SEG_MAX_S)}s clips)")
    # gap_tolerance 2.0s — Weverse/vlog 형식의 자연 pause 허용
    stitched = phase_stitch(root, urls, min_dur=PROMPT_SECONDS - 2, max_dur=SEG_MAX_S, gap_tolerance=2.0)

    log("[7/8] filter (VAD + SNR)")
    kept = phase_filter(root, urls)

    log("[8/8] concat + transcribe")
    total_s, chunks = phase_concat(root)
    text = ""
    if not args.skip_transcribe:
        text = phase_transcribe(root)

    manifest = {
        "idol": args.idol,
        "urls": urls,
        "target_speaker": target,
        "matched_per_url": matched,
        "filtered_count": kept,
        "prompt_seconds": total_s,
        "prompt_chunks": chunks,
        "prompt_text": text,
        "dataset_dir": str(root / "filtered"),
    }
    (root / "manifest.json").write_text(json.dumps(manifest, indent=2, ensure_ascii=False))

    log("=== done ===")
    log(f"  dataset_dir: {root}/filtered")
    log(f"  prompt.wav:  {root}/prompt.wav")
    log(f"  prompt.txt:  {root}/prompt.txt")
    log(f"  manifest:    {root}/manifest.json")


if __name__ == "__main__":
    main()
