"""ImageReviewService — 이미지 리뷰/재생성/primary/검증 서비스 (W5 F23 Phase 2).

ImageService facade(image_service.py)에서 review/primary/validation 관련 8개 공개 메서드를
literal lift. facade는 얇은 delegate만 유지.

이관된 메서드:
- get_image, update_review, regenerate_image, regenerate_needs_fix
- set_primary_image
- _create_validator, validate_image, get_validation
"""

import json
import logging
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.errors import AppError
from app.i18n.loader import t
from app.logging.activity_logger import ActivityLogger
from app.models.project import EntityCanon, ImageAsset, SceneStill
from app.modules.image_validator import ImageValidator
from app.services.image_service_helpers import image_to_dict

logger = logging.getLogger(__name__)


class ImageReviewService:
    """이미지 리뷰/재생성/primary/검증 서비스."""

    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

    def get_image(self, image_id: str) -> ImageAsset:
        """Get a single image asset by ID."""
        img = (
            self._db.query(ImageAsset)
            .filter(ImageAsset.id == image_id, ImageAsset.project_id == self._project_id)
            .first()
        )
        if not img:
            raise AppError(
                code="image.not_found",
                message=t("image.not_found"),
                status_code=404,
            )
        return img

    def update_review(
        self,
        image_id: str,
        status: str,
        notes: str,
        ip: Optional[str] = None,
    ) -> ImageAsset:
        """Update review status of an image."""
        if status not in ("approved", "needs_fix"):
            raise AppError(
                code="image.invalid_review_status",
                message=t("image.invalid_review_status"),
                status_code=400,
            )

        img = self.get_image(image_id)
        img.status = status
        img.review_notes = notes
        self._db.commit()

        self._logger.log(
            actor_id=self._actor_id,
            action="image.review",
            resource_type="image",
            resource_id=image_id,
            project_id=self._project_id,
            detail={"status": status, "notes": notes},
            ip_address=ip,
        )

        return img

    def regenerate_image(
        self,
        image_id: str,
        ip: Optional[str] = None,
    ) -> None:
        """Mark an image for regeneration."""
        img = self.get_image(image_id)
        img.status = "regenerating"
        self._db.commit()

        self._logger.log(
            actor_id=self._actor_id,
            action="image.regenerate",
            resource_type="image",
            resource_id=image_id,
            project_id=self._project_id,
            ip_address=ip,
        )

    def regenerate_needs_fix(
        self,
        episode_id: str,
        ip: Optional[str] = None,
    ) -> int:
        """Mark all needs_fix images for an episode for regeneration."""
        images = (
            self._db.query(ImageAsset)
            .filter(
                ImageAsset.project_id == self._project_id,
                ImageAsset.episode_id == episode_id,
                ImageAsset.status == "needs_fix",
            )
            .all()
        )

        if not images:
            raise AppError(
                code="image.no_needs_fix",
                message=t("image.no_needs_fix"),
                status_code=400,
            )

        for img in images:
            img.status = "regenerating"
        self._db.commit()

        self._logger.log(
            actor_id=self._actor_id,
            action="image.regenerate_batch",
            resource_type="episode",
            resource_id=episode_id,
            project_id=self._project_id,
            detail={"count": len(images)},
            ip_address=ip,
        )

        return len(images)

    def set_primary_image(
        self,
        image_id: str,
        ip: Optional[str] = None,
    ) -> Dict[str, Any]:
        """Set an image as the primary image for its entity/still, unsetting others."""
        img = self.get_image(image_id)

        # Unset primary: composite는 같은 composite_key끼리만, 비-composite는 같은 entity끼리
        import re as _re_sp
        prompt = img.prompt_used or ""
        is_composite = prompt.startswith("[composite:") or prompt.startswith("[outlook_id:")
        composite_key = None
        if is_composite:
            # [composite:char_id:outlook_id] 또는 [outlook_id:xxx] 에서 키 추출
            m = _re_sp.match(r'\[(composite:[a-f0-9-]+:[a-f0-9-]+|outlook_id:[a-f0-9-]+)\]', prompt)
            composite_key = m.group(1) if m else None

        if img.entity_id:
            siblings = (
                self._db.query(ImageAsset)
                .filter(
                    ImageAsset.project_id == self._project_id,
                    ImageAsset.entity_id == img.entity_id,
                    ImageAsset.id != image_id,
                )
                .all()
            )
            for sib in siblings:
                sib_prompt = sib.prompt_used or ""
                sib_is_composite = sib_prompt.startswith("[composite:") or sib_prompt.startswith("[outlook_id:")
                if is_composite and sib_is_composite:
                    # composite끼리: 같은 composite_key일 때만 해제
                    if composite_key and composite_key in sib_prompt:
                        sib.is_primary = 0
                elif not is_composite and not sib_is_composite:
                    # 비-composite끼리: 같은 entity
                    sib.is_primary = 0

        if img.still_id:
            siblings = (
                self._db.query(ImageAsset)
                .filter(
                    ImageAsset.project_id == self._project_id,
                    ImageAsset.still_id == img.still_id,
                    ImageAsset.id != image_id,
                )
                .all()
            )
            for sib in siblings:
                sib.is_primary = 0

        img.is_primary = 1
        self._db.commit()

        self._logger.log(
            actor_id=self._actor_id,
            action="image.set_primary",
            resource_type="image",
            resource_id=image_id,
            project_id=self._project_id,
            detail={"entity_id": img.entity_id, "still_id": img.still_id},
            ip_address=ip,
        )

        return image_to_dict(img)

    # ------------------------------------------------------------------
    # Validation
    # ------------------------------------------------------------------

    def _create_validator(self) -> Optional[ImageValidator]:
        from app.core.openai_keys import has_openai_key
        if not has_openai_key():
            logger.info("Skipping image validation -- OpenAI API key not configured.")
            return None
        return ImageValidator()

    def validate_image(
        self,
        image_id: str,
        ip: Optional[str] = None,
    ) -> Dict[str, Any]:
        from app.core.openai_keys import has_openai_key
        if not has_openai_key():
            raise AppError(
                code="image.openai_key_missing",
                message=t("image.openai_key_missing"),
                status_code=400,
            )

        img = self.get_image(image_id)
        fp = Path(img.file_path)
        if not fp.exists():
            raise AppError(
                code="image.not_found",
                message=t("image.not_found"),
                status_code=404,
            )

        image_bytes = fp.read_bytes()
        validator = ImageValidator()

        if img.asset_type == "reference":
            entity_info: Dict[str, Any] = {}
            if img.entity_id:
                entity = (
                    self._db.query(EntityCanon)
                    .filter(EntityCanon.id == img.entity_id)
                    .first()
                )
                if entity:
                    entity_info = {
                        "name": entity.name,
                        "entity_type": entity.entity_type,
                        "description": entity.description or "",
                        "stable_traits": entity.stable_traits or "{}",
                    }
            result = validator.validate_reference_image(image_bytes, entity_info)
        else:
            scene_info: Dict[str, Any] = {}
            entity_names: List[str] = []
            if img.still_id:
                still = (
                    self._db.query(SceneStill)
                    .filter(SceneStill.id == img.still_id)
                    .first()
                )
                if still:
                    scene_info = {
                        "scene_heading": still.screenplay_scene_heading or "",
                        "beat_title": still.beat_title or "",
                        "still_frame_prompt": still.still_frame_prompt or "",
                    }
                    try:
                        visible_ids = json.loads(still.visible_entities_json or "[]")
                        for v in visible_ids:
                            eid = (v.get("id") or v.get("entity_id", "")) if isinstance(v, dict) else v
                            entity = (
                                self._db.query(EntityCanon)
                                .filter(EntityCanon.id == eid)
                                .first()
                            )
                            if entity:
                                entity_names.append(entity.name)
                    except (json.JSONDecodeError, TypeError) as exc:
                        logger.warning(
                            "visible_entities_json parse failed for still %s: %s — entity_names 없이 검증 계속",
                            img.still_id, exc,
                        )
            result = validator.validate_scene_image(image_bytes, scene_info, entity_names)

        score = result.get("score", 0)
        img.validation_score = score
        img.validation_result = json.dumps(result, ensure_ascii=False)
        if score < 60:
            issues = result.get("issues", [])
            issue_text = "; ".join(issues) if issues else "Low validation score"
            img.status = "needs_fix"
            img.review_notes = f"[auto-validation] score={score}: {issue_text}"
        self._db.commit()

        self._logger.log(
            actor_id=self._actor_id,
            action="image.validate",
            resource_type="image",
            resource_id=image_id,
            project_id=self._project_id,
            detail={"score": score, "passed": result.get("passed", False)},
            ip_address=ip,
        )

        return result

    def get_validation(self, image_id: str) -> Dict[str, Any]:
        img = self.get_image(image_id)
        result: Dict[str, Any] = {
            "image_id": image_id,
            "score": img.validation_score,
            "passed": None,
            "issues": [],
            "description": "",
        }
        if img.validation_result:
            try:
                parsed = json.loads(img.validation_result)
                result["passed"] = parsed.get("passed")
                result["issues"] = parsed.get("issues", [])
                result["description"] = parsed.get("description", "")
            except json.JSONDecodeError as exc:
                logger.warning("validation_result parse failed for image %s: %s", img.id, exc)
        return result
