#!/usr/bin/env python3
"""GPT vs Gemini Pro 전체 캐스케이드 비교.

기존 프로젝트에서 step 4(visual_world_rules) 이후 전체 파이프라인을 재실행.
6개 대규모 입력 단계만 gemini-pro, 나머지는 기본 모델 유지.

사용법:
    cd backend && .venv/bin/python ../tests/model_comparison/run_gemini_cascade.py --scenario srd_ep1
    cd backend && .venv/bin/python ../tests/model_comparison/run_gemini_cascade.py --scenario srd_ep1 --dry-run
    cd backend && .venv/bin/python ../tests/model_comparison/run_gemini_cascade.py --scenario srd_ep1 --step visual_world_rules
"""
import argparse
import json
import logging
import os
import shutil
import sys
import time
from pathlib import Path

BACKEND_DIR = Path(__file__).resolve().parent.parent.parent / "backend"
sys.path.insert(0, str(BACKEND_DIR))
os.chdir(str(BACKEND_DIR))

logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
)
logger = logging.getLogger("gemini_cascade")

COMP_DIR = Path(__file__).resolve().parent

# ── 프로젝트 ID 매핑 ──
PROJECTS = {
    "srd_ep1": {
        "project_id": "184bd3f1-9415-40c6-80a0-9f6d4901fa9f",
        "episode_id": "224296aa-91cb-4749-a42c-92b14a39a0bb",
        "name": "SRD v4.3",
    },
    "mia_ep1": {
        "project_id": "0f6fa335-1a98-4b15-887e-f0478a40d830",
        "episode_id": "fe88d126-102e-401e-846a-bad6cc3eb016",
        "name": "미아 v1 Pipeline",
    },
}

# ── 재실행 단계 (파이프라인 순서) ──
STEPS_TO_RERUN = [
    "visual_world_rules",        # order 4  — GEMINI PRO
    "scene_summary",             # order 7  — 재실행 (VWR 입력 변경)
    "entity_extract_character",  # order 8  — GEMINI PRO
    "entity_extract_location",   # order 9  — GEMINI PRO
    "entity_extract_prop",       # order 10 — GEMINI PRO
    "entity_filter",             # order 11 — 재실행 (entity 입력 변경)
    "entity_review",             # order 12 — 재실행
    "entity_detail",             # order 13 — 재실행
    "entity_t2i",                # order 14 — 재실행
    # ── mid-pipeline DB sync ──
    "scene_director",            # order 15 — 재실행 (entity_t2i 변경)
    "scene_cinematography",      # order 16 — GEMINI PRO
    "scene_dependency",          # order 17 — GEMINI PRO
    "outlook_extraction",        # order 18 — 재실행 (scene_director 변경)
    "scene_detail",              # order 19 — 재실행 (모든 입력 변경)
    "scene_verify",              # order 20 — 재실행
]

# 6개 대규모 입력 단계 — gemini-pro 오버라이드
GEMINI_OVERRIDES = {
    "visual_world_rules": {"model": "gemini-pro"},
    "entity_extract": {"model": "gemini-pro"},       # entity_extractor_v4에서 사용하는 step 이름
    "scene_cinematography": {"model": "gemini-pro"},
    "scene_dependency": {"model": "gemini-pro"},
}


def get_step_runner(step_id, project_id, episode_id, db, project_config):
    """StepRunner 인스턴스 생성 (steps.py의 _get_step_runner 복제)."""
    from app.core.steps import STEP_CLASSES as V3_CLASSES
    try:
        from app.core.steps.analysis_steps import ANALYSIS_STEP_CLASSES
    except ImportError:
        ANALYSIS_STEP_CLASSES = {}
    try:
        from app.core.steps.image_steps import IMAGE_STEP_CLASSES
    except ImportError:
        IMAGE_STEP_CLASSES = {}

    step_classes = {**ANALYSIS_STEP_CLASSES, **IMAGE_STEP_CLASSES, **V3_CLASSES}
    cls = step_classes.get(step_id)
    if not cls:
        raise ValueError(f"Step class not found: {step_id}")

    return cls(
        step_id=step_id,
        project_id=project_id,
        episode_id=episode_id,
        db=db,
        project_config=project_config,
    )


