"""scene_director 결과 검증 스크립트.

각 씬에 대해 (앞 2씬 + 현재 씬) 컨텍스트와 전체 엔티티 목록을 제공하고,
Gemini Pro / GPT 5.4 양쪽에 현재 scene_director 결과가 맞는지 물어봄.

Usage:
    cd backend
    .venv/bin/python scripts/verify_scene_director.py <project_id> <episode_id> [--scenes 11,29,30]
"""

import argparse
import json
import logging
import sys
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path

# backend를 PYTHONPATH에 추가
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from app.core.config import settings
from app.modules.llm.llm_client import call_text

logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s")
logger = logging.getLogger(__name__)

SYSTEM_PROMPT = """너는 시나리오 분석 전문가다. 아래 씬 텍스트를 읽고, 해당 씬에서 **물리적으로 존재하여 카메라에 보이는** 엔티티(인물, 배경, 소품)를 판별하라.

규칙:
- "물리적 존재" = 해당 씬에서 실제로 그 장소에 있어 촬영 가능한 대상
- 대화 속 언급, 회상, 상상 속 인물은 제외
- 빙의/변신 등으로 육체가 존재하면 포함
- 엔티티 목록에 없는 인물(단역, 엑스트라)은 무시

현재 AI가 판별한 결과가 제시된다. 이 결과가 맞는지 검증하고:
1. 누락된 엔티티가 있다면 short_id와 이유를 제시
2. 잘못 포함된 엔티티가 있다면 short_id와 이유를 제시
3. 모두 맞다면 "정확함"이라고 답변

반드시 한국어로 답변하라."""


def load_data(project_id: str, episode_id: str):
    """체크포인트 + 원문 로드."""
    base = Path(settings.projects_dir) / project_id / "checkpoints" / "episodes" / episode_id

    def _load(step):
        p = base / step / "manifest.json"
        return json.loads(p.read_text(encoding="utf-8")) if p.exists() else None

    seg_cp = _load("scene_segmentation")
    segments = seg_cp.get("data", {}).get("segments", []) if seg_cp else []

    summary_cp = _load("scene_summary")
    summaries = {
        s["scene_index"]: s.get("scene_summary", "")
        for s in (summary_cp.get("data", {}).get("summaries", []) if summary_cp else [])
    }

    director_cp = _load("scene_director")
    director_scenes = {
        s["scene_index"]: s
        for s in (director_cp.get("data", {}).get("scenes", []) if director_cp else [])
    }

    t2i_cp = _load("entity_t2i")
    entities = t2i_cp.get("data", {}) if t2i_cp else {}

    # 원문 로드
    ep_dir = Path(settings.projects_dir) / project_id / "assets" / "screenplays"
    fulltext = ""
    text_path = base / "scene_segmentation" / "fulltext.txt"
    if text_path.exists():
        fulltext = text_path.read_text(encoding="utf-8")
    else:
        # cleaned_text 체크포인트에서 로드
        cleaned_cp = _load("cleaned_text")
        if cleaned_cp:
            fulltext = cleaned_cp.get("data", {}).get("cleaned_text", "")
        if not fulltext:
            # episode text에서 로드
            from app.core.database import SessionLocal
            from sqlalchemy import text as sql_text
            db = SessionLocal()
            try:
                row = db.execute(sql_text(
                    "SELECT fulltext FROM episode WHERE id = :eid"
                ), {"eid": episode_id}).fetchone()
                if row:
                    fulltext = row[0] or ""
            finally:
                db.close()

    return segments, summaries, director_scenes, entities, fulltext


def build_entity_block(entities: dict) -> str:
    """전체 엔티티 목록 텍스트 구성."""
    lines = []
    for etype in ["characters", "locations", "props"]:
        for e in entities.get(etype, []):
            sid = e.get("short_id", e.get("name", "?"))
            desc = e.get("description", "")
            lines.append(f"- {sid}: {e.get('name', '')} ({etype[:-1]}) — {desc}")
    return "\n".join(lines)


def build_scene_context(
    scene_index: int,
    segments: list,
    summaries: dict,
    director_scenes: dict,
    fulltext: str,
) -> str:
    """앞 2씬 + 현재 씬 컨텍스트 + 현재 director 결과."""
    seg_by_idx = {s.get("scene_index", i + 1): s for i, s in enumerate(segments)}

    context_parts = []
    # 앞 2씬 (요약만)
    for offset in [-2, -1]:
        prev_si = scene_index + offset
        if prev_si >= 1 and prev_si in summaries:
            context_parts.append(f"[씬 {prev_si} 요약] {summaries[prev_si]}")

    # 현재 씬 (전문)
    seg = seg_by_idx.get(scene_index)
    if seg:
        start = seg.get("start_char", 0)
        end = seg.get("end_char", len(fulltext))
        scene_text = fulltext[start:end]
        heading = seg.get("heading", "")
        context_parts.append(f"[씬 {scene_index} 전문] {heading}\n{scene_text}")

    # 현재 director 결과
    ds = director_scenes.get(scene_index, {})
    present = ds.get("present_entity_ids", [])
    context_parts.append(
        f"\n--- AI 판별 결과 ---\n"
        f"씬 {scene_index}에 물리적으로 존재하는 엔티티: {present if present else '없음 (빈 배열)'}"
    )

    return "\n\n".join(context_parts)


