# -*- coding: utf-8 -*-
"""경청(傾聽) 로컬 인덱싱 킷 — 녹음 파일 폴더 → 검색 인덱스.

STT(faster-whisper)와 임베딩(bge-m3)이 **이 PC(GPU 권장)에서** 돌고,
서버에는 전사 청크 텍스트 + 벡터만 올라간다. **원본 음성은 서버로 가지 않는다.**

설치:  pip install faster-whisper transformers torch
사용:
  python indexer_speech.py --src <녹음 폴더> --server https://gyeongcheong.14dimension.com --key <API키>
  # 완전 로컬(서버 없이):  --state-dir <상태폴더>  (검색은 query_doc_local.py)

옵션:
  --model small|medium|large-v3   STT 모델 (기본 small — Phase 0 실측으로 검색엔 충분,
                                  전화녹취 등 저음질이면 medium 이상 권장)
  --lang ko|auto                  언어 (기본 ko)
  --limit N                       파일 N개까지만
"""
import os
os.environ.setdefault("HYEAN_MODALITY", "doc")   # bge-m3 임베더 (온고와 동일 공간)

import argparse
import sys
import time
from pathlib import Path

try:
    sys.stdout.reconfigure(encoding="utf-8", line_buffering=True)
except Exception:
    pass

SERVER_DIR = Path(__file__).resolve().parent
if str(SERVER_DIR) not in sys.path:
    sys.path.insert(0, str(SERVER_DIR))

AUDIO_EXTS = {".mp3", ".wav", ".m4a", ".ogg", ".flac", ".aac", ".wma", ".opus", ".webm", ".mp4"}
SEP = "::c"
# 청크 = 위스퍼 세그먼트를 창으로 병합. 실측(Phase 0): 세그먼트 단위도 충분했지만
# 너무 짧은 조각은 문맥이 얇아지므로 ~350자/40초 창으로 묶는다.
MAX_CHARS = 350
MAX_SEC = 40.0


def mmss(s: float) -> str:
    m, sec = divmod(int(s), 60)
    h, m = divmod(m, 60)
    return f"{h}:{m:02d}:{sec:02d}" if h else f"{m}:{sec:02d}"


def scan(srcs, limit):
    out, seen = [], set()
    for src in srcs:
        root = Path(src)
        if not root.exists():
            print(f"!! 소스 폴더 없음: {root}")
            continue
        for p in sorted(root.rglob("*")):
            if p.suffix.lower() in AUDIO_EXTS and p.is_file():
                s = str(p.resolve())
                if s not in seen:
                    seen.add(s)
                    out.append(s)
    return out[:limit] if limit else out