def dry_run(scenario: str, project_config: dict):
    """드라이런 — 게이트 상태와 실행 계획만 확인."""
    from app.core.database import SessionLocal
    from sqlalchemy import text
    from app.core.config import settings

    info = PROJECTS[scenario]
    pid, eid = info["project_id"], info["episode_id"]

    db = SessionLocal()
    try:
        print(f"\n{'='*60}")
        print(f"  드라이런: {info['name']} ({scenario})")
        print(f"{'='*60}\n")

        for step_id in STEPS_TO_RERUN:
            # step_run 상태
            row = db.execute(text(
                "SELECT status, resolved_model FROM step_run "
                "WHERE project_id = :pid AND episode_id = :eid AND step_id = :sid"
            ), {"pid": pid, "eid": eid, "sid": step_id}).fetchone()
            status = row[0] if row else "없음"
            model = row[1] if row else "-"

            # 체크포인트 존재
            cp_path = (
                Path(settings.projects_dir) / pid
                / "checkpoints" / "episodes" / eid / step_id / "manifest.json"
            )
            cp_exists = cp_path.exists()

            # 새 모델
            from app.modules.llm.llm_client import _resolve_model
            new_model = _resolve_model(step_id, project_config)

            marker = "★" if new_model == "gemini-pro" and model != "gemini-pro" else " "
            print(f"  {marker} {step_id:30s}  현재={status:12s} 모델={model:12s} → {new_model:12s}  cp={'✓' if cp_exists else '✗'}")

        # GPT 백업 확인
        backup_dir = COMP_DIR / scenario
        print(f"\n  GPT 백업: {backup_dir}")
        print(f"  백업 파일: {len(list(backup_dir.glob('*.json')))} files")

    finally:
        db.close()


def save_gemini_checkpoints(scenario: str, project_id: str, episode_id: str):
    """Gemini 결과 체크포인트를 비교 디렉토리에 복사."""
    from app.core.config import settings

    out_dir = COMP_DIR / f"{scenario}_gemini"
    out_dir.mkdir(exist_ok=True)

    cp_base = (
        Path(settings.projects_dir) / project_id
        / "checkpoints" / "episodes" / episode_id
    )

    for step_id in STEPS_TO_RERUN:
        src = cp_base / step_id / "manifest.json"
        if src.exists():
            dst = out_dir / f"{step_id}.json"
            shutil.copy2(src, dst)
            logger.info("Copied: %s → %s", step_id, dst.name)
        else:
            logger.warning("Missing checkpoint: %s", step_id)


def compare_results(scenario: str):
    """GPT vs Gemini 전체 비교 리포트."""
    gpt_dir = COMP_DIR / scenario
    gem_dir = COMP_DIR / f"{scenario}_gemini"

    if not gem_dir.exists():
        logger.warning("Gemini 결과 없음: %s", scenario)
        return

    lines = [f"\n{'='*70}", f"  GPT vs Gemini Pro 비교: {scenario}", f"{'='*70}\n"]

    for step_id in STEPS_TO_RERUN:
        gpt_path = gpt_dir / f"{step_id}.json"
        gem_path = gem_dir / f"{step_id}.json"

        if not gpt_path.exists() or not gem_path.exists():
            lines.append(f"  [{step_id}] 데이터 없음")
            continue

        gpt = json.loads(gpt_path.read_text(encoding="utf-8"))
        gem = json.loads(gem_path.read_text(encoding="utf-8"))

        gpt_model = gpt.get("resolved_model", "?")
        gem_model = gem.get("resolved_model", "?")
        lines.append(f"  [{step_id}]  GPT={gpt_model}  Gemini={gem_model}")

        _compare_step_data(step_id, gpt.get("data", {}), gem.get("data", {}), lines)
        lines.append("")

    report = "\n".join(lines)
    print(report)

    report_path = gem_dir / "comparison_report.txt"
    report_path.write_text(report, encoding="utf-8")
    logger.info("Report saved: %s", report_path)


