"""파이프라인 작업의 프로비저닝을 기록하는 헬퍼 — Opik 호환."""

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

from sqlalchemy.orm import Session as OrmSession

from app.core.version_registry import MODULE_VERSIONS, get_module_info
from app.models.project import OperationLog

logger = logging.getLogger(__name__)

_MAX_SUMMARY_LEN = 2000


def _truncate_json(data: dict, max_len: int = _MAX_SUMMARY_LEN) -> str:
    """Serialize dict to JSON, truncating if necessary."""
    raw = json.dumps(data, ensure_ascii=False, default=str)
    if len(raw) > max_len:
        return raw[: max_len - 3] + "..."
    return raw


class OperationContext:
    """Context manager for recording operation provenance."""

    def __init__(
        self,
        recorder: "ProvenanceRecorder",
        operation_type: str,
        module_name: str,
        project_id: str = "",
        episode_id: Optional[str] = None,
        prompt_name: Optional[str] = None,
        prompt_version: Optional[str] = None,
        prompt_hash: Optional[str] = None,
    ) -> None:
        self._recorder = recorder
        self._op_id = str(uuid.uuid4())
        self._operation_type = operation_type
        self._module_name = module_name
        self._module_version = MODULE_VERSIONS.get(module_name, "0.0.0")
        self._project_id = project_id
        self._episode_id = episode_id
        self._prompt_name = prompt_name
        self._prompt_version = prompt_version
        self._prompt_hash = prompt_hash
        self._input_summary: Optional[str] = None
        self._output_summary: Optional[str] = None
        self._token_usage: str = "{}"
        self._metadata: str = "{}"
        self._start_time = time.time()
        self._status: Optional[str] = None
        self._error_message: Optional[str] = None

    @property
    def operation_id(self) -> str:
        return self._op_id

    def set_input(self, summary: dict) -> None:
        """Record input summary (truncated to max length)."""
        self._input_summary = _truncate_json(summary)

    def set_output(self, summary: dict) -> None:
        """Record output summary (truncated to max length)."""
        self._output_summary = _truncate_json(summary)

    def set_token_usage(
        self,
        input_tokens: int = 0,
        output_tokens: int = 0,
        cost_usd: float = 0.0,
    ) -> None:
        """Record token usage information."""
        self._token_usage = json.dumps(
            {
                "input_tokens": input_tokens,
                "output_tokens": output_tokens,
                "cost_usd": cost_usd,
            }
        )

    def set_metadata(self, metadata: dict) -> None:
        """Record extra metadata."""
        self._metadata = _truncate_json(metadata)

    def finish(self, status: str = "success") -> None:
        """Mark operation as finished with given status."""
        self._status = status
        self._save()

    def fail(self, error_message: str) -> None:
        """Mark operation as failed."""
        self._status = "error"
        self._error_message = str(error_message)[:2000]
        self._save()

    def _save(self) -> None:
        """Persist the operation log entry."""
        elapsed_ms = int((time.time() - self._start_time) * 1000)
        now = datetime.now(timezone.utc).isoformat()

        log_entry = OperationLog(
            id=self._op_id,
            project_id=self._project_id,
            operation_type=self._operation_type,
            episode_id=self._episode_id,
            module_name=self._module_name,
            module_version=self._module_version,
            prompt_name=self._prompt_name,
            prompt_version=self._prompt_version,
            prompt_hash=self._prompt_hash,
            input_summary=self._input_summary,
            output_summary=self._output_summary,
            status=self._status or "error",
            error_message=self._error_message,
            duration_ms=elapsed_ms,
            token_usage=self._token_usage,
            metadata_json=self._metadata,
            created_at=now,
        )
        self._recorder._db.add(log_entry)
        self._recorder._db.flush()

        logger.info(
            "Operation logged: type=%s, module=%s, status=%s, duration=%dms",
            self._operation_type,
            self._module_name,
            self._status,
            elapsed_ms,
        )

    def __enter__(self) -> "OperationContext":
        return self

    def __exit__(self, exc_type, exc_val, exc_tb) -> bool:
        if self._status is not None:
            # Already saved (finish/fail called explicitly)
            return False
        if exc_type is not None:
            self.fail(str(exc_val))
            return False
        self.finish()
        return False


