"""T2I 생성 결과 추적 모듈 — Opik 스타일 생성 로그."""

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

from sqlalchemy.orm import Session as OrmSession

from app.models.project import GenerationTrace

logger = logging.getLogger(__name__)


class GenerationTracker:
    """T2I 생성 결과를 추적하는 Opik 스타일 모듈."""

    def __init__(self, db_session: OrmSession, project_id: str = "") -> None:
        self._db = db_session
        self._project_id = project_id

    def record(
        self,
        *,
        still_id: Optional[str] = None,
        entity_id: Optional[str] = None,
        image_asset_id: Optional[str] = None,
        attempt_number: int,
        prompt_used: str,
        prompt_version: str,
        model_name: Optional[str] = None,
        status: str,
        block_reason: Optional[str] = None,
        block_categories: Optional[List[str]] = None,
        response_time_ms: Optional[int] = None,
        sanitizer_feedback: Optional[str] = None,
    ) -> GenerationTrace:
        """Record a T2I generation attempt."""
        trace = GenerationTrace(
            id=str(uuid.uuid4()),
            project_id=self._project_id,
            image_asset_id=image_asset_id,
            still_id=still_id,
            entity_id=entity_id,
            attempt_number=attempt_number,
            prompt_used=prompt_used,
            prompt_version=prompt_version,
            model_name=model_name,
            status=status,
            block_reason=block_reason,
            block_categories=json.dumps(block_categories or [], ensure_ascii=False),
            response_time_ms=response_time_ms,
            sanitizer_feedback=sanitizer_feedback,
            created_at=datetime.now(timezone.utc).isoformat(),
        )
        self._db.add(trace)
        self._db.flush()

        logger.info(
            "Generation trace recorded: status=%s, attempt=%d, version=%s",
            status,
            attempt_number,
            prompt_version,
        )

        return trace

    def get_traces(
        self,
        still_id: Optional[str] = None,
        entity_id: Optional[str] = None,
        image_asset_id: Optional[str] = None,
    ) -> List[Dict[str, Any]]:
        """Get all traces for a still, entity, or image asset."""
        query = self._db.query(GenerationTrace)
        if still_id:
            query = query.filter(GenerationTrace.still_id == still_id)
        if entity_id:
            query = query.filter(GenerationTrace.entity_id == entity_id)
        if image_asset_id:
            query = query.filter(GenerationTrace.image_asset_id == image_asset_id)
        query = query.order_by(GenerationTrace.created_at)
        traces = query.all()
        return [self._trace_to_dict(t) for t in traces]

    def get_rejection_stats(self) -> Dict[str, Any]:
        """Get aggregated rejection statistics."""
        query = self._db.query(GenerationTrace)
        if self._project_id:
            query = query.filter(GenerationTrace.project_id == self._project_id)
        all_traces = query.all()

        total = len(all_traces)
        success = sum(1 for t in all_traces if t.status == "success")
        blocked = sum(1 for t in all_traces if t.status == "moderation_blocked")
        error = sum(1 for t in all_traces if t.status in ("error", "timeout"))

        block_reasons: Dict[str, int] = {}
        for t in all_traces:
            if t.status == "moderation_blocked" and t.block_reason:
                block_reasons[t.block_reason] = block_reasons.get(t.block_reason, 0) + 1

        return {
            "total": total,
            "success": success,
            "blocked": blocked,
            "error": error,
            "block_reasons": block_reasons,
        }

    def get_prompt_evolution(self, still_id: str) -> List[Dict[str, Any]]:
        """Get the prompt evolution for a still (original -> sanitized -> ...)."""
        traces = (
            self._db.query(GenerationTrace)
            .filter(GenerationTrace.still_id == still_id)
            .order_by(GenerationTrace.attempt_number)
            .all()
        )
        return [self._trace_to_dict(t) for t in traces]

    @staticmethod
    def _trace_to_dict(trace: GenerationTrace) -> Dict[str, Any]:
        return {
            "id": trace.id,
            "image_asset_id": trace.image_asset_id,
            "still_id": trace.still_id,
            "entity_id": trace.entity_id,
            "attempt_number": trace.attempt_number,
            "prompt_used": trace.prompt_used,
            "prompt_version": trace.prompt_version,
            "model_name": trace.model_name,
            "status": trace.status,
            "block_reason": trace.block_reason,
            "block_categories": trace.block_categories,
            "response_time_ms": trace.response_time_ms,
            "sanitizer_feedback": trace.sanitizer_feedback,
            "created_at": trace.created_at,
        }