def _compare_step_data(step_id: str, gpt_d: dict, gem_d: dict, lines: list):
    """단계별 상세 비교."""
    if step_id == "visual_world_rules":
        lines.append(f"    GPT: {len(gpt_d.get('rules',[]))} rules, {len(gpt_d.get('director_notes',[]))} notes")
        lines.append(f"    Gem: {len(gem_d.get('rules',[]))} rules, {len(gem_d.get('director_notes',[]))} notes")

    elif step_id.startswith("entity_extract_"):
        etype = step_id.rsplit("_", 1)[-1]
        key = f"{etype}s" if etype != "prop" else "props"
        gpt_ents = gpt_d.get(key, [])
        gem_ents = gem_d.get(key, [])
        gpt_names = {e.get("name") for e in gpt_ents}
        gem_names = {e.get("name") for e in gem_ents}
        lines.append(f"    GPT: {len(gpt_ents)} / Gem: {len(gem_ents)}")
        only_gpt = gpt_names - gem_names
        only_gem = gem_names - gpt_names
        if only_gpt:
            lines.append(f"    GPT only: {', '.join(sorted(only_gpt))}")
        if only_gem:
            lines.append(f"    Gem only: {', '.join(sorted(only_gem))}")

    elif step_id == "entity_filter":
        lines.append(f"    GPT removed: {gpt_d.get('removed_count', '?')} / Gem removed: {gem_d.get('removed_count', '?')}")

    elif step_id == "entity_review":
        gpt_un = gpt_d.get("unnecessary_entities", gpt_d.get("review", {}).get("unnecessary_entities", []))
        gem_un = gem_d.get("unnecessary_entities", gem_d.get("review", {}).get("unnecessary_entities", []))
        if isinstance(gpt_un, list) and isinstance(gem_un, list):
            lines.append(f"    GPT 불필요: {len(gpt_un)} / Gem 불필요: {len(gem_un)}")

    elif step_id == "entity_t2i":
        gpt_c = len(gpt_d.get("characters", []))
        gem_c = len(gem_d.get("characters", []))
        gpt_l = len(gpt_d.get("locations", []))
        gem_l = len(gem_d.get("locations", []))
        gpt_p = len(gpt_d.get("props", []))
        gem_p = len(gem_d.get("props", []))
        lines.append(f"    GPT: C{gpt_c}/L{gpt_l}/P{gpt_p} / Gem: C{gem_c}/L{gem_l}/P{gem_p}")

    elif step_id == "scene_director":
        gpt_scenes = gpt_d.get("scenes", [])
        gem_scenes = gem_d.get("scenes", [])
        lines.append(f"    GPT: {len(gpt_scenes)} scenes / Gem: {len(gem_scenes)} scenes")
        # VE 차이
        diffs = 0
        for gs in gpt_scenes:
            si = gs.get("scene_index")
            gms = next((s for s in gem_scenes if s.get("scene_index") == si), {})
            if set(gs.get("present_entity_ids", [])) != set(gms.get("present_entity_ids", [])):
                diffs += 1
        if diffs:
            lines.append(f"    VE 차이 있는 씬: {diffs}/{len(gpt_scenes)}")

    elif step_id == "scene_cinematography":
        gpt_s = gpt_d.get("scenes", [])
        gem_s = gem_d.get("scenes", [])
        lines.append(f"    GPT: {len(gpt_s)} / Gem: {len(gem_s)} scenes")
        diffs = 0
        for gs in gpt_s:
            si = gs.get("scene_index")
            gms = next((s for s in gem_s if s.get("scene_index") == si), {})
            if gs.get("shot_1") != gms.get("shot_1") or gs.get("shot_2") != gms.get("shot_2"):
                diffs += 1
        lines.append(f"    Shot 차이: {diffs} scenes")

    elif step_id == "scene_dependency":
        lines.append(f"    GPT: {len(gpt_d.get('dependencies',[]))} / Gem: {len(gem_d.get('dependencies',[]))} deps")

    elif step_id == "outlook_extraction":
        gpt_ol = gpt_d.get("outlooks", [])
        gem_ol = gem_d.get("outlooks", [])
        gpt_sa = gpt_d.get("scene_assignments", [])
        gem_sa = gem_d.get("scene_assignments", [])
        lines.append(f"    GPT: {len(gpt_ol)} outlooks, {len(gpt_sa)} assignments")
        lines.append(f"    Gem: {len(gem_ol)} outlooks, {len(gem_sa)} assignments")

    elif step_id == "scene_detail":
        gpt_s = gpt_d.get("scenes", [])
        gem_s = gem_d.get("scenes", [])
        lines.append(f"    GPT: {len(gpt_s)} / Gem: {len(gem_s)} scenes")

    elif step_id == "scene_verify":
        gpt_s = gpt_d.get("scenes", [])
        gem_s = gem_d.get("scenes", [])
        lines.append(f"    GPT: {len(gpt_s)} / Gem: {len(gem_s)} scenes")


