"""H3 DoD 3 2단 — 산출물 발화 구간이 **어느 레퍼런스 음색**에 붙었는지 관측한다.

1단(오디오 ref 없이)이 못 닫은 것은 "여성 2인이 갈렸나" 였다. 그때 막은 것은 두 가지다:
정답 앵커가 없었고(누가 어떤 목소리여야 하는지 모름), 지표가 f0 자기상관 하나였다.
여기선 앵커가 생겼다 — **합성 ref 3개가 곧 정답**이다. 지표도 아래처럼 캘리브레이션해서 쓴다.

지표
  ⓐ **MFCC(13) 평균 벡터의 코사인** — 단 도메인(녹음 채널)별 CMVN 을 먼저 뺀다.
     TTS ref(24kHz 합성)와 H3 산출물(32kHz 생성)은 채널이 달라 CMVN 없이 비교하면
     채널 차가 화자 차를 덮는다. **화자별이 아니라 도메인별**로 빼는 것이 핵심 —
     화자별로 빼면 화자 정보 자체가 0 이 된다.
  ⓑ **중앙 f0**(자기상관·옥타브 보정) — 채널에 둔감한 보조 앵커. 1단이 쓴 그 지표다.
  ⓒ **스펙트럼 기울기**(저역/고역 에너지 비) — 밝음/허스키 축의 대리지표.

캘리브레이션 (지표를 믿어도 되는지 먼저 본다)
  각 ref 를 반으로 갈라 **동일 화자쌍**의 코사인을, 서로 다른 ref 끼리 **이종 화자쌍**의 코사인을
  낸다. 두 분포가 마진을 두고 갈리지 않으면 이 지표로는 아무 판정도 하지 않는다.

판정 순서
  1. 산출물 내부 — 발화 구간끼리 갈리나 (같은 채널이라 비교가 깨끗하다)
  2. 산출물 ↔ ref — 각 구간이 어느 ref 에 가장 가깝나, 2등과의 마진은 얼마나 되나
"""
import json
import subprocess
import sys
from pathlib import Path

import numpy as np

SR = 16000
FRAME, HOP = 400, 160          # 25ms / 10ms
N_FFT, N_MEL, N_MFCC = 512, 26, 13
F0_MIN, F0_MAX = 60.0, 400.0


# ---------- io ----------

def load(path, start=None, dur=None):
    """ffmpeg 로 16kHz mono float32 로 통일해 읽는다 (wav·mp4 공통 경로)."""
    cmd = ["ffmpeg", "-v", "error"]
    if start is not None:
        cmd += ["-ss", f"{start:.3f}"]
    if dur is not None:
        cmd += ["-t", f"{dur:.3f}"]
    cmd += ["-i", str(path), "-ac", "1", "-ar", str(SR), "-f", "f32le", "-"]
    raw = subprocess.run(cmd, capture_output=True, check=True).stdout
    return np.frombuffer(raw, dtype=np.float32).astype(np.float64)


# ---------- features ----------

