"""shot_selection canary — Flash 3.5 vs 3.7 이 **고르는 샷**이 갈리나.

왜 이 스텝인가: Flash 교체가 그림에 닿는 자리는 실측으로 둘뿐이었다 —
`shot_selection`(무엇을 그릴지)과 `shot_dependency_t2i`(참조 선택).
후자는 이미 봤고(edge 6/6 동일) 이것만 남았다. 고르는 샷이 갈리면
그림이 통째로 달라지므로 영향이 가장 크다.

지키는 것은 dependency canary 와 같다 — 이미지 0장, CP 안 건드림
(`_execute` 만 부른다), 조립을 다시 만들지 않고 `call_structured` 를
감싸 나가는 것을 기록만 한다. 모드는 **프로세스로** 가른다.

    python scratchpad/canary_shot_selection.py --model gemini-3.5-flash [--dry]
    python scratchpad/canary_shot_selection.py --compare A.json B.json
"""
import argparse
import json
import os
import sys
from pathlib import Path

ROOT = Path("/Users/manta/Documents/Projects/TheRoad-I1")
BACKEND = ROOT / "backend"
OUT_DIR = ROOT / "scratchpad" / "canary_shot_selection"

PROJECT_ID = "da049582-2c6d-492c-979d-f468d61bab6e"
EPISODE_ID = "fb7a883f-baac-4145-9131-732ce628d474"


def run_one(model: str, dry: bool) -> Path:
    os.environ["GEMINI_FLASH_MODEL"] = model
    sys.path.insert(0, str(BACKEND))
    os.chdir(BACKEND)  # config.py 의 env_file=".env" 는 상대 경로다

    from app.core.config import settings
    assert settings.gemini_flash_model == model, (
        f"env 가 안 먹었다: {settings.gemini_flash_model!r}")

    from app.core.database import SessionLocal
    from app.core.steps import shot_selection_step as mod
    from app.services.analysis_dispatch_service import get_step_runner
    from app.services.step_execution_service import _load_project_config

    calls = []
    real_call = mod.call_structured

    def _spy(**kw):
        # ★이 스텝은 **병렬**이라 호출 순서가 씬 순서가 아니다. 씬을 안
        #  남기면 답을 엉뚱한 씬에 붙여 읽게 된다(실제로 한 번 그랬다).
        usr = kw.get("user_prompt") or ""
        rec = {"usr_len": len(usr), "sys_len": len(kw.get("system_prompt") or ""),
               "usr_head": usr[:300]}
        if dry:
            rec["result"] = "<dry>"
            calls.append(rec)
            return {"selected_shots": []}
        result = real_call(**kw)
        rec["result"] = result
        calls.append(rec)
        return result

    mod.call_structured = _spy
    db = SessionLocal()
    try:
        runner = get_step_runner(
            "shot_selection", PROJECT_ID, EPISODE_ID, db,
            _load_project_config(db, PROJECT_ID),
            opik_context={"run_tag": f"canary_flash_{model}",
                          "project_name": "canary_flash_selection",
                          "episode_title": "골목 끝"},
        )
        result = runner._execute(mode="resume")
    finally:
        mod.call_structured = real_call
        db.close()

    picks = {f"S{s['scene_index']}": s.get("selected_shot_indices", [])
             for s in result.get("data", {}).get("scenes", [])}
    payload = {"model": model, "dry": dry,
               "llm_calls": len(calls),
               "usr_chars": sum(c["usr_len"] for c in calls),
               "picks": dict(sorted(picks.items())),
               "calls": calls}
    OUT_DIR.mkdir(parents=True, exist_ok=True)
    out = OUT_DIR / f"{model}_{'dry' if dry else 'live'}.json"
    out.write_text(json.dumps(payload, ensure_ascii=False, indent=2),
                   encoding="utf-8")
    print(f"[{model} {'dry' if dry else 'live'}] calls={len(calls)} "
          f"usr_chars={payload['usr_chars']}")
    for k, v in payload["picks"].items():
        print(f"  {k}: {v}")
    print(f"  saved: {out}")
    return out


def compare(paths) -> int:
    docs = [json.loads(Path(p).read_text(encoding="utf-8")) for p in paths]
    keys = sorted({k for d in docs for k in d["picks"]})
    print("\n=== 고른 샷 대조 ===")
    print(f"{'scene':8s} " + " ".join(f"{d['model']:<22s}" for d in docs))
    diff = []
    for k in keys:
        cells = [d["picks"].get(k) for d in docs]
        same = len({str(c) for c in cells}) == 1
        print(f"{k:8s} " + " ".join(f"{str(c):<22s}" for c in cells)
              + ("" if same else "   <-- 다름"))
        if not same:
            diff.append(k)
    print("\n=== 발송량 ===")
    for d in docs:
        print(f"  {d['model']}: calls={d['llm_calls']} usr_chars={d['usr_chars']}")
    print(f"\n다른 씬: {diff if diff else '없음 — 고른 샷 동일'}")
    return len(diff)


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--model")
    ap.add_argument("--dry", action="store_true")
    ap.add_argument("--compare", nargs="+")
    a = ap.parse_args()
    if a.compare:
        return compare(a.compare)
    if not a.model:
        ap.error("--model 또는 --compare 가 필요하다")
    run_one(a.model, a.dry)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