def verify_one_scene(
    scene_index: int,
    entity_block: str,
    scene_context: str,
    model_override: str,
    project_config: dict,
) -> dict:
    """한 씬에 대해 특정 모델로 검증."""
    user_prompt = f"엔티티 목록:\n{entity_block}\n\n{scene_context}"

    # model_override로 강제 지정
    config = dict(project_config) if project_config else {}
    config["scene_director_verify"] = {"model": model_override}

    try:
        result = call_text(
            step="scene_director_verify",
            system_prompt=SYSTEM_PROMPT,
            user_prompt=user_prompt,
            project_config=config,
            temperature=0.1,
        )
        return {
            "scene_index": scene_index,
            "model": model_override,
            "response": result.strip(),
        }
    except Exception as exc:
        return {
            "scene_index": scene_index,
            "model": model_override,
            "response": f"ERROR: {exc}",
        }


def main():
    parser = argparse.ArgumentParser(description="scene_director 결과 검증")
    parser.add_argument("project_id")
    parser.add_argument("episode_id")
    parser.add_argument("--scenes", default=None, help="검증할 씬 번호 (콤마 구분, 기본=전체)")
    parser.add_argument("--workers", type=int, default=10)
    args = parser.parse_args()

    segments, summaries, director_scenes, entities, fulltext = load_data(
        args.project_id, args.episode_id
    )
    entity_block = build_entity_block(entities)

    if args.scenes:
        target_scenes = [int(s.strip()) for s in args.scenes.split(",")]
    else:
        target_scenes = sorted(director_scenes.keys())

    logger.info("검증 대상: %d씬, 모델: gemini-pro + gpt", len(target_scenes))

    # 프로젝트 설정 로드
    from app.core.database import SessionLocal
    from sqlalchemy import text as sql_text
    db = SessionLocal()
    try:
        row = db.execute(sql_text(
            "SELECT llm_config_json FROM project_settings WHERE project_id = :pid"
        ), {"pid": args.project_id}).fetchone()
        project_config = json.loads(row[0]) if row and row[0] else {}
    finally:
        db.close()

    # 태스크 구성: 각 씬 × 2 모델
    tasks = []
    for si in target_scenes:
        scene_ctx = build_scene_context(si, segments, summaries, director_scenes, fulltext)
        for model in ["gemini-pro", "gpt"]:
            tasks.append((si, entity_block, scene_ctx, model, project_config))

    logger.info("총 %d 호출 (씬 %d × 모델 2), workers=%d", len(tasks), len(target_scenes), args.workers)

    results = []
    with ThreadPoolExecutor(max_workers=args.workers) as executor:
        futures = {
            executor.submit(verify_one_scene, si, eb, sc, m, pc): (si, m)
            for si, eb, sc, m, pc in tasks
        }
        for future in as_completed(futures):
            si, model = futures[future]
            try:
                r = future.result()
                results.append(r)
                # 간단 출력
                short = r["response"][:120].replace("\n", " ")
                logger.info("씬 %d [%s]: %s...", si, model, short)
            except Exception as exc:
                logger.error("씬 %d [%s] 실패: %s", si, model, exc)

    # 결과 정리 — 씬별로 그룹
    results.sort(key=lambda x: (x["scene_index"], x["model"]))
    print("\n" + "=" * 80)
    print("검증 결과")
    print("=" * 80)

    current_si = None
    for r in results:
        if r["scene_index"] != current_si:
            current_si = r["scene_index"]
            ds = director_scenes.get(current_si, {})
            present = ds.get("present_entity_ids", [])
            print(f"\n{'─' * 60}")
            print(f"씬 {current_si} | 현재 결과: {present if present else '[]'}")
            print(f"{'─' * 60}")
        print(f"\n  [{r['model']}]")
        print(f"  {r['response']}")

    # JSON 저장
    out_path = Path(settings.projects_dir) / args.project_id / "verify_scene_director.json"
    out_path.write_text(json.dumps(results, ensure_ascii=False, indent=2), encoding="utf-8")
    logger.info("결과 저장: %s", out_path)


if __name__ == "__main__":
    main()
