"""BackgroundClassifyStep — Phase 7 Step 1.

Spec D1: building group clustering + chain_bg/prev_shot_ref classification — single LLM call.

흐름:
  1. shot_validator/shot_selection/entity_merge/entity_detail/visual_world_rules 로드
  2. shot_count(loc_id별 출현 횟수) 집계
  3. _build_locations_from_entities — entity_merge.locations + entity_detail.kind/summary
     를 flat list로 만듦 (pre-grouping 없음 — LLM이 cluster 결정)
  4. build_classify_user_prompt → run_background_classify (3회 retry)

체크포인트 data:
  data.building_groups[] = [{group_id, members[], anchor_loc, kind, rationale}]
"""
from __future__ import annotations

import json
import logging
from collections import defaultdict
from pathlib import Path
from typing import Any, Dict, List, Optional

from app.core.step_runner import StepRunner

logger = logging.getLogger(__name__)

SCHEMA_VERSION = 1
# C4 v1 — v4 prompt: indoor/outdoor 판단의 닫힌 한국어 keyword 목록을 generic
# semantic principle 로 추상화 (Prompt Closed-List Ban). config_hash invalidation.
# SCHEMA_VERSION 1 유지.
PROMPT_VERSION = "4"


class BackgroundClassifyStep(StepRunner):
    def _config_hash(self) -> str:
        import hashlib, json as _json
        from app.core.config import settings
        payload = {
            "background_mode": settings.background_mode,
            "model": settings.openai_model,
            "schema_version": SCHEMA_VERSION,
            "prompt_version": PROMPT_VERSION,
        }
        return hashlib.sha256(
            _json.dumps(payload, sort_keys=True).encode("utf-8")
        ).hexdigest()[:16]

    def _load_prev_checkpoint(self, step_id: str) -> Optional[Dict[str, Any]]:
        from app.core.config import settings
        cp = (
            Path(settings.projects_dir) / self.project_id
            / "checkpoints" / "episodes" / self.episode_id
            / step_id / "manifest.json"
        )
        if cp.exists():
            try:
                return json.loads(cp.read_text(encoding="utf-8"))
            except Exception as exc:
                logger.warning("background_classify: %s parse failed: %s", step_id, exc)
        return None

    def _execute(self, mode: str = "resume") -> Dict[str, Any]:
        from app.core.config import settings
        if settings.background_mode not in {"on", "floor_plan_anchored"}:
            logger.info(
                "background_classify: skipped — background_mode=%s",
                settings.background_mode,
            )
            return {
                "applicable_count": 0,
                "completed_count": 0,
                "failed_count": 0,
                "schema_version": SCHEMA_VERSION,
                "config_hash": self._config_hash(),
                "data": {},
            }

        shot_validator_cp = self._load_prev_checkpoint("shot_validator")
        shot_selection_cp = self._load_prev_checkpoint("shot_selection")
        entity_merge_cp = self._load_prev_checkpoint("entity_merge")
        entity_detail_cp = self._load_prev_checkpoint("entity_detail")
        rules_cp = self._load_prev_checkpoint("visual_world_rules")
        scene_director_cp = self._load_prev_checkpoint("scene_director")

        # scene → primary_location fallback (shot에 location_id 없는 경우, Phase 5 패턴)
        scene_primary: Dict[int, str] = {}
        if scene_director_cp:
            for sc in (scene_director_cp.get("data", {}) or {}).get("scenes", []) or []:
                si = sc.get("scene_index")
                primary = sc.get("primary_location", "") or ""
                if si is not None and primary:
                    scene_primary[int(si)] = primary

        # location별 shot count (selected shots만, scene primary fallback)
        shot_counts = _count_selected_shots_by_loc(
            shot_validator_cp, shot_selection_cp, scene_primary,
        )
        locations = _build_locations_from_entities(
            entity_merge_cp, entity_detail_cp, shot_counts
        )

        if not locations:
            logger.info("background_classify: no locations — empty plan")
            return {
                "applicable_count": 1,
                "completed_count": 1,
                "failed_count": 0,
                "schema_version": SCHEMA_VERSION,
                "config_hash": self._config_hash(),
                "data": {"building_groups": []},
            }

        rules_text = ""
        if rules_cp:
            data = rules_cp.get("data", {}) or {}
            rules_text = data.get("rules_text", "") or data.get("text", "") or ""

        from app.modules.pipeline.background_classify import (
            build_classify_user_prompt, run_background_classify, ClassifyError,
        )
        from app.modules.llm.llm_client import call_structured

        all_loc_ids = sorted({loc["loc_id"] for loc in locations})
        # indoor_loc_ids: 휴리스틱 판단 불가능한 케이스(entity_detail에 kind 없음)가
        # 많아 LLM 응답의 is_indoor를 신뢰. None 전달 → validator가 LLM member 사용.

        user_prompt = build_classify_user_prompt(locations, rules_text)

        try:
            result = run_background_classify(
                user_prompt=user_prompt,
                all_loc_ids=all_loc_ids,
                indoor_loc_ids=None,
                shot_counts=shot_counts,
                call_structured_fn=call_structured,
                project_config=self.project_config,
                opik_metadata=self.build_opik_metadata(),
            )
        except ClassifyError as exc:
            logger.error("background_classify: exhausted: %s", exc)
            return {
                "applicable_count": 1,
                "completed_count": 0,
                "failed_count": 1,
                "schema_version": SCHEMA_VERSION,
                "config_hash": self._config_hash(),
                "data": {"error": str(exc), "building_groups": []},
            }

        # LLM이 클러스터링 + 분류를 모두 결정 — 출력을 그대로 보존.
        # 각 그룹의 members는 LLM이 echo한 입력 필드(loc_id/label/shot_count/is_indoor)를
        # 사용. 추가로 entity_detail의 summary가 있으면 합쳐 downstream에 전달.
        summary_by_loc = {loc["loc_id"]: loc.get("summary", "") for loc in locations}
        out_groups: List[Dict[str, Any]] = []
        for cls in result.get("building_groups", []):
            gid = cls["group_id"]
            members_out: List[Dict[str, Any]] = []
            for m in cls.get("members", []) or []:
                loc_id = m["loc_id"]
                members_out.append({
                    "loc_id": loc_id,
                    "label": m.get("label", "") or "",
                    "shot_count": int(m.get("shot_count", 0) or 0),
                    "is_indoor": bool(m.get("is_indoor", False)),
                    "summary": summary_by_loc.get(loc_id, ""),
                })
            out_groups.append({
                "group_id": gid,
                "members": members_out,
                "anchor_loc": cls["anchor_loc"],
                "kind": cls["kind"],
                "rationale": cls.get("rationale", ""),
            })

        return {
            "applicable_count": 1,
            "completed_count": 1,
            "failed_count": 0,
            "schema_version": SCHEMA_VERSION,
            "config_hash": self._config_hash(),
            "data": {"building_groups": out_groups},
        }


