"""Shot 추출 실험 — beat 기반으로 씬별 shot 추출.

gemini-pro beat 결과 + 원본 씬 텍스트 → 3000자 번들 → 씬별 shot 최소 1개.

Usage:
    cd backend
    .venv/bin/python experiments/shot_extract_test.py \
        --project-id 84e9e90b-... --episode-id 78bc78d1-... \
        --beats experiments/beat_results_gemini_pro.json \
        --model gpt-mini --threads 4
"""
import argparse
import json
import sys
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path

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

from app.modules.llm.llm_client import call_structured
from app.modules.pdf_parser import extract_text_from_pdf

BUNDLE_TARGET = 3000

SHOT_SYSTEM = """당신은 시나리오를 스틸컷(shot)으로 분해하는 전문가입니다.
각 씬의 beat(상태 변화) 정보와 원문을 참고하여, 스틸컷으로 찍을 핵심 순간을 추출합니다."""

SHOT_USER_TEMPLATE = """\
아래 씬들의 원문과 beat 분석 결과를 참고하여, 각 씬에서 "Shot"을 추출하라.

[Shot 정의]
Shot은 하나의 스틸컷(정지 화면)으로 촬영할 수 있는 한 순간이다.
- 한 Shot = 한 시공간, 한 순간
- 각 씬에서 최소 1개 이상의 Shot을 추출한다
- beat의 전환점 중 시각적으로 가장 인상적인 순간을 선택한다
- Shot의 순서는 시나리오 시간 순서를 따른다

[중요 규칙]
- beat가 0개인 씬도 최소 1개의 Shot을 추출한다 (상황 묘사 기반)
- 하나의 Shot에 여러 beat를 합치지 말 것 — 각 Shot은 단일 순간
- 대사 장면보다 행동/반응/상황 변화가 보이는 순간을 우선한다
- 각 Shot에 등장하는 주요 인물을 명시한다
- Shot description은 카메라가 보는 장면을 묘사한다 (감정이 아닌 시각적 묘사)
- **based_on_beat는 반드시 해당 Shot의 근거가 된 Beat의 beat_index를 명시한다. beat 없이 원문에서 직접 추출한 경우에만 0으로 표기한다.**

{scenes_section}
"""

SHOT_SCHEMA = {
    "type": "object",
    "properties": {
        "scenes": {
            "type": "array",
            "items": {
                "type": "object",
                "properties": {
                    "scene_index": {"type": "integer"},
                    "scene_heading": {"type": "string"},
                    "shots": {
                        "type": "array",
                        "items": {
                            "type": "object",
                            "properties": {
                                "shot_index": {"type": "integer"},
                                "description": {"type": "string"},
                                "characters": {
                                    "type": "array",
                                    "items": {"type": "string"},
                                },
                                "based_on_beat": {
                                    "type": "integer",
                                    "description": "이 Shot의 근거가 된 Beat의 beat_index. beat 없이 원문에서 추출한 경우 0."
                                },
                            },
                            "required": ["shot_index", "description", "characters", "based_on_beat"],
                            "additionalProperties": False,
                        },
                    },
                },
                "required": ["scene_index", "scene_heading", "shots"],
                "additionalProperties": False,
            },
        },
    },
    "required": ["scenes"],
    "additionalProperties": False,
}


def load_scenes(project_dir: Path, episode_id: str):
    cp_path = (project_dir / "checkpoints" / "episodes" / episode_id
               / "scene_segmentation" / "manifest.json")
    cp = json.loads(cp_path.read_text())
    segments = cp["data"]["segments"]

    pdf_dir = project_dir / "assets" / "screenplays"
    pdfs = list(pdf_dir.glob("*.pdf"))
    fulltext, _ = extract_text_from_pdf(pdfs[0])

    scene_texts = []
    for seg in segments:
        start = seg["start_char"]
        end = seg["end_char"]
        scene_texts.append({
            "idx": seg["scene_index"],
            "heading": seg.get("heading", ""),
            "text": fulltext[start:end],
            "length": end - start,
        })
    return scene_texts


def load_beats(beats_path: str):
    """beat 결과를 scene_index → beats 딕셔너리로 변환."""
    beats_data = json.loads(Path(beats_path).read_text())
    beats_by_scene = {}
    for s in beats_data:
        idx = s["scene_index"]
        # 같은 scene_index가 여러 번 나올 수 있음 (heading으로 구분)
        key = (idx, s.get("scene_heading", ""))
        beats_by_scene[key] = s.get("beats", [])
    return beats_data, beats_by_scene


def _format_beats_for_scene(beats):
    if not beats:
        return "  (beat 없음 — 상황 묘사 기반으로 Shot 추출)"
    lines = []
    for b in beats:
        lines.append(f"  Beat {b['beat_index']}: [{b['change_type']}] "
                      f"{b['before_state']} → {b['after_state']}")
    return "\n".join(lines)