def frames(x):
    n = 1 + max(0, (len(x) - FRAME) // HOP)
    idx = np.arange(FRAME)[None, :] + HOP * np.arange(n)[:, None]
    return x[idx] * np.hamming(FRAME)[None, :]


def _mel_fb():
    def hz2mel(f):
        return 2595.0 * np.log10(1.0 + f / 700.0)

    def mel2hz(m):
        return 700.0 * (10.0 ** (m / 2595.0) - 1.0)

    lo, hi = hz2mel(50.0), hz2mel(SR / 2)
    pts = mel2hz(np.linspace(lo, hi, N_MEL + 2))
    bins = np.floor((N_FFT + 1) * pts / SR).astype(int)
    fb = np.zeros((N_MEL, N_FFT // 2 + 1))
    for m in range(N_MEL):
        l, c, r = bins[m], bins[m + 1], bins[m + 2]
        if c == l:
            c = l + 1
        if r == c:
            r = c + 1
        fb[m, l:c] = (np.arange(l, c) - l) / (c - l)
        fb[m, c:r] = (r - np.arange(c, r)) / (r - c)
    return fb


MEL_FB = _mel_fb()
DCT = np.cos(np.pi / N_MEL * (np.arange(N_MEL) + 0.5)[None, :] * np.arange(1, N_MFCC + 1)[:, None])


def mfcc_frames(x):
    """프레임별 MFCC(13) + 프레임 에너지(dB). c0 는 버린다 (음량축)."""
    f = frames(np.append(x[0], x[1:] - 0.97 * x[:-1]))       # pre-emphasis
    if len(f) == 0:
        return np.zeros((0, N_MFCC)), np.zeros(0)
    spec = np.abs(np.fft.rfft(f, N_FFT)) ** 2
    mel = np.log(MEL_FB @ spec.T + 1e-10)                     # (N_MEL, T)
    return (DCT @ mel).T, 10.0 * np.log10(spec.sum(axis=1) + 1e-10)


def voiced_mask(energy_db, floor_offset=25.0):
    """상위 에너지 대비 -25dB 이상을 유성으로 본다 (무음 프레임이 평균을 흐리지 않게)."""
    if len(energy_db) == 0:
        return np.zeros(0, dtype=bool)
    return energy_db > (np.percentile(energy_db, 95) - floor_offset)


def f0_median(x):
    """자기상관 f0 (옥타브 보정) — 유성 프레임만."""
    vals = []
    for i in range(0, max(0, len(x) - 1024), 512):
        w = x[i:i + 1024]
        if np.sqrt((w ** 2).mean()) < 1e-3:
            continue
        w = w - w.mean()
        ac = np.correlate(w, w, "full")[len(w) - 1:]
        lo, hi = int(SR / F0_MAX), int(SR / F0_MIN)
        if hi >= len(ac):
            continue
        seg = ac[lo:hi]
        if seg.max() <= 0 or ac[0] <= 0:
            continue
        lag = lo + int(seg.argmax())
        if ac[lag] / ac[0] < 0.3:                              # 약한 주기성은 버린다
            continue
        f = SR / lag
        half = lag // 2                                        # 옥타브 보정
        if half >= lo and ac[half] / ac[0] > 0.8 * ac[lag] / ac[0]:
            f = SR / half
        vals.append(f)
    return float(np.median(vals)) if vals else float("nan")


def tilt(x):
    """스펙트럼 기울기 = log10(고역 2~8kHz / 저역 0~1kHz) 에너지 비."""
    f = frames(x)
    if len(f) == 0:
        return float("nan")
    spec = np.abs(np.fft.rfft(f, N_FFT)) ** 2
    freq = np.fft.rfftfreq(N_FFT, 1 / SR)
    lo = spec[:, (freq > 0) & (freq <= 1000)].sum()
    hi = spec[:, (freq >= 2000) & (freq <= 8000)].sum()
    return float(np.log10((hi + 1e-12) / (lo + 1e-12)))


# ---------- segmentation ----------

def speech_segments(x, min_dur=0.5, gap=0.30):
    """에너지 VAD → 발화 덩이. 3인이 차례로 말하므로 덩이 ≈ 화자 턴이다."""
    _, e = mfcc_frames(x)
    v = voiced_mask(e, floor_offset=22.0)
    segs, run = [], None
    gap_frames = int(gap * SR / HOP)
    silence = 0
    for i, on in enumerate(v):
        if on:
            if run is None:
                run = i
            silence = 0
        elif run is not None:
            silence += 1
            if silence >= gap_frames:
                segs.append((run, i - silence))
                run = None
    if run is not None:
        segs.append((run, len(v)))
    out = []
    for s, e_ in segs:
        t0, t1 = s * HOP / SR, (e_ * HOP + FRAME) / SR
        if t1 - t0 >= min_dur:
            out.append((round(t0, 2), round(t1, 2)))
    return out


# ---------- domain-normalised speaker vector ----------

class Domain:
    """한 녹음 채널(도메인)의 CMVN 통계. 도메인 안에서만 화자 차이를 남긴다."""

    def __init__(self, waves):
        acc = []
        for w in waves:
            m, e = mfcc_frames(w)
            if len(m):
                acc.append(m[voiced_mask(e)])
        allf = np.concatenate(acc) if acc else np.zeros((1, N_MFCC))
        self.mu = allf.mean(axis=0)
        self.sd = allf.std(axis=0) + 1e-9

    def vec(self, w):
        m, e = mfcc_frames(w)
        v = voiced_mask(e)
        if v.sum() < 3:
            return None
        z = (m[v] - self.mu) / self.sd
        return z.mean(axis=0)


def cos(a, b):
    return float(np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b) + 1e-12))


def xcorr_max(a, b):
    """정규화 상호상관 최대값 — 산출물이 ref 파형을 그대로 통과시켰는지(§5.6 ⓑ) 본다."""
    n = min(len(a), len(b))
    if n < SR // 2:
        return float("nan")
    a = (a[:n] - a[:n].mean()) / (a[:n].std() + 1e-12)
    b = (b[:n] - b[:n].mean()) / (b[:n].std() + 1e-12)
    c = np.correlate(a, b, "full") / n
    return round(float(np.abs(c).max()), 3)


# ---------- main ----------

def main():
    base = Path(sys.argv[1] if len(sys.argv) > 1 else ".")
    refs = {
        "F1_hanna_bright": base / "refs/v1-hanna-female-bright.wav",
        "F2_yumin_mid": base / "refs/v2-yumin-female-mid.wav",
        "M1_jiho_deep": base / "refs/v3-jiho-male-deep.wav",
    }
    out = {"refs": {}, "calibration": {}, "clips": {}}

    ref_wav = {k: load(p) for k, p in refs.items()}
    ref_dom = Domain(list(ref_wav.values()))
    ref_vec = {k: ref_dom.vec(w) for k, w in ref_wav.items()}
    for k, w in ref_wav.items():
        out["refs"][k] = {"dur_s": round(len(w) / SR, 2), "f0_med_hz": round(f0_median(w), 1),
                          "tilt": round(tilt(w), 3)}

    # 캘리브레이션 — 같은 화자(전반/후반) vs 다른 화자
    halves = {}
    for k, w in ref_wav.items():
        h = len(w) // 2
        halves[k] = [ref_dom.vec(w[:h]), ref_dom.vec(w[h:])]
    same = {k: round(cos(a, b), 3) for k, (a, b) in halves.items()}
    diff = {}
    keys = list(refs)
    for i in range(len(keys)):
        for j in range(i + 1, len(keys)):
            diff[f"{keys[i]}~{keys[j]}"] = round(cos(ref_vec[keys[i]], ref_vec[keys[j]]), 3)
    out["calibration"] = {
        "same_speaker_halves": same, "cross_speaker": diff,
        "margin": round(min(same.values()) - max(diff.values()), 3),
        "usable": bool(min(same.values()) - max(diff.values()) > 0.15),
    }

    # 산출물 — 발화 구간별 귀속
    for label, path in [a.split("=", 1) for a in sys.argv[2:]]:
        wav = load(path)
        segs = speech_segments(wav)
        clip_dom = Domain([wav])                              # 산출물은 자기 채널로 정규화
        rows = []
        for (t0, t1) in segs:
            seg = wav[int(t0 * SR):int(t1 * SR)]
            v_self = clip_dom.vec(seg)
            v_ref = ref_dom.vec(seg)                          # ref 도메인 기준 (교차비교용)
            if v_self is None or v_ref is None:
                continue
            sims = {k: round(cos(v_ref, ref_vec[k]), 3) for k in refs}
            rank = sorted(sims.items(), key=lambda kv: -kv[1])
            # 이 클립 도메인 안에서의 **동일 화자 바닥값** — 구간을 반으로 갈라 자기 자신과 비교한다.
            # ref 도메인 캘리브레이션 값을 클립 내부 비교에 그대로 갖다 대면 도메인이 달라 틀린다.
            h = len(seg) // 2
            va, vb = clip_dom.vec(seg[:h]), clip_dom.vec(seg[h:])
            rows.append({
                "t": [t0, t1], "dur": round(t1 - t0, 2),
                "f0_med_hz": round(f0_median(seg), 1), "tilt": round(tilt(seg), 3),
                "self_halves_cos": round(cos(va, vb), 3) if va is not None and vb is not None else None,
                "sim_to_ref": sims, "best": rank[0][0], "margin": round(rank[0][1] - rank[1][1], 3),
                "xcorr_to_ref": {k: xcorr_max(seg, ref_wav[k]) for k in refs},
                "_v_self": v_self.tolist(),
            })
        # 산출물 내부 구간끼리 (같은 채널 = 깨끗한 비교)
        inner = {}
        for i in range(len(rows)):
            for j in range(i + 1, len(rows)):
                inner[f"{i}~{j}"] = round(cos(np.array(rows[i]["_v_self"]), np.array(rows[j]["_v_self"])), 3)
        for r in rows:
            del r["_v_self"]
        # ★VAD 덩이는 화자 턴과 1:1 이 아니다 — 실제로 한 덩이의 반쪽끼리 코사인이 -0.44 로 갈렸다
        #   (= 덩이 안에서 화자가 바뀐다). 그래서 **0.6s 창을 0.2s 씩 밀며** 음색 궤적을 따로 찍는다.
        #   턴 경계를 VAD 에 맡기지 않고 궤적의 변화로 본다.
        win, hop_s = 0.6, 0.2
        track = []
        t = 0.0
        while t + win <= len(wav) / SR:
            seg = wav[int(t * SR):int((t + win) * SR)]
            v = ref_dom.vec(seg)
            if v is not None and np.sqrt((seg ** 2).mean()) > 0.01:
                sims = {k: round(cos(v, ref_vec[k]), 3) for k in refs}
                rank = sorted(sims.items(), key=lambda kv: -kv[1])
                track.append({"t": round(t, 2), "f0": round(f0_median(seg), 1),
                              "sim": sims, "best": rank[0][0],
                              "margin": round(rank[0][1] - rank[1][1], 3)})
            t += hop_s

        floors = [r["self_halves_cos"] for r in rows if r["self_halves_cos"] is not None]
        out["clips"][label] = {
            "dur_s": round(len(wav) / SR, 2), "segments": rows,
            "inner_pairwise_cos": inner, "track": track,
            # 판정 기준 — 구간쌍 코사인이 **같은 구간 반쪽끼리의 최소값보다 뚜렷이 낮으면** 다른 화자다.
            "same_speaker_floor": round(min(floors), 3) if floors else None,
        }

    print(json.dumps(out, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()
