"""Embedding generation service for table metadata."""

from enum import Enum
from typing import Any, List

import numpy as np
from loguru import logger

from quber.settings import get_settings


class EmbeddingProvider(Enum):
    """Available embedding providers."""

    LOCAL = "local"
    OPENAI = "openai"


class EmbeddingService:
    """Generate embeddings for table metadata using local or OpenAI models."""

    def __init__(
        self,
        provider: str | EmbeddingProvider = "local",
        device: str = "cuda",
        model_name: str | None = None,
    ):
        """
        Initialize the embedding service.

        Args:
            provider: Embedding provider ('local' or 'openai')
            device: Device for local model ('cuda' or 'cpu')
            model_name: Optional custom model name (defaults per provider)
        """
        if isinstance(provider, str):
            provider = EmbeddingProvider(provider)

        self.provider = provider
        self.device = device
        self.model: Any = None
        self.client: Any = None
        self.model_name: str = ""
        self.embedding_dim: int = 0

        if provider == EmbeddingProvider.LOCAL:
            self.init_local_model(model_name or "BAAI/bge-large-en-v1.5", device)
        elif provider == EmbeddingProvider.OPENAI:
            self.init_openai_client(model_name or "text-embedding-3-small")
        else:
            raise ValueError(f"Unsupported provider: {provider}")

        logger.info(f"Initialized {provider.value} embedding service")

    def init_local_model(self, model_name: str, device: str):
        """Initialize local sentence-transformers model."""
        try:
            from sentence_transformers import SentenceTransformer
        except ImportError as err:
            raise ImportError(
                "sentence-transformers not installed. Install with: pip install sentence-transformers"
            ) from err

        logger.info(f"Loading local model: {model_name} on {device}")
        self.model = SentenceTransformer(model_name)

        # Move to specified device (CUDA if available)
        if device == "cuda":
            import torch

            if torch.cuda.is_available():
                self.model.to(device)
                logger.success(f"Model loaded on {device} (GPU: {torch.cuda.get_device_name(0)})")
            else:
                logger.warning("CUDA not available, using CPU")
                self.model.to("cpu")
                self.device = "cpu"
        else:
            self.model.to("cpu")

        self.model_name = model_name
        self.embedding_dim = self.model.get_sentence_embedding_dimension()
        logger.info(f"Model embedding dimension: {self.embedding_dim}")

    def init_openai_client(self, model_name: str):
        """Initialize OpenAI client."""
        try:
            from openai import OpenAI
        except ImportError as err:
            raise ImportError("openai not installed. Install with: pip install openai") from err

        api_key = get_settings().llm.openai_api_key
        if not api_key:
            raise ValueError("OPENAI_API_KEY environment variable not set")

        self.client = OpenAI(api_key=api_key)
        self.model_name = model_name
        # OpenAI text-embedding-3-small is 1536 dims
        self.embedding_dim = 1536 if "3-small" in model_name else 3072
        logger.info(f"Using OpenAI model: {model_name} ({self.embedding_dim} dims)")

    def embed(
        self, texts: List[str], batch_size: int = 32, show_progress: bool = False
    ) -> "np.ndarray[Any, Any]":
        """
        Generate embeddings for a list of texts.

        Args:
            texts: List of text strings to embed
            batch_size: Batch size for processing (local model only)
            show_progress: Show progress bar (local model only)

        Returns:
            NumPy array of shape (len(texts), embedding_dim)
        """
        if not texts:
            return np.array([])

        if self.provider == EmbeddingProvider.LOCAL:
            return self.embed_local(texts, batch_size, show_progress)
        else:
            return self.embed_openai(texts)

    def embed_local(self, texts: List[str], batch_size: int, show_progress: bool) -> "np.ndarray[Any, Any]":
        """Generate embeddings using local model."""
        logger.debug(f"Generating {len(texts)} embeddings (batch_size={batch_size})")
        embeddings = self.model.encode(
            texts, batch_size=batch_size, show_progress_bar=show_progress, convert_to_numpy=True
        )
        return embeddings

    def embed_openai(self, texts: List[str]) -> "np.ndarray[Any, Any]":
        """Generate embeddings using OpenAI API."""
        logger.debug(f"Generating {len(texts)} embeddings via OpenAI")

        # OpenAI has a limit of 2048 texts per request
        # For simplicity, batch in chunks of 100
        all_embeddings = []
        chunk_size = 100

        for i in range(0, len(texts), chunk_size):
            chunk = texts[i : i + chunk_size]
            response = self.client.embeddings.create(input=chunk, model=self.model_name)
            embeddings = [item.embedding for item in response.data]
            all_embeddings.extend(embeddings)

        return np.array(all_embeddings)

    def embed_table_metadata(self, title: str, description: str) -> "np.ndarray[Any, Any]":
        """
        Generate embedding for table metadata.

        Combines title and description into a single text for embedding.

        Args:
            title: Table title (llm_title)
            description: Table description (llm_description)

        Returns:
            Embedding vector as NumPy array
        """
        # Combine title and description
        # Title is more important, so it comes first
        text = f"{title}\n\n{description}"

        embeddings = self.embed([text])
        return embeddings[0]


def get_embedding_service(provider: str | None = None, device: str | None = None) -> EmbeddingService:
    """
    Factory function to create an embedding service based on environment config.

    Args:
        provider: Override provider (defaults to EMBEDDING_PROVIDER env var or 'local')
        device: Override device (defaults to EMBEDDING_DEVICE env var or 'cuda')

    Returns:
        Configured EmbeddingService instance
    """
    embedding_settings = get_settings().embeddings
    if provider is None:
        provider = embedding_settings.provider

    if device is None:
        device = embedding_settings.device

    return EmbeddingService(provider=provider, device=device)