def _build_bundles(scene_texts, beats_data):
    """씬을 BUNDLE_TARGET 이하로 묶되, beat 정보도 함께."""
    # 순서 기반 1:1 매칭 (beats_data[i] ↔ scene_texts[i])
    beats_by_order = [bd.get("beats", []) for bd in beats_data]
    while len(beats_by_order) < len(scene_texts):
        beats_by_order.append([])

    bundles = []
    i = 0
    while i < len(scene_texts):
        bundle = []
        bundle_len = 0
        while i < len(scene_texts) and bundle_len + scene_texts[i]["length"] <= BUNDLE_TARGET:
            st = scene_texts[i]
            bundle.append({**st, "beats": beats_by_order[i]})
            bundle_len += st["length"]
            i += 1
        if not bundle and i < len(scene_texts):
            st = scene_texts[i]
            bundle.append({**st, "beats": beats_by_order[i]})
            bundle_len = st["length"]
            i += 1
        bundles.append({"scenes": bundle, "bundle_len": bundle_len})
    return bundles


def _call_one_bundle(bundle_info, call_idx, model):
    bundle = bundle_info["scenes"]
    bundle_len = bundle_info["bundle_len"]

    scenes_parts = []
    for s in bundle:
        beat_text = _format_beats_for_scene(s["beats"])
        scenes_parts.append(
            f"--- Scene {s['idx']}: {s['heading']} ---\n"
            f"[Beats]\n{beat_text}\n\n"
            f"[원문]\n{s['text']}"
        )

    scenes_section = "\n\n".join(scenes_parts)
    user_prompt = SHOT_USER_TEMPLATE.format(scenes_section=scenes_section)

    scene_range = f"{bundle[0]['idx']}~{bundle[-1]['idx']}"
    call_start = time.time()

    try:
        result = call_structured(
            step="shot_extract_test",
            system_prompt=SHOT_SYSTEM,
            user_prompt=user_prompt,
            response_schema=SHOT_SCHEMA,
            project_config={"default_model": model},
            schema_name=f"shot_extract_{call_idx}",
        )
        elapsed = time.time() - call_start
        scenes_result = result.get("scenes", [])
        total_shots = sum(len(s.get("shots", [])) for s in scenes_result)
        print(f"  call {call_idx} (scenes {scene_range}, {bundle_len}chars, {elapsed:.1f}s): "
              f"{len(scenes_result)} scenes, {total_shots} shots", flush=True)
        return call_idx, scenes_result, None
    except Exception as exc:
        elapsed = time.time() - call_start
        print(f"  call {call_idx} (scenes {scene_range}) FAILED ({elapsed:.1f}s): {exc}", flush=True)
        return call_idx, [], str(exc)


def extract_shots(scene_texts, beats_data, model: str, max_threads: int = 4):
    bundles = _build_bundles(scene_texts, beats_data)
    print(f"  {len(bundles)} bundles prepared", flush=True)

    total_start = time.time()
    results_by_idx = {}

    if max_threads <= 1:
        for i, b in enumerate(bundles):
            idx, scenes, err = _call_one_bundle(b, i + 1, model)
            results_by_idx[idx] = scenes
    else:
        with ThreadPoolExecutor(max_workers=max_threads) as pool:
            futures = {
                pool.submit(_call_one_bundle, b, i + 1, model): i + 1
                for i, b in enumerate(bundles)
            }
            for fut in as_completed(futures):
                idx, scenes, err = fut.result()
                results_by_idx[idx] = scenes

    all_results = []
    for i in sorted(results_by_idx):
        all_results.extend(results_by_idx[i])

    total_elapsed = time.time() - total_start
    total_shots = sum(len(s.get("shots", [])) for s in all_results)
    print(f"\n  Total: {len(all_results)} scenes, {total_shots} shots, "
          f"{len(bundles)} calls, {total_elapsed:.1f}s", flush=True)
    return all_results


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--project-id", required=True)
    parser.add_argument("--episode-id", required=True)
    parser.add_argument("--beats", required=True, help="beat 결과 JSON 경로")
    parser.add_argument("--model", default="gpt-mini")
    parser.add_argument("--output", default=None)
    parser.add_argument("--threads", type=int, default=4)
    args = parser.parse_args()

    projects_root = Path(__file__).resolve().parent.parent.parent / "projects"
    project_dir = projects_root / args.project_id

    print(f"Loading scenes from {args.project_id}...")
    scene_texts = load_scenes(project_dir, args.episode_id)
    print(f"  {len(scene_texts)} original scenes")

    print(f"Loading beats from {args.beats}...")
    beats_data, beats_by_scene = load_beats(args.beats)
    print(f"  {len(beats_data)} scene beat records")

    print(f"\nExtracting shots with model={args.model}, threads={args.threads}...")
    results = extract_shots(scene_texts, beats_data, args.model, args.threads)

    out_path = args.output or f"experiments/shot_results_{args.model}.json"
    Path(out_path).parent.mkdir(parents=True, exist_ok=True)
    Path(out_path).write_text(json.dumps(results, ensure_ascii=False, indent=2))
    print(f"\nSaved to {out_path}")

    print("\n=== Summary ===")
    for s in results:
        shots = s.get("shots", [])
        print(f"  Scene {s['scene_index']:3d} ({s.get('scene_heading','')[:30]:30s}): {len(shots)} shots")


if __name__ == "__main__":
    main()
