"""Beat 추출 실험 — 원본 씬에서 beat 추출 (5000자 묶음 방식).

Usage:
    cd backend
    .venv/bin/python experiments/beat_extract_test.py \
        --project-id 84e9e90b-88a5-4539-8d93-4404145369fa \
        --episode-id 78bc78d1-7d73-4163-9fc1-a4ab93ec3b7c \
        --model gpt-mini
        --threads 4
"""
import argparse
import json
import sys
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path

# backend를 sys.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   # 분석 대상 씬 최대 글자수
REF_MAX = 2000         # 앞쪽 참조 씬 최대 글자수

BEAT_SYSTEM = """당신은 시나리오 분석 전문가입니다. 주어진 씬들에서 Beat를 추출합니다."""

BEAT_USER_TEMPLATE = """\
다음 시나리오를 분석하여 Scene별로 "Beat"를 추출하라.

[Beat 정의]
Beat는 씬 안에서 인물의 "상태(state)"가 변화하는 최소 단위이다.

상태 변화는 다음 중 하나 이상이 바뀔 때 발생한다:
1. 행동 변화 (action)
2. 반응 변화 (reaction)
3. 감정 변화 (emotion)
4. 정보/인지 변화 (information)
5. 의도/결심 변화 (decision)
6. 관계 변화 (relationship)
7. 상황/환경 변화 (situation)

[중요 규칙]
- 각 Beat는 반드시 "이전 상태 → 이후 상태"의 변화가 명확해야 한다
- 변화가 없는 지속 상태는 Beat로 만들지 않는다
- 하나의 Beat에는 하나의 핵심 변화만 포함한다
- 단순한 대사 pause, 리듬상의 쉼은 Beat로 간주하지 않는다
- 시간 순서를 유지한다

{ref_section}
[분석 대상 씬]
{bundle_text}
"""

BEAT_SCHEMA = {
    "type": "object",
    "properties": {
        "scenes": {
            "type": "array",
            "items": {
                "type": "object",
                "properties": {
                    "scene_index": {"type": "integer"},
                    "scene_heading": {"type": "string"},
                    "beats": {
                        "type": "array",
                        "items": {
                            "type": "object",
                            "properties": {
                                "beat_index": {"type": "integer"},
                                "change_type": {
                                    "type": "string",
                                    "enum": ["action", "reaction", "emotion",
                                             "information", "decision",
                                             "relationship", "situation"]
                                },
                                "before_state": {"type": "string"},
                                "after_state": {"type": "string"},
                                "description": {"type": "string"},
                            },
                            "required": ["beat_index", "change_type",
                                         "before_state", "after_state",
                                         "description"],
                            "additionalProperties": False,
                        },
                    },
                },
                "required": ["scene_index", "scene_heading", "beats"],
                "additionalProperties": False,
            },
        },
    },
    "required": ["scenes"],
    "additionalProperties": False,
}


def load_scenes(project_dir: Path, episode_id: str):
    """scene_segmentation 체크포인트에서 원본 씬 로드 + PDF에서 fulltext 추출."""
    cp_path = (project_dir / "checkpoints" / "episodes" / episode_id
               / "scene_segmentation" / "manifest.json")
    cp = json.loads(cp_path.read_text())
    segments = cp["data"]["segments"]

    # PDF에서 fulltext
    pdf_dir = project_dir / "assets" / "screenplays"
    pdfs = list(pdf_dir.glob("*.pdf"))
    if not pdfs:
        raise FileNotFoundError(f"No PDF in {pdf_dir}")
    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, fulltext


def _build_bundles(scene_texts):
    """씬을 BUNDLE_TARGET 이하로 묶어 번들 리스트 반환."""
    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:
            bundle.append(scene_texts[i])
            bundle_len += scene_texts[i]["length"]
            i += 1
        if not bundle and i < len(scene_texts):
            bundle.append(scene_texts[i])
            bundle_len = scene_texts[i]["length"]
            i += 1

        # 앞쪽 참조 (REF_MAX 이하)
        ref_parts = []
        ref_len = 0
        j = bundle[0]["idx"] - 2
        while j >= 0 and j < len(scene_texts) and ref_len + scene_texts[j]["length"] <= REF_MAX:
            ref_parts.insert(0, scene_texts[j]["text"])
            ref_len += scene_texts[j]["length"]
            j -= 1

        bundles.append({
            "scenes": bundle,
            "bundle_len": bundle_len,
            "ref_parts": ref_parts,
        })
    return bundles


def _call_one_bundle(bundle_info, call_idx, model):
    """단일 번들의 beat 추출 (스레드에서 실행)."""
    bundle = bundle_info["scenes"]
    bundle_len = bundle_info["bundle_len"]
    ref_parts = bundle_info["ref_parts"]

    ref_section = ""
    if ref_parts:
        ref_section = f"[앞쪽 씬 — 참조만, 이 씬들의 beat는 추출하지 마세요]\n{''.join(r['text'] if isinstance(r, dict) else r for r in ref_parts)}\n\n"

    bundle_text = "\n".join(
        f"--- Scene {s['idx']}: {s['heading']} ---\n{s['text']}"
        for s in bundle
    )

    user_prompt = BEAT_USER_TEMPLATE.format(
        ref_section=ref_section,
        bundle_text=bundle_text,
    )

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

    try:
        result = call_structured(
            step="beat_extract_test",
            system_prompt=BEAT_SYSTEM,
            user_prompt=user_prompt,
            response_schema=BEAT_SCHEMA,
            project_config={"default_model": model},
            schema_name=f"beat_extract_{call_idx}",
        )
        elapsed = time.time() - call_start
        scenes_result = result.get("scenes", [])
        total_beats = sum(len(s.get("beats", [])) 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_beats} beats", 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_beats(scene_texts, model: str, max_threads: int = 1):
    """씬 묶음으로 beat 추출. max_threads > 1이면 병렬."""
    bundles = _build_bundles(scene_texts)
    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_beats = sum(len(s.get("beats", [])) for s in all_results)
    failed = sum(1 for v in results_by_idx.values() if not v)
    print(f"\n  Total: {len(all_results)} scenes, {total_beats} beats, "
          f"{len(bundles)} calls, {total_elapsed:.1f}s"
          f"{f' ({failed} failed)' if failed else ''}", 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("--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, fulltext = load_scenes(project_dir, args.episode_id)
    print(f"  {len(scene_texts)} original scenes, fulltext {len(fulltext)} chars")

    print(f"\nExtracting beats with model={args.model}, threads={args.threads}...")
    results = extract_beats(scene_texts, args.model, max_threads=args.threads)

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

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


if __name__ == "__main__":
    main()
