"""선정 판정의 분산 측정 — flip 이 흔들림을 줄이는가.

실측 배경: 같은 이미지·같은 팩인데 회차마다 승자가 갈렸다(S104sh7 은 4회 중
1회만 정답, S62sh4 는 개별 판정이 같은데 합의가 B·B·A 로 갈렸다). 3롤 선정은
지금 **flip 없이 1회**다 — 정순·역순 결합은 A/B 2택1 경로와 fix-rejudge 에만
걸려 있다.

두 조건을 같은 샷에 각각 N회 돌려 승자 분포를 본다.
  · noflip : 프로덕션 현행 — judge_fn 1회
  · flip   : 정순 + 역순(후보 순서를 뒤집어) → combine_flip_verdicts

판정 함수는 `make_gemini_judge_fn` 을 그대로 쓴다 — 조건부 이중(1심이
0.30 이상 앞서면 2심 생략)까지 프로덕션과 같아야 비교가 의미를 갖는다.
이미지 생성 0.

사용
  .venv/bin/python judge_variance.py <out.json> [--rounds 5] [--stems S104sh7,S62sh4]
"""
from __future__ import annotations

import argparse
import json
import sys
import threading
import traceback
from collections import Counter
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path
from typing import Any, Dict, List

sys.path.insert(0, str(Path(__file__).resolve().parent))

from app.modules.pipeline.multiroll_gemini import (  # noqa: E402
    STILL_JUDGE_PACK_VERSION, make_gemini_judge_fn, resolve_judge_texts,
)
from app.modules.pipeline.multiroll_select import (  # noqa: E402
    build_judge_schema, combine_flip_verdicts, flip_display_to_canonical,
    normalize_flip_verdict, roll_labels,
)
from reselect_v6 import (  # noqa: E402
    PROJ, RECIPE, build_size_index, resolve_refs,
)

# 정답을 아는 샷만 쓴다 — 모르면 분산은 재도 정확도는 못 잰다.
TRUTH = {
    "S104sh7": "B",   # 사용자 육안 지목
    "S62sh4": "B",    # B 가 한국 원화 (실물 확인)
}


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("out", type=Path)
    ap.add_argument("--rounds", type=int, default=5)
    ap.add_argument("--stems", default=",".join(TRUTH))
    ap.add_argument("--workers", type=int, default=4)
    args = ap.parse_args()

    records = json.loads((RECIPE / "records.json").read_text("utf-8"))
    size_idx = build_size_index()
    stems = [s.strip() for s in args.stems.split(",") if s.strip()]

    jobs = []
    for stem in stems:
        rec = records.get(stem) or {}
        labels = roll_labels(len(rec.get("verdicts") or []))
        cands = [RECIPE / f"{stem}_{lab.lower()}.png" for lab in labels]
        if not labels or any(not p.exists() for p in cands):
            print(f"건너뜀 {stem}: 후보 결손", flush=True)
            continue
        refs, unresolved = resolve_refs(rec.get("refs") or [], size_idx)
        if unresolved:
            print(f"건너뜀 {stem}: 참조 복원 실패", flush=True)
            continue
        for mode in ("noflip", "flip"):
            for r in range(args.rounds):
                jobs.append((stem, rec, labels, cands, refs, mode, r))

    print(f"측정 {len(jobs)}회 (샷 {len(stems)} × 2조건 × {args.rounds}회)",
          flush=True)

    results: List[Dict[str, Any]] = []
    lock, done = threading.Lock(), [0]

    def run(job):
        stem, rec, labels, cands, refs, mode, rnd = job
        texts = resolve_judge_texts(
            len(labels), judge_name="judge_still",
            pack_version=STILL_JUDGE_PACK_VERSION)
        judge_fn = make_gemini_judge_fn(
            judge_sys=texts["judge_sys"],
            judge_schema=build_judge_schema(labels, with_physics=True),
            project_config=None,
            step_tag="judge_variance",
        )
        prompt = rec.get("prompt", "")
        if mode == "noflip":
            res = judge_fn(f"{stem}_v{rnd}", prompt, refs, cands, labels)
            winner = res.get("winner")
        else:
            fwd = normalize_flip_verdict(
                judge_fn(f"{stem}_f{rnd}", prompt, refs, cands, labels),
                {lab: lab for lab in labels}, labels)
            rev = normalize_flip_verdict(
                judge_fn(f"{stem}_r{rnd}", prompt, refs,
                         list(reversed(cands)), labels),
                flip_display_to_canonical(labels), labels)
            winner, _ = combine_flip_verdicts(fwd, rev, labels, labels)
        return {"stem": stem, "mode": mode, "round": rnd, "winner": winner}

    with ThreadPoolExecutor(max_workers=args.workers) as ex:
        futs = {ex.submit(run, j): j for j in jobs}
        for f in as_completed(futs):
            j = futs[f]
            try:
                results.append(f.result())
            except Exception as exc:  # noqa: BLE001
                results.append({"stem": j[0], "mode": j[5], "round": j[6],
                                "error": repr(exc),
                                "trace": traceback.format_exc()[-600:]})
            with lock:
                done[0] += 1
                print(f"  [{done[0]}/{len(jobs)}] {j[0]} {j[5]} #{j[6]}",
                      flush=True)

    args.out.parent.mkdir(parents=True, exist_ok=True)
    args.out.write_text(json.dumps(
        {"rounds": args.rounds, "truth": TRUTH, "results": results},
        ensure_ascii=False, indent=1), "utf-8")

    print("\n" + "=" * 62)
    print(f"{'샷':10s} {'조건':8s} {'승자 분포':28s} {'정답':4s} {'적중'}")
    print("=" * 62)
    for stem in stems:
        for mode in ("noflip", "flip"):
            ws = [r["winner"] for r in results
                  if r.get("stem") == stem and r.get("mode") == mode
                  and r.get("winner")]
            if not ws:
                continue
            c = Counter(ws)
            dist = " ".join(f"{k}×{v}" for k, v in sorted(c.items()))
            truth = TRUTH.get(stem, "?")
            hit = sum(1 for w in ws if w == truth)
            print(f"{stem:10s} {mode:8s} {dist:28s} {truth:4s} "
                  f"{hit}/{len(ws)}")
    errs = [r for r in results if "error" in r]
    if errs:
        print(f"\n오류 {len(errs)}건 — 첫 건: {errs[0]['error'][:200]}")
    print(f"\n기록 → {args.out}")


if __name__ == "__main__":
    main()
