"""Embedding backends — pluggable text → vector encoders.

Ships a deterministic ``HashEmbedder`` so semantic similarity tests work
offline. Real backends (OpenAI, Cohere, local sentence-transformers) plug
in with the same protocol.
"""

from __future__ import annotations

import hashlib
import math
from collections.abc import Sequence
from typing import Protocol


class EmbeddingError(RuntimeError):
    """Raised on registration errors or embedding failures."""


class Embedder(Protocol):
    name: str
    dimensions: int

    def embed(self, text: str) -> list[float]: ...

    def embed_many(self, texts: Sequence[str]) -> list[list[float]]: ...


class HashEmbedder:
    """Deterministic embedder based on token hashes.

    Vector length defaults to 32 — small enough for fast tests, large
    enough to keep collisions rare. Cosine similarity between two strings
    sharing many tokens trends toward 1.0.
    """

    name = "hash"

    def __init__(self, dimensions: int = 32) -> None:
        if dimensions < 4:
            raise EmbeddingError("dimensions must be >= 4")
        self.dimensions = dimensions

    def _vector_for_token(self, token: str) -> list[float]:
        digest = hashlib.blake2b(token.encode("utf-8"), digest_size=self.dimensions).digest()
        # Map each byte to a centred float in [-1, 1].
        return [(b - 128) / 128.0 for b in digest]

    def embed(self, text: str) -> list[float]:
        tokens = [t for t in text.lower().split() if t]
        if not tokens:
            return [0.0] * self.dimensions
        acc = [0.0] * self.dimensions
        for token in tokens:
            vec = self._vector_for_token(token)
            for i in range(self.dimensions):
                acc[i] += vec[i]
        norm = math.sqrt(sum(x * x for x in acc)) or 1.0
        return [x / norm for x in acc]

    def embed_many(self, texts: Sequence[str]) -> list[list[float]]:
        return [self.embed(t) for t in texts]


def cosine_similarity(a: Sequence[float], b: Sequence[float]) -> float:
    if len(a) != len(b):
        raise EmbeddingError("vectors must have equal length")
    dot = sum(x * y for x, y in zip(a, b, strict=True))
    na = math.sqrt(sum(x * x for x in a)) or 1.0
    nb = math.sqrt(sum(x * x for x in b)) or 1.0
    return dot / (na * nb)


_EMBEDDERS: dict[str, Embedder] = {}


def register_embedder(embedder: Embedder) -> Embedder:
    name = getattr(embedder, "name", "")
    if not name:
        raise EmbeddingError("embedder must expose a non-empty .name")
    _EMBEDDERS[name] = embedder
    return embedder


def get_embedder(name: str) -> Embedder:
    try:
        return _EMBEDDERS[name]
    except KeyError as exc:
        raise EmbeddingError(f"unknown embedder {name!r}") from exc


def list_embedders() -> list[str]:
    return sorted(_EMBEDDERS)


def reset_for_tests() -> None:
    _EMBEDDERS.clear()
    register_default_embeddings()


def register_default_embeddings() -> None:
    if "hash" not in _EMBEDDERS:
        register_embedder(HashEmbedder())