def _count_selected_shots_by_loc(
    shot_validator_cp: Optional[Dict[str, Any]],
    shot_selection_cp: Optional[Dict[str, Any]],
    scene_primary: Optional[Dict[int, str]] = None,
) -> Dict[str, int]:
    """Selected shots를 loc_id별로 집계.

    shot에 location_id가 있으면 우선, 없으면 scene_primary[scene_index]로 fallback
    (Phase 5 background_planner와 동일 패턴). 둘 다 없으면 카운트하지 않음.
    """
    if not shot_validator_cp or not shot_selection_cp:
        return {}
    sel_map: Dict[int, set] = {}
    for s in (shot_selection_cp.get("data", {}) or {}).get("scenes", []) or []:
        si = s.get("scene_index")
        if si is None:
            continue
        sel_map[int(si)] = set(s.get("selected_shot_indices", []) or [])

    primary_map = scene_primary or {}
    counts: Dict[str, int] = defaultdict(int)
    for s in (shot_validator_cp.get("data", {}) or {}).get("scenes", []) or []:
        si = s.get("scene_index")
        if si is None:
            continue
        sel = sel_map.get(int(si), set())
        for sh in s.get("shots", []) or []:
            shi = sh.get("shot_index")
            if shi is None or shi not in sel:
                continue
            loc_id = sh.get("location_id") or primary_map.get(int(si), "") or ""
            if loc_id:
                counts[loc_id] += 1
    return dict(counts)


def _build_locations_from_entities(
    entity_merge_cp: Optional[Dict[str, Any]],
    entity_detail_cp: Optional[Dict[str, Any]],
    shot_counts: Dict[str, int],
) -> List[Dict[str, Any]]:
    """entity_merge.locations + entity_detail.kind/summary을 flat list로 변환.

    LLM이 직접 cluster + classify할 수 있도록 pre-grouping 없이 raw list 반환.

    Returns:
        [{loc_id, label, shot_count, is_indoor, summary}, ...]
    """
    if not entity_merge_cp:
        return []
    locations = (entity_merge_cp.get("data", {}) or {}).get("locations", []) or []

    # entity_merge.locations에 description/visual_traits 모두 있으므로 직접 사용.
    # entity_detail의 entity_details dict는 location name별 dict (key='name:location')라
    # short_id 조인이 어렵고, description도 entity_merge와 중복이라 무시.
    # is_indoor는 휴리스틱하게 미리 set하지 않는다 (entity_detail에 kind 정보가
    # 없는 경우가 많아 false positive 발생). LLM에 description 전체를 넘겨 자율 판단.

    out: List[Dict[str, Any]] = []
    for loc in locations:
        sid = loc.get("short_id") or ""
        if not sid:
            continue
        # description + visual_traits 묶어서 LLM 입력에 충분한 컨텍스트 제공
        desc = loc.get("description", "") or ""
        vts = loc.get("visual_traits", []) or []
        if isinstance(vts, list):
            vt_text = "; ".join(str(v) for v in vts)
        else:
            vt_text = str(vts)
        summary = desc
        if vt_text:
            summary = f"{desc} (특징: {vt_text})" if desc else f"특징: {vt_text}"
        out.append({
            "loc_id": sid,
            "label": loc.get("name") or sid,
            "shot_count": int(shot_counts.get(sid, 0)),
            "is_indoor": False,  # placeholder — LLM이 응답에서 결정
            "summary": summary,
        })
    # deterministic order by loc_id
    out.sort(key=lambda x: x["loc_id"])
    return out
