"""ReferencePipelineOrchestrator — Phase 1/2/3 참조 이미지 생성 orchestration.

Phase 3b.2 ~ F24.3: reference_image_service.py의 ``generate_reference_images_only``
(666 LOC) 본체를 literal lift. facade ``ReferenceImageService``는 본 서비스의
``run(...)``에 얇게 위임한다.

이후 F24.4.2~4에서 Phase1/2/3 서비스로 분해될 예정 (현재는 단일 run 메서드).
"""
from __future__ import annotations

import hashlib
import json
import logging
import uuid
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Dict, List, Optional

from sqlalchemy.orm import Session as OrmSession

from app.core.config import settings
from app.core.entity_protection import (
    _collect_reference_required_ids,
    compute_variant_pole_ids,
    should_skip_low_freq,
)
from app.core.errors import AppError
from app.i18n.loader import t
from app.logging.activity_logger import ActivityLogger
from app.models.project import (
    CharacterOutlook,
    Episode,
    EntityCanon,
    EntityEpisodeLink,
    ImageAsset,
    RelationFact,
    RelationParticipant,
    SceneStill,
    WorldGuide,
)
from app.modules.entity_dependency import (
    build_visual_dependency_graph,
    topological_sort_entities,
)
from app.modules.llm.gemini_image_client import GeminiImageClient
from app.modules.llm.gemini_key_pool import key_count as gemini_key_count
from app.modules.llm.openai_client import OpenAIClient
from app.modules.progress_tracker import ProgressTracker
from app.modules.world_guide_generator import WorldGuideGenerator
from app.services.reference_phase1_service import ReferencePhase1Service
from app.services.reference_phase2_service import ReferencePhase2Service
from app.services.reference_phase3_service import ReferencePhase3Service
from app.services.reference_pipeline_context import ReferencePipelineContext

logger = logging.getLogger(__name__)


def _now() -> str:
    return datetime.now(timezone.utc).isoformat()


def _new_id() -> str:
    return str(uuid.uuid4())


