"""C 실행 223쌍 오프라인 fix-rejudge — [수정 전 원본 vs 수정본] 블라인드 재판정.

기존 방식 그대로 재현한다(multiroll_select._critique_and_fix 의 재판정 절):
canonical A=원본 / B=수정본, 정순 1회 + 역순 1회(flip) → combine_flip_verdicts,
동점 우선순위 ["B","A"](동점=수정본). 이미지 생성 없음 — 판정 콜만.

기록의 참조가 `<bytes:N>` 로 남은 것(인라인 첨부)은 **바이트 길이 정확 일치**로
프로젝트 이미지에서 복원한다. 복원 실패 시 그 쌍은 건너뛰고 사유를 남긴다 —
참조를 비운 채 판정하면 원 판정과 다른 계약이 된다.
"""
from __future__ import annotations

import json
import os
import sys
import threading
import traceback
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path

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

from app.modules.pipeline.multiroll_gemini import (  # noqa: E402
    load_fix_rejudge_header, 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, _compose_critique_prompt,
)

PROJ = "e716bafb-24bb-42b7-aea0-fdb383844ee8"
EPI = "d6a9aa85-b75e-400c-980c-4ee7e876a15b"
ROOT = Path(__file__).resolve().parent.parent
RECIPE = ROOT / f"projects/{PROJ}/images/{EPI}/scene/recipe"
OUT = Path(sys.argv[1]) if len(sys.argv) > 1 else Path("rejudge_result.json")
WORKERS = int(os.environ.get("REJUDGE_WORKERS", "6"))
LIMIT = int(os.environ.get("REJUDGE_LIMIT", "0"))


def build_size_index() -> dict[int, list[Path]]:
    idx: dict[int, list[Path]] = {}
    for p in (ROOT / f"projects/{PROJ}").rglob("*.png"):
        idx.setdefault(p.stat().st_size, []).append(p)
    return idx


def resolve_refs(entries, size_idx):
    """기록의 refs → [(label, Path)] — `<bytes:N>` 은 크기 일치로 복원."""
    out, unresolved = [], []
    for e in entries:
        label, path = e.get("label", ""), str(e.get("path", ""))
        if path.startswith("<bytes:"):
            n = int(path[len("<bytes:"):-1])
            cands = size_idx.get(n, [])
            if len(cands) >= 1:
                out.append((label, cands[0]))
            else:
                unresolved.append((label, n))
            continue
        p = Path(path)
        if p.exists():
            out.append((label, p))
        else:
            unresolved.append((label, path))
    return out, unresolved


def main() -> None:
    records = json.loads((RECIPE / "records.json").read_text("utf-8"))
    size_idx = build_size_index()

    texts = resolve_judge_texts(2, judge_name="judge_still")
    header = load_fix_rejudge_header()
    judge_fn = make_gemini_judge_fn(
        judge_sys=texts["judge_sys"],
        judge_schema=build_judge_schema(roll_labels(2)),
        project_config=None,
        step_tag="still_recipe_fix_rejudge",
        prompt_header=header,
    )

    jobs = []
    skipped = []
    for stem, rec in records.items():
        if not isinstance(rec, dict) or "fix_prompt" not in rec:
            continue
        sel = str(rec.get("selected", "")).strip()
        orig = RECIPE / f"{stem}_{sel.lower()}.png"
        fixed = RECIPE / f"{stem}_fix.png"
        if not (orig.exists() and fixed.exists()):
            skipped.append({"stem": stem, "why": "이미지 결손"})
            continue
        roll_prompts = rec.get("roll_prompts")
        prompt = _compose_critique_prompt(
            rec.get("prompt", ""), roll_prompts, sel, shared_prompt=None)
        raw_refs = (rec.get("roll_refs") or {}).get(sel) or rec.get("refs") or []
        refs, unresolved = resolve_refs(raw_refs, size_idx)
        if unresolved:
            skipped.append({"stem": stem, "why": "참조 복원 실패",
                            "detail": unresolved})
            continue
        jobs.append((stem, sel, prompt, refs, orig, fixed,
                     rec.get("critique") or {}))
    if LIMIT:
        jobs = jobs[:LIMIT]

    print(f"판정 대상 {len(jobs)}쌍 · 건너뜀 {len(skipped)}", flush=True)
    results, lock, done = {}, threading.Lock(), [0]

    def run(job):
        stem, sel, prompt, refs, orig, fixed, crit = job
        labels2 = ["A", "B"]
        cand = {"A": orig, "B": fixed}
        fwd_raw = judge_fn(f"{stem}_fixjudge", prompt, refs,
                           [cand["A"], cand["B"]], labels2)
        fwd = normalize_flip_verdict(fwd_raw, {x: x for x in labels2}, labels2)
        rev_raw = judge_fn(f"{stem}_fixjudge_rev", prompt, refs,
                           [cand["B"], cand["A"]], labels2)
        rev = normalize_flip_verdict(
            rev_raw, flip_display_to_canonical(labels2), labels2)
        winner, combined = combine_flip_verdicts(fwd, rev, labels2, ["B", "A"])
        return stem, {
            "selected_roll": sel,
            "orig": str(orig), "fixed": str(fixed),
            "forward_normalized": fwd, "reverse_normalized": rev,
            "combined": combined, "winner": winner, "fix_won": winner == "B",
            "issues": [i.get("issue_ko", "") for i in (crit.get("issues") or [])],
        }

    with ThreadPoolExecutor(max_workers=WORKERS) as ex:
        futs = {ex.submit(run, j): j[0] for j in jobs}
        for f in as_completed(futs):
            stem = futs[f]
            try:
                k, v = f.result()
                results[k] = v
            except Exception as exc:  # noqa: BLE001
                results[stem] = {"error": repr(exc),
                                 "trace": traceback.format_exc()[-800:]}
            with lock:
                done[0] += 1
                if done[0] % 10 == 0 or done[0] == len(jobs):
                    ok = sum(1 for v in results.values() if "winner" in v)
                    fw = sum(1 for v in results.values() if v.get("fix_won"))
                    print(f"  {done[0]}/{len(jobs)}  성공 {ok}  "
                          f"수정본 승 {fw}", flush=True)
                    OUT.write_text(json.dumps(
                        {"results": results, "skipped": skipped},
                        ensure_ascii=False, indent=1), "utf-8")

    OUT.write_text(json.dumps({"results": results, "skipped": skipped},
                              ensure_ascii=False, indent=1), "utf-8")
    ok = [v for v in results.values() if "winner" in v]
    fw = sum(1 for v in ok if v["fix_won"])
    print(f"\n완료 — 판정 성공 {len(ok)} / 오류 {len(results)-len(ok)}")
    if ok:
        print(f"  수정본(B) 승 {fw} ({100*fw/len(ok):.1f}%)  "
              f"원본(A) 승 {len(ok)-fw} ({100*(len(ok)-fw)/len(ok):.1f}%)")


if __name__ == "__main__":
    main()
