"""GPT vs Gemini Pro 비교 — 6개 대규모 문맥 단계.

사용법:
    cd backend && python -m tests.model_comparison.run_gemini_comparison --scenario srd_ep1
    cd backend && python -m tests.model_comparison.run_gemini_comparison --scenario mia_ep1
"""
import argparse
import json
import logging
import os
import sys
import time
from pathlib import Path

# backend/ 를 sys.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("model_comparison")

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


def load_checkpoint(scenario: str, step: str) -> dict:
    """저장된 체크포인트 로드."""
    fpath = COMP_DIR / scenario / f"{step}.json"
    if not fpath.exists():
        raise FileNotFoundError(f"Checkpoint not found: {fpath}")
    return json.loads(fpath.read_text(encoding="utf-8"))


def load_fulltext(scenario: str) -> str:
    """저장된 fulltext 로드."""
    fpath = COMP_DIR / scenario / "fulltext.txt"
    return fpath.read_text(encoding="utf-8")


def save_result(scenario: str, step: str, result: dict, model: str, elapsed: float):
    """결과 저장."""
    out_dir = COMP_DIR / f"{scenario}_gemini"
    out_dir.mkdir(exist_ok=True)
    result_with_meta = {
        "model": model,
        "elapsed_seconds": round(elapsed, 1),
        "data": result,
    }
    fpath = out_dir / f"{step}.json"
    fpath.write_text(
        json.dumps(result_with_meta, ensure_ascii=False, indent=2),
        encoding="utf-8",
    )
    logger.info("Saved: %s (%.1fs)", fpath.name, elapsed)


# ── 개별 단계 실행 ──

def run_visual_world_rules(scenario: str, project_config: dict):
    """#4 visual_world_rules — gemini-pro."""
    from app.modules.pipeline.visual_world_rules import extract_visual_rules

    fulltext = load_fulltext(scenario)
    ep_cp = load_checkpoint(scenario, "episode_summary")
    episode_summary = ep_cp["data"].get("summary", "")

    logger.info("=== visual_world_rules [gemini-pro] ===")
    logger.info("  fulltext: %d chars, summary: %d chars", len(fulltext), len(episode_summary))

    t0 = time.time()
    result = extract_visual_rules(
        fulltext=fulltext,
        episode_summary=episode_summary,
        project_config=project_config,
    )
    elapsed = time.time() - t0

    rules_count = len(result.get("rules", []))
    notes_count = len(result.get("director_notes", []))
    logger.info("  Result: %d rules, %d director_notes, era=%s",
                rules_count, notes_count, result.get("era", "?"))

    save_result(scenario, "visual_world_rules", result, "gemini-pro", elapsed)
    return result


def run_entity_extract(scenario: str, entity_type: str, project_config: dict):
    """#8/#9/#10 entity_extract — gemini-pro."""
    from app.modules.pipeline.entity_extractor_v4 import extract_entities_by_type

    fulltext = load_fulltext(scenario)
    vwr_cp = load_checkpoint(scenario, "visual_world_rules")
    visual_rules = json.dumps(vwr_cp["data"], ensure_ascii=False)
    ss_cp = load_checkpoint(scenario, "scene_save")
    segments_json = json.dumps(ss_cp["data"]["segments"], ensure_ascii=False)

    logger.info("=== entity_extract_%s [gemini-pro] ===", entity_type)
    logger.info("  fulltext: %d chars, segments: %d", len(fulltext), len(ss_cp["data"]["segments"]))

    t0 = time.time()
    entities = extract_entities_by_type(
        fulltext=fulltext,
        entity_type=entity_type,
        visual_rules=visual_rules,
        segments_json=segments_json,
        project_config=project_config,
    )
    elapsed = time.time() - t0

    logger.info("  Result: %d %ss extracted (after 2+ scene filter)", len(entities), entity_type)
    for e in entities:
        scenes = e.get("scene_appearances", [])
        logger.info("    - %s: %d scenes", e.get("name", "?"), len(scenes))

    save_result(scenario, f"entity_extract_{entity_type}", entities, "gemini-pro", elapsed)
    return entities