class ReferencePipelineOrchestrator:
    """Phase 1/2/3 참조 이미지 생성 orchestrator (F24.4)."""

    def __init__(
        self,
        db: OrmSession,
        project_id: str,
        actor_id: str,
        activity_logger: ActivityLogger,
    ) -> None:
        self._db = db
        self._project_id = project_id
        self._actor_id = actor_id
        self._logger = activity_logger
        self._phase1_svc = ReferencePhase1Service(
            db=db,
            project_id=project_id,
            actor_id=actor_id,
            activity_logger=activity_logger,
        )
        self._phase2_svc = ReferencePhase2Service(
            db=db,
            project_id=project_id,
            actor_id=actor_id,
            activity_logger=activity_logger,
        )
        self._phase3_svc = ReferencePhase3Service(
            db=db,
            project_id=project_id,
            actor_id=actor_id,
            activity_logger=activity_logger,
        )

    # ------------------------------------------------------------------
    # Internal helpers
    # ------------------------------------------------------------------

    def _get_episode(self, episode_id: str) -> Episode:
        ep = (
            self._db.query(Episode)
            .filter(Episode.id == episode_id, Episode.project_id == self._project_id)
            .first()
        )
        if not ep:
            raise AppError(
                code="episode.not_found",
                message=t("episode.not_found"),
                status_code=404,
            )
        return ep

    # ------------------------------------------------------------------
    # Public — 참조 이미지 orchestration
    # ------------------------------------------------------------------

    def run(
        self,
        episode_id: str,
        ip: Optional[str] = None,
        mode: str = "resume",
        *,
        skip_composite: bool = False,
    ) -> Dict[str, Any]:
        """엔티티 참조 이미지 생성 orchestrator — Phase 1 (base ref) + Phase 2/3 (composite).

        Phase 3b.4: 신규 호출자는 `generate_base_references` / `generate_composites`
        public API를 사용할 것. 본 함수는 기존 호출자(api/v1/images.py 및
        backend/scripts/experiment_*.py) 호환을 위해 `skip_composite` 플래그 시그니처로
        유지된다.

        mode: "resume" 만 valid. "full" 모드는 feedback_never_delete_images
              규칙 위반 (ImageAsset 일괄 DELETE + cp.clear) 으로 영구 차단.
              강제 재생성은 RefImageGenStep mode='force' 사용.

        skip_composite (Phase 3.1 step 경계 재정의):
            True → Phase 1(base reference)까지만 수행, Phase 2(character+outlook composite) 생략.
            RefImageGenStep이 이 플래그로 호출하여 자기 책임(ref)만 수행.
            CompositeImageGenStep은 기본값(False)로 호출하여 resume 모드에서 이미 완료된
            reference는 skip되고 composite만 생성됨.
        """
        # feedback_never_delete_images: image step 재실행은 절대 일괄 DELETE 금지.
        # mode='full' 의 destructive 동작 (ImageAsset DELETE + ref_cp.clear) 영구 차단.
        # early validation — db/episode 의존성 lookup 전에 fail-fast.
        if mode not in ("resume",):
            raise ValueError(
                f"ReferencePipelineOrchestrator.run: mode={mode!r} unsupported. "
                "Only 'resume' is allowed (feedback_never_delete_images). "
                "Use RefImageGenStep mode='force' for non-destructive regeneration."
            )

        episode = self._get_episode(episode_id)
        if not episode.fulltext:
            raise AppError(code="analysis.no_text", message=t("analysis.no_text"), status_code=400)

        if episode.status != "analyzed":
            raise AppError(
                code="image.analysis_not_complete",
                message=f"분석이 완료되지 않았습니다 (현재: {episode.status}). 분석을 먼저 완료하세요.",
                status_code=400,
            )
        entity_count = self._db.query(EntityCanon).filter(
            EntityCanon.project_id == self._project_id,
        ).count()
        if entity_count == 0:
            raise AppError(
                code="image.no_entities",
                message="추출된 요소가 없습니다. 분석을 먼저 실행하세요.",
                status_code=400,
            )
        if not settings.gemini_api_key and gemini_key_count() == 0:
            raise AppError(code="image.gemini_key_missing", message=t("image.gemini_key_missing"), status_code=400)

        language = episode.language or "ko"
        project_dir = Path(settings.projects_dir) / self._project_id
        reference_dir = project_dir / "images" / episode_id / "reference"

        links = self._db.query(EntityEpisodeLink).filter(
            EntityEpisodeLink.project_id == self._project_id,
            EntityEpisodeLink.episode_id == episode_id,
        ).all()
        canon_ids = [link.canon_id for link in links]
        entities_orm = self._db.query(EntityCanon).filter(
            EntityCanon.id.in_(canon_ids)
        ).all() if canon_ids else []
        entities = [
            {"id": e.id, "name": e.name, "entity_type": e.entity_type,
             "short_id": e.short_id or "",  # W21B-W7 W-B: prop anchor overlay 조인 키
             "description": e.description or "", "stable_traits": e.stable_traits or "{}",
             "t2i_prompt": e.t2i_prompt or ""}
            for e in entities_orm
        ]

        ref_image_map: Dict[str, bytes] = {}
        already_done: set = set()
        if mode == "resume":
            for entity in entities:
                existing = (
                    self._db.query(ImageAsset)
                    .filter(
                        ImageAsset.project_id == self._project_id,
                        ImageAsset.entity_id == entity["id"],
                        ImageAsset.asset_type == "reference",
                        ImageAsset.is_primary == 1,
                    )
                    .order_by(ImageAsset.created_at.desc())
                    .first()
                )
                if existing:
                    fp = Path(existing.file_path)
                    if fp.exists():
                        already_done.add(entity["id"])
                        ref_image_map[entity["id"]] = fp.read_bytes()
                    else:
                        logger.warning("Ref image file missing, will regenerate: %s → %s", entity["name"], existing.file_path)

        openai_client = OpenAIClient()
        stills_orm = self._db.query(SceneStill).filter(
            SceneStill.project_id == self._project_id,
            SceneStill.episode_id == episode_id,
            SceneStill.is_selected == True,      # noqa: E712
            SceneStill.still_index >= 0,
            SceneStill.status != "stale",
        ).all()
        stills = [{"still_frame_prompt": s.still_frame_prompt or ""} for s in stills_orm]
        _ft = episode.fulltext or ""
        wg_hash = hashlib.md5(f"{_ft}:{len(entities)}:{len(stills)}".encode()).hexdigest()

        existing_wg = self._db.query(WorldGuide).filter(
            WorldGuide.project_id == self._project_id,
            WorldGuide.episode_id == episode_id,
        ).order_by(WorldGuide.created_at.desc()).first()

        if existing_wg and existing_wg.source_hash == wg_hash:
            world_guide = json.loads(existing_wg.guide_json)
        else:
            wg_gen = WorldGuideGenerator(llm_client=openai_client)
            world_guide = wg_gen.generate(
                fulltext=episode.fulltext, language=language,
                source_file=episode.source_filename or "episode",
                entities=entities, stills=stills,
            )
            wg_record = WorldGuide(
                id=_new_id(), project_id=self._project_id, episode_id=episode_id,
                guide_json=json.dumps(world_guide, ensure_ascii=False),
                source_hash=wg_hash, created_at=_now(),
            )
            self._db.add(wg_record)
            self._db.commit()

        gemini_client = GeminiImageClient(model=settings.gemini_image_model)
        from app.modules.image_checkpoint import ImageCheckpointManager

        cp_dir = project_dir / "checkpoints" / "images" / episode_id
        ref_cp = ImageCheckpointManager(cp_dir, "reference")
        # mode='full' 분기 제거됨 (run() 첫 줄 ValueError 가드) — 단일 resume 경로.
        for cp_id in ref_cp.get_completed_ids():
            cp_data = ref_cp._data.get("completed", {}).get(cp_id, {})
            cp_path = cp_data.get("file_path", "")
            if cp_path and Path(cp_path).exists():
                already_done.add(cp_id)
            else:
                logger.warning("Checkpoint ref file missing, will regenerate: %s", cp_id[:8])

        relations_orm = self._db.query(RelationFact).filter(RelationFact.project_id == self._project_id).all()
        rel_ids = [r.id for r in relations_orm]
        participants_orm = self._db.query(RelationParticipant).filter(
            RelationParticipant.relation_id.in_(rel_ids)
        ).all() if rel_ids else []
        deps = build_visual_dependency_graph(
            entities,
            [{"id": r.id, "relation_family": r.relation_family} for r in relations_orm],
            [{"relation_id": p.relation_id, "canon_id": p.canon_id} for p in participants_orm],
        )
        batches = topological_sort_entities(entities, deps)

        # ── 저빈도 요소 참조이미지 스킵 (G4.6 RC-C: 보호 cascade 추가) ──
        _t2i_count_map: Dict[str, int] = {}
        for _lnk in self._db.query(EntityEpisodeLink).filter(
            EntityEpisodeLink.project_id == self._project_id,
            EntityEpisodeLink.episode_id == episode_id,
        ).all():
            _t2i_count_map[_lnk.canon_id] = _lnk.t2i_appearance_count or 0

        _reverse_dep_ids: set = set()
        for _dep_set in deps.values():
            _reverse_dep_ids.update(_dep_set)

        # Phase 1 — reference 생성 보호 = scene_detail.required_refs 단독 SOT.
        # broad 4-source union (scene_director.present / shot_validator /
        # shot_director.visible / scene_detail.visible_entities) 은 reference
        # 보호에서 제외 — Phase 0 audit 가 차집합 안전성 검증 완료.
        _required_short_ids = _collect_reference_required_ids(
            self._project_id, episode_id,
        )
        _required_canon_uuids: set = set()
        if _required_short_ids:
            _rows = self._db.query(EntityCanon.id).filter(
                EntityCanon.project_id == self._project_id,
                EntityCanon.short_id.in_(_required_short_ids),
            ).all()
            _required_canon_uuids = {r[0] for r in _rows}

        # variant_self 는 character-character + identity/transformation 만 (RO-1
        # 좁힘 — possession prop / location-location 은 variant_self 미적용).
        _variant_pole_uuids = compute_variant_pole_ids(
            entities,
            [{"id": r.id, "relation_family": r.relation_family} for r in relations_orm],
            [{"relation_id": p.relation_id, "canon_id": p.canon_id} for p in participants_orm],
        )

        _low_freq_skip_ids: set = set()
        _skip_decisions: list = []
        for e in entities:
            eid = e["id"]
            etype = e.get("entity_type", "")
            if etype in ("location", "outlook"):
                continue
            count = _t2i_count_map.get(eid, 0)
            is_base_for_variant = eid in _reverse_dep_ids
            is_variant_self = eid in _variant_pole_uuids
            required = eid in _required_canon_uuids
            skipped = should_skip_low_freq(
                e, count, is_base_for_variant, is_variant_self, required,
            )
            if skipped:
                _low_freq_skip_ids.add(eid)
                _reason = "no required_ref + low t2i count + not variant"
            elif required:
                _reason = "protected: scene_detail.required_refs"
            elif count > 1:
                _reason = f"protected: t2i_count={count}"
            elif is_base_for_variant:
                _reason = "protected: base_for_variant"
            elif is_variant_self:
                _reason = "protected: variant_self"
            else:
                _reason = "protected"
            _skip_decisions.append({
                "canon_id": eid,
                "name": e.get("name"),
                "entity_type": etype,
                "skipped": skipped,
                "reason": _reason,
                "t2i_count": count,
            })
            logger.info(
                "Low-freq decision: %s (%s, t2i_count=%d, base_for_variant=%s, "
                "variant_self=%s, required=%s, skipped=%s)",
                e.get("name"), etype, count, is_base_for_variant,
                is_variant_self, required, skipped,
            )

        from app.core.low_freq_skip import save_low_freq_skip_report
        save_low_freq_skip_report(
            self._project_id, episode_id, _skip_decisions,
        )

        entity_by_id = {e["id"]: e for e in entities}
        generated_count = 0
        failed_count = 0
        low_freq_skipped = len(_low_freq_skip_ids)
        skipped_count = len(already_done) + low_freq_skipped
        max_concurrent = settings.max_concurrent_image_gen
        total_to_gen = len(entities) - skipped_count
        progress = ProgressTracker(self._db, episode_id, "reference_image_generation", self._project_id)

        # O00 (Null Outlook) 캐릭터 ID 세트 — 전신 이미지 생성용
        _o00_outlook_ids = {
            e.id for e in self._db.query(EntityCanon).filter(
                EntityCanon.project_id == self._project_id,
                EntityCanon.short_id == "O00",
                EntityCanon.entity_type == "outlook",
            ).all()
        }
        _o00_char_ids = set()
        if _o00_outlook_ids:
            for co in self._db.query(CharacterOutlook).filter(
                CharacterOutlook.project_id == self._project_id,
                CharacterOutlook.outlook_id.in_(_o00_outlook_ids),
            ).all():
                _o00_char_ids.add(co.character_id)

        # ── Phase 1: base reference (ReferencePhase1Service로 위임, F24.4.2) ──
        ctx = ReferencePipelineContext(
            episode_id=episode_id,
            episode=episode,
            mode=mode,
            skip_composite=skip_composite,
            ip=ip,
            entities=entities,
            entity_by_id=entity_by_id,
            deps=deps,
            batches=batches,
            reference_dir=reference_dir,
            gemini_client=gemini_client,
            max_concurrent=max_concurrent,
            progress=progress,
            ref_cp=ref_cp,
            ref_image_map=ref_image_map,
            already_done=already_done,
            low_freq_skip_ids=_low_freq_skip_ids,
            o00_char_ids=_o00_char_ids,
            generated_count=0,
            failed_count=0,
            low_freq_skipped=low_freq_skipped,
            skipped_count=skipped_count,
        )
        early_return = self._phase1_svc.run(ctx)
        if early_return is not None:
            return early_return

        # Phase 1 결과를 local로 sync (Phase 2/3은 F24.4.3/4에서 ctx 전환 예정)
        generated_count = ctx.generated_count
        failed_count = ctx.failed_count

        # ── Phase 1/2 경계: skip_composite 조기 종료 ──
        if skip_composite:
            progress.complete()
            return {
                "generated": generated_count,
                "skipped": skipped_count,
                "low_freq_skipped": low_freq_skipped,
                "failed": failed_count,
                "total": len(entities),
                "outlook_generated": 0,
                "outlook_skipped": 0,
                "skipped_composite_phase": True,
            }

        # ── Phase 2: outfit 단독 (ReferencePhase2Service, F24.4.3) ──
        self._phase2_svc.run(ctx)
        outlook_generated = ctx.outlook_generated
        outlook_skipped = ctx.outlook_skipped

        # ── Phase 3: composite (ReferencePhase3Service, F24.4.4) ──
        self._phase3_svc.run(ctx)
        composite_generated = ctx.composite_generated
        composite_skipped = ctx.composite_skipped

        progress.complete()
        self._logger.log(
            actor_id=self._actor_id, action="episode.generate_reference_images",
            resource_type="episode", resource_id=episode_id,
            project_id=self._project_id,
            detail={
                "generated": generated_count,
                "skipped": skipped_count,
                "low_freq_skipped": low_freq_skipped,
                "outlook_generated": outlook_generated,
                "outlook_skipped": outlook_skipped,
                "composite_generated": composite_generated,
                "composite_skipped": composite_skipped,
            },
            ip_address=ip,
        )
        return {
            "reference_count": generated_count,
            "skipped": skipped_count,
            "low_freq_skipped": low_freq_skipped,
            "outlook_generated": outlook_generated,
            "outlook_skipped": outlook_skipped,
        }