def run_cascade(scenario: str, project_config: dict, start_from: str = None):
    """전체 캐스케이드 실행."""
    from app.core.database import SessionLocal
    from app.api.v1.steps import _sync_checkpoints_to_db

    info = PROJECTS[scenario]
    pid, eid = info["project_id"], info["episode_id"]

    db = SessionLocal()
    try:
        steps = STEPS_TO_RERUN
        if start_from:
            idx = steps.index(start_from)
            steps = steps[idx:]

        total = len(steps)
        logger.info("=" * 60)
        logger.info("  캐스케이드 시작: %s (%d steps)", info["name"], total)
        logger.info("=" * 60)

        t_total = time.time()

        for i, step_id in enumerate(steps, 1):
            # mid-pipeline DB sync (entity_t2i 완료 후, scene_director 전)
            if step_id == "scene_director":
                logger.info("── mid-pipeline DB sync ──")
                try:
                    _sync_checkpoints_to_db(pid, eid, db)
                except Exception as e:
                    logger.warning("Mid sync failed (non-fatal): %s", e)
                    db.rollback()

            logger.info("── [%d/%d] %s ──", i, total, step_id)
            t0 = time.time()

            try:
                runner = get_step_runner(step_id, pid, eid, db, project_config)

                # 첫 단계(VWR)만 force → 하위 무효화, 나머지는 resume
                mode = "force" if step_id == steps[0] else "force"
                # 실제로는 모든 단계를 force로 돌림 (확실한 재실행)
                result = runner.run(mode=mode)

                elapsed = time.time() - t0
                status = result.get("status", "?")
                logger.info("  ✓ %s: %s (%.1fs)", step_id, status, elapsed)

                if status in ("failed",):
                    logger.error("  단계 실패, 중단합니다.")
                    break

            except Exception as exc:
                elapsed = time.time() - t0
                logger.error("  ✗ %s 에러 (%.1fs): %s", step_id, elapsed, exc)
                import traceback
                traceback.print_exc()
                break

        total_elapsed = time.time() - t_total
        logger.info("=" * 60)
        logger.info("  캐스케이드 완료: %.1f분", total_elapsed / 60)
        logger.info("=" * 60)

        # 최종 DB sync
        logger.info("── final DB sync ──")
        try:
            _sync_checkpoints_to_db(pid, eid, db)
        except Exception as e:
            logger.error("Final sync failed: %s", e)
            db.rollback()

        # 결과 저장
        save_gemini_checkpoints(scenario, pid, eid)

    finally:
        db.close()


def main():
    parser = argparse.ArgumentParser(description="GPT vs Gemini Pro cascade comparison")
    parser.add_argument("--scenario", required=True, choices=["srd_ep1", "mia_ep1"])
    parser.add_argument("--step", default=None,
                        help="특정 단계부터 시작 (e.g., scene_director)")
    parser.add_argument("--dry-run", action="store_true",
                        help="실행 없이 계획만 확인")
    parser.add_argument("--compare", action="store_true",
                        help="비교 리포트만 생성")
    args = parser.parse_args()

    project_config = dict(GEMINI_OVERRIDES)

    if args.dry_run:
        dry_run(args.scenario, project_config)
        return

    if args.compare:
        compare_results(args.scenario)
        return

    run_cascade(args.scenario, project_config, start_from=args.step)
    compare_results(args.scenario)


if __name__ == "__main__":
    main()