def chunk_segments(segs):
    """[(start,end,text)] → [(start,end,text)] 병합 창."""
    out, buf, t0, t1 = [], [], None, None
    for st, en, tx in segs:
        tx = tx.strip()
        if not tx:
            continue
        if buf and (sum(len(b) for b in buf) + len(tx) > MAX_CHARS or en - t0 > MAX_SEC):
            out.append((t0, t1, " ".join(buf)))
            buf, t0 = [], None
        if t0 is None:
            t0 = st
        buf.append(tx)
        t1 = en
    if buf:
        out.append((t0, t1, " ".join(buf)))
    return out


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--src", action="append", required=True)
    ap.add_argument("--server")
    ap.add_argument("--state-dir")
    ap.add_argument("--key", required=True)
    ap.add_argument("--model", default="small")
    ap.add_argument("--lang", default="ko")
    ap.add_argument("--limit", type=int)
    args = ap.parse_args()
    if not args.server and not args.state_dir:
        ap.error("--server 또는 --state-dir 중 하나가 필요합니다")

    files = scan(args.src, args.limit)
    print(f"스캔: 녹음 파일 {len(files)}건")
    if not files:
        return

    # 기존 인덱스 (증분 — 재실행해도 중복 없음)
    if args.server:
        import urllib.request, json as _json
        req = urllib.request.Request(args.server.rstrip("/") + "/index/list")
        req.add_header("x-api-key", args.key)
        try:
            with urllib.request.urlopen(req, timeout=60) as r:
                have_docs = {it["path"].rsplit(SEP, 1)[0]
                             for it in _json.loads(r.read()).get("items", [])}
        except Exception as e:
            print(f"!! 서버 조회 실패: {e}")
            return
        print(f"서버 인덱스에 기존 문서 {len(have_docs)}건")
    else:
        import tenants
        index = tenants.get_index(Path(args.state_dir), args.key)
        have_docs = {r["path"].rsplit(SEP, 1)[0] for r in index.records}

    todo = [f for f in files if f not in have_docs]
    if not todo:
        print("새로 인덱싱할 파일이 없습니다.")
        return
    print(f"신규 {len(todo)}건 — STT({args.model}) 시작")

    from faster_whisper import WhisperModel
    try:
        stt = WhisperModel(args.model, device="cuda", compute_type="float16")
    except Exception:
        print("[stt] GPU 사용 불가 — CPU(int8)로 진행 (느립니다)")
        stt = WhisperModel(args.model, device="cpu", compute_type="int8")

    import embedder as E     # bge-m3 (지연 로드)
    import tenants as T

    lang = None if args.lang == "auto" else args.lang
    total_chunks = 0
    t_all = time.time()
    for fi, f in enumerate(todo, 1):
        t0 = time.time()
        try:
            segs, info = stt.transcribe(f, language=lang, vad_filter=True)
            windows = chunk_segments([(s.start, s.end, s.text) for s in segs])
        except Exception as e:
            print(f"  !! STT 실패 {Path(f).name}: {e}")
            continue
        if not windows:
            print(f"  ({fi}/{len(todo)}) {Path(f).name}: 발화 없음 — 건너뜀")
            continue
        texts = [w[2] for w in windows]
        C = E.encode_texts(texts)
        label = T.fname_label(f)
        F = E.encode_label_texts([label])[0]
        items = []
        for i, (st_, en_, tx) in enumerate(windows):
            items.append({"path": f"{f}{SEP}{i:03d}", "doc": f,
                          "loc": f"{mmss(st_)}–{mmss(en_)}", "start_sec": float(st_),
                          "text": tx, "txt_vec": C[i].tolist(), "fname_vec": F.tolist()})
        if args.server:
            import urllib.request, json as _json
            for s0 in range(0, len(items), 32):
                body = _json.dumps({"items": items[s0:s0 + 32]}).encode()
                req = urllib.request.Request(args.server.rstrip("/") + "/index/add_speech",
                                             data=body, method="POST")
                req.add_header("Content-Type", "application/json")
                req.add_header("x-api-key", args.key)
                with urllib.request.urlopen(req, timeout=120) as r:
                    _json.loads(r.read())
        else:
            import numpy as np
            recs = [{"path": it["path"], "thumb": it["text"].encode("utf-8"),
                     "extra": {"doc": it["doc"], "loc": it["loc"], "kind": "speech",
                               "start_sec": it["start_sec"], "snippet": it["text"][:200]}}
                    for it in items]
            Cv = np.asarray([it["txt_vec"] for it in items], dtype=np.float32)
            Fv = np.asarray([it["fname_vec"] for it in items], dtype=np.float32)
            index.add_batch(recs, np.zeros_like(Cv), Cv, Fv)
        total_chunks += len(items)
        dur = getattr(info, "duration", 0) or 0
        print(f"  ({fi}/{len(todo)}) {Path(f).name}: {dur:.0f}s 오디오 → {len(items)}청크 "
              f"({time.time()-t0:.0f}s)")
    print(f"\n완료 — 파일 {len(todo)}건, 청크 {total_chunks}개, {time.time()-t_all:.0f}s")
    print("원본 음성은 업로드되지 않았습니다 — 전사 청크와 벡터만 서버에 있습니다.")


if __name__ == "__main__":
    main()
