"""Embedding services for DSCons knowledge retrieval."""

from __future__ import annotations

import hashlib
import os
from collections.abc import Sequence
from pathlib import Path

try:
    from sentence_transformers import SentenceTransformer
except ImportError:  # pragma: no cover
    SentenceTransformer = None  # type: ignore[assignment]

from app.core.settings import Settings


class EmbeddingService:
    """PhoBERT-oriented embedding wrapper with safe fallback behavior."""

    def __init__(self, settings: Settings) -> None:
        """Initialize the embedding backend if available."""
        self._settings = settings
        self._model_name = settings.embedding_model_name or "vinai/phobert-base-v2"
        self._fallback_dimension = 768
        self._model = None
        self._cache_dir = Path(settings.embedding_cache_dir)
        self._cache_dir.mkdir(parents=True, exist_ok=True)
        os.environ.setdefault("HF_HOME", str(self._cache_dir))
        os.environ.setdefault("TRANSFORMERS_CACHE", str(self._cache_dir))
        os.environ.setdefault("SENTENCE_TRANSFORMERS_HOME", str(self._cache_dir))
        self._disable_model_load = os.getenv(
            "DSCONS_DISABLE_EMBEDDING_MODEL_LOAD", "0"
        ).lower() in {
            "1",
            "true",
            "yes",
            "on",
        }

        if (
            SentenceTransformer is not None
            and not self._disable_model_load
            and settings.embedding_eager_load
        ):
            self._load_model()

    def _load_model(self) -> None:
        """Load the embedding model once and reuse it across requests."""
        if self._model is not None or SentenceTransformer is None:
            return

        try:
            # Direct SentenceTransformer attempt
            self._model = SentenceTransformer(
                self._model_name,
                cache_folder=str(self._cache_dir),
                local_files_only=self._settings.embedding_local_files_only,
            )
        except Exception:
            # Fallback for complex multi-module models like BAAI/bge-m3
            try:
                from sentence_transformers.models import Normalize, Pooling, Transformer

                word_model = Transformer(
                    self._model_name, cache_dir=str(self._cache_dir)
                )
                dim = (
                    getattr(word_model, "get_word_embedding_dimension", None)
                    or word_model.get_embedding_dimension
                )
                pooling_model = Pooling(dim())
                self._model = SentenceTransformer(
                    modules=[word_model, pooling_model, Normalize()]
                )
            except Exception:
                self._model = None

        if self._model is not None and hasattr(self._model, "max_seq_length"):
            self._model.max_seq_length = min(self._model.max_seq_length or 1024, 8192)

    @property
    def model_name(self) -> str:
        """Return the configured embedding model name."""
        return self._model_name

    @property
    def vector_size(self) -> int:
        """Return the active embedding dimensionality."""
        if self._model is None and not self._disable_model_load:
            self._load_model()

        if self._model is not None:
            dimension = self._model.get_sentence_embedding_dimension()
            if isinstance(dimension, int) and dimension > 0:
                return dimension
        return self._fallback_dimension

    def embed_texts(self, texts: Sequence[str]) -> list[list[float]]:
        """Embed a batch of texts."""
        normalized_texts = [
            text if isinstance(text, str) else str(text) for text in texts
        ]
        if not normalized_texts:
            return []

        if self._model is None and not self._disable_model_load:
            self._load_model()

        if self._model is not None:
            vectors = self._model.encode(
                normalized_texts,
                convert_to_numpy=False,
                normalize_embeddings=True,
            )
            return [list(map(float, vector)) for vector in vectors]

        return [self._fallback_embed(text) for text in normalized_texts]

    def embed_query(self, text: str) -> list[float]:
        """Embed a single search query."""
        vectors = self.embed_texts([text])
        return vectors[0] if vectors else [0.0] * self.vector_size

    def _fallback_embed(self, text: str) -> list[float]:
        """Generate a deterministic fallback embedding from text content."""
        vector = [0.0] * self._fallback_dimension
        encoded = text.encode("utf-8")
        if not encoded:
            return vector

        digest_stream = b""
        counter = 0
        required_bytes = self._fallback_dimension * 4

        while len(digest_stream) < required_bytes:
            digest_stream += hashlib.sha256(
                encoded + counter.to_bytes(4, "little")
            ).digest()
            counter += 1

        for index in range(self._fallback_dimension):
            start = index * 4
            chunk = digest_stream[start : start + 4]
            value = int.from_bytes(chunk, "little", signed=False)
            vector[index] = (value / 4294967295.0) * 2.0 - 1.0

        norm = sum(component * component for component in vector) ** 0.5
        if norm == 0:
            return vector
        return [component / norm for component in vector]