def run_scene_cinematography(scenario: str, project_config: dict):
    """#16 scene_cinematography — gemini-pro."""
    from app.modules.llm.llm_client import call_structured
    from app.modules.prompt_loader import load_prompt, load_schema

    fulltext = load_fulltext(scenario)
    ss_cp = load_checkpoint(scenario, "scene_save")
    segments = ss_cp["data"]["segments"]

    system_prompt = load_prompt("scene_cinematography", "system")
    analyze_template = load_prompt("scene_cinematography", "analyze")
    schema = load_schema("scene_cinematography", "analyze_schema")

    # DB에서 shot_types 로드
    from app.core.database import SessionLocal
    from sqlalchemy import text as sql_text
    db = SessionLocal()
    try:
        shot_rows = db.execute(
            sql_text(
                "SELECT name, category, description FROM shot_type "
                "WHERE is_active = true ORDER BY sort_order"
            )
        ).fetchall()
    finally:
        db.close()

    shot_types_block = "\n".join(f"- {r[0]} [{r[1]}]: {r[2]}" for r in shot_rows)

    scenes_block_items = []
    for seg in segments:
        scene_text = fulltext[seg["start_char"]:seg["end_char"]]
        scenes_block_items.append(f"씬 {seg['scene_index']}: {seg['heading']}\n{scene_text}")
    scenes_block = "\n\n".join(scenes_block_items)

    user_prompt = analyze_template.format(
        shot_types_block=shot_types_block,
        scenes_block=scenes_block,
    )

    logger.info("=== scene_cinematography [gemini-pro] ===")
    logger.info("  %d scenes, %d shot types", len(segments), len(shot_rows))

    t0 = time.time()
    result = call_structured(
        step="scene_cinematography",
        system_prompt=system_prompt,
        user_prompt=user_prompt,
        response_schema=schema,
        project_config=project_config,
        schema_name="scene_cinematography",
    )
    elapsed = time.time() - t0

    scenes = result.get("scenes", [])
    logger.info("  Result: %d scenes with shot types", len(scenes))

    save_result(scenario, "scene_cinematography", result, "gemini-pro", elapsed)
    return result


def run_scene_dependency(scenario: str, project_config: dict):
    """#17 scene_dependency — gemini-pro."""
    from app.modules.pipeline.scene_dependency_v2 import extract_dependencies

    fulltext = load_fulltext(scenario)
    ss_cp = load_checkpoint(scenario, "scene_save")
    segments = ss_cp["data"]["segments"]

    director_cp = load_checkpoint(scenario, "scene_director")
    director_result = director_cp["data"]

    entity_cp = load_checkpoint(scenario, "entity_t2i")
    entities = entity_cp["data"]

    logger.info("=== scene_dependency [gemini-pro] ===")
    logger.info("  %d segments, %d director scenes", len(segments), len(director_result.get("scenes", [])))

    t0 = time.time()
    result = extract_dependencies(
        segments=segments,
        fulltext=fulltext,
        director_result=director_result,
        entities=entities,
        project_config=project_config,
    )
    elapsed = time.time() - t0

    deps = result.get("dependencies", [])
    logger.info("  Result: %d dependencies", len(deps))

    save_result(scenario, "scene_dependency", result, "gemini-pro", elapsed)
    return result


# ── 비교 리포트 ──

def compare_results(scenario: str):
    """GPT(기존) vs Gemini(새) 비교 리포트."""
    gemini_dir = COMP_DIR / f"{scenario}_gemini"
    if not gemini_dir.exists():
        logger.warning("Gemini results not found for %s", scenario)
        return

    report_lines = [f"\n{'='*60}", f"  비교 리포트: {scenario}", f"{'='*60}\n"]

    steps = [
        ("visual_world_rules", "시각적 세계관 규칙"),
        ("entity_extract_character", "인물 추출"),
        ("entity_extract_location", "배경 추출"),
        ("entity_extract_prop", "소품 추출"),
        ("scene_cinematography", "촬영 감독"),
        ("scene_dependency", "씬 연관 분석"),
    ]

    for step_id, label in steps:
        gpt_path = COMP_DIR / scenario / f"{step_id}.json"
        gemini_path = gemini_dir / f"{step_id}.json"

        if not gpt_path.exists() or not gemini_path.exists():
            report_lines.append(f"  [{label}] 데이터 없음\n")
            continue

        gpt_data = json.loads(gpt_path.read_text(encoding="utf-8"))
        gemini_data = json.loads(gemini_path.read_text(encoding="utf-8"))

        report_lines.append(f"  [{label}] ({step_id})")
        report_lines.append(f"    GPT model: {gpt_data.get('resolved_model', 'gpt-5.4')}")
        report_lines.append(f"    Gemini elapsed: {gemini_data.get('elapsed_seconds', '?')}s")

        _compare_step(step_id, gpt_data, gemini_data, report_lines)
        report_lines.append("")

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

    # 리포트 저장
    report_path = gemini_dir / "comparison_report.txt"
    report_path.write_text(report, encoding="utf-8")
    logger.info("Report saved: %s", report_path)