class ProvenanceRecorder:
    """파이프라인 작업의 프로비저닝을 기록하는 헬퍼."""

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

    def start_operation(
        self,
        operation_type: str,
        module_name: str,
        episode_id: Optional[str] = None,
        prompt_path: Optional[str] = None,
    ) -> OperationContext:
        """작업 시작. context manager로 사용.

        Args:
            operation_type: e.g., entity_extraction, scene_still_extraction
            module_name: key from MODULE_VERSIONS
            episode_id: optional episode ID
            prompt_path: optional path to a prompt file for hashing

        Returns:
            OperationContext context manager.
        """
        prompt_name = None
        prompt_version = None
        prompt_hash = None

        if prompt_path:
            prompt_hash = self.compute_prompt_hash(prompt_path)
            # Extract prompt_name and prompt_version from path
            # e.g., ".../entity_extraction/v5/chunk_system.md" -> "entity_extraction/v5"
            p = Path(prompt_path)
            if p.exists():
                parts = p.parts
                # Find the version directory (v1, v2, etc.)
                for i, part in enumerate(parts):
                    if part.startswith("v") and part[1:].isdigit():
                        prompt_version = part
                        if i > 0:
                            prompt_name = f"{parts[i-1]}/{part}"
                        break

        # Fallback: get prompt dependency from module info
        if not prompt_name:
            info = get_module_info(module_name)
            if info.get("prompt_dependency"):
                prompt_name = info["prompt_dependency"]
                # Extract version from prompt_name like "entity_extraction/v5"
                if "/" in prompt_name:
                    prompt_version = prompt_name.split("/")[-1]

        return OperationContext(
            recorder=self,
            operation_type=operation_type,
            module_name=module_name,
            project_id=self._project_id,
            episode_id=episode_id,
            prompt_name=prompt_name,
            prompt_version=prompt_version,
            prompt_hash=prompt_hash,
        )

    @staticmethod
    def compute_prompt_hash(prompt_path: str) -> Optional[str]:
        """프롬프트 파일의 SHA256 해시 계산."""
        p = Path(prompt_path)
        if not p.exists():
            return None
        content = p.read_text(encoding="utf-8")
        return hashlib.sha256(content.encode("utf-8")).hexdigest()

    def get_operations(
        self,
        operation_type: Optional[str] = None,
        episode_id: Optional[str] = None,
        offset: int = 0,
        limit: int = 50,
    ) -> list:
        """Get operation logs with optional filters."""
        query = self._db.query(OperationLog)
        if self._project_id:
            query = query.filter(OperationLog.project_id == self._project_id)
        if operation_type:
            query = query.filter(OperationLog.operation_type == operation_type)
        if episode_id:
            query = query.filter(OperationLog.episode_id == episode_id)
        query = query.order_by(OperationLog.created_at.desc())
        return query.offset(offset).limit(limit).all()

    def get_operation(self, op_id: str) -> Optional[OperationLog]:
        """Get a single operation by ID."""
        return (
            self._db.query(OperationLog)
            .filter(OperationLog.id == op_id)
            .first()
        )

    def get_summary(self) -> list:
        """Get aggregated stats per operation type."""
        query = self._db.query(OperationLog)
        if self._project_id:
            query = query.filter(OperationLog.project_id == self._project_id)
        all_ops = query.all()
        stats: Dict[str, Dict[str, Any]] = {}
        for op in all_ops:
            key = op.operation_type
            if key not in stats:
                stats[key] = {
                    "operation_type": key,
                    "total": 0,
                    "success": 0,
                    "error": 0,
                    "total_duration_ms": 0,
                }
            stats[key]["total"] += 1
            if op.status == "success":
                stats[key]["success"] += 1
            elif op.status == "error":
                stats[key]["error"] += 1
            if op.duration_ms:
                stats[key]["total_duration_ms"] += op.duration_ms

        result = []
        for s in stats.values():
            avg = s["total_duration_ms"] / s["total"] if s["total"] else 0.0
            result.append({
                "operation_type": s["operation_type"],
                "total": s["total"],
                "success": s["success"],
                "error": s["error"],
                "avg_duration_ms": round(avg, 1),
            })
        return result