def _compare_step(step_id: str, gpt_data: dict, gemini_data: dict, lines: list):
    """단계별 비교."""
    gpt_d = gpt_data.get("data", gpt_data)
    gemini_d = gemini_data.get("data", gemini_data)

    if step_id == "visual_world_rules":
        gpt_rules = gpt_d.get("rules", [])
        gem_rules = gemini_d.get("rules", [])
        gpt_notes = gpt_d.get("director_notes", [])
        gem_notes = gemini_d.get("director_notes", [])
        lines.append(f"    GPT: {len(gpt_rules)} rules, {len(gpt_notes)} director_notes, era={gpt_d.get('era','?')}")
        lines.append(f"    Gem: {len(gem_rules)} rules, {len(gem_notes)} director_notes, era={gemini_d.get('era','?')}")

    elif step_id.startswith("entity_extract_"):
        etype = step_id.split("_")[-1]
        key = f"{etype}s" if etype != "prop" else "props"
        # GPT data is inside data.characters/locations/props
        gpt_entities = gpt_d.get(key, gpt_d.get("entities", []))
        gem_entities = gemini_d if isinstance(gemini_d, list) else gemini_d.get(key, [])
        lines.append(f"    GPT: {len(gpt_entities)} {etype}s")
        lines.append(f"    Gem: {len(gem_entities)} {etype}s")
        # 이름 비교
        gpt_names = {e.get("name", "?") for e in gpt_entities}
        gem_names = {e.get("name", "?") for e in gem_entities}
        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 == "scene_cinematography":
        gpt_scenes = gpt_d.get("scenes", [])
        gem_scenes = gemini_d.get("scenes", [])
        lines.append(f"    GPT: {len(gpt_scenes)} scenes")
        lines.append(f"    Gem: {len(gem_scenes)} scenes")
        # shot type 차이 비교
        diffs = 0
        for gs, gms in zip(gpt_scenes, gem_scenes):
            if gs.get("shot_1") != gms.get("shot_1") or gs.get("shot_2") != gms.get("shot_2"):
                diffs += 1
        lines.append(f"    Shot type 차이: {diffs}/{min(len(gpt_scenes), len(gem_scenes))} scenes")

    elif step_id == "scene_dependency":
        gpt_deps = gpt_d.get("dependencies", [])
        gem_deps = gemini_d.get("dependencies", [])
        lines.append(f"    GPT: {len(gpt_deps)} dependencies")
        lines.append(f"    Gem: {len(gem_deps)} dependencies")


def main():
    parser = argparse.ArgumentParser(description="GPT vs Gemini Pro model comparison")
    parser.add_argument("--scenario", required=True, choices=["srd_ep1", "mia_ep1"],
                        help="시나리오 선택")
    parser.add_argument("--step", default="all",
                        choices=["all", "visual_world_rules",
                                 "entity_extract_character", "entity_extract_location",
                                 "entity_extract_prop", "scene_cinematography",
                                 "scene_dependency", "compare"],
                        help="실행할 단계 (기본: all)")
    args = parser.parse_args()

    # gemini-pro 모델 오버라이드
    project_config = {
        "visual_world_rules": {"model": "gemini-pro"},
        "entity_extract": {"model": "gemini-pro"},
        "scene_cinematography": {"model": "gemini-pro"},
        "scene_dependency": {"model": "gemini-pro"},
    }

    scenario = args.scenario
    step = args.step

    logger.info("Starting comparison: scenario=%s, step=%s", scenario, step)

    if step == "compare":
        compare_results(scenario)
        return

    if step in ("all", "visual_world_rules"):
        run_visual_world_rules(scenario, project_config)

    if step in ("all", "entity_extract_character"):
        run_entity_extract(scenario, "character", project_config)

    if step in ("all", "entity_extract_location"):
        run_entity_extract(scenario, "location", project_config)

    if step in ("all", "entity_extract_prop"):
        run_entity_extract(scenario, "prop", project_config)

    if step in ("all", "scene_cinematography"):
        run_scene_cinematography(scenario, project_config)

    if step in ("all", "scene_dependency"):
        run_scene_dependency(scenario, project_config)

    # 비교 리포트
    if step == "all":
        compare_results(scenario)


if __name__ == "__main__":
    main()
