File size: 4,570 Bytes
aad7814
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
"""Embedding abstraction for the v2 backend.

Default backend is a free, local sentence-transformers model (MiniLM). Setting
``embedding_provider='openai'`` switches to OpenAI embeddings. The active
embedder exposes ``embed_dim`` so the FAISS store can size its index, and a
deterministic fake backend is used when neither dependency is available (tests).
"""

from __future__ import annotations

import hashlib
import logging
import struct

from backend.config import settings

logger = logging.getLogger(__name__)


class Embedder:
    """Minimal embedding interface used by the RAG store."""

    embed_dim: int

    def embed_documents(self, texts: list[str]) -> list[list[float]]:  # pragma: no cover - interface
        raise NotImplementedError

    def embed_query(self, text: str) -> list[float]:  # pragma: no cover - interface
        raise NotImplementedError


class LocalEmbedder(Embedder):
    """sentence-transformers backend (default)."""

    def __init__(self, model_name: str) -> None:
        from sentence_transformers import SentenceTransformer

        self._model = SentenceTransformer(model_name)
        self.embed_dim = int(self._model.get_sentence_embedding_dimension())
        logger.info("LocalEmbedder loaded %s (dim=%d)", model_name, self.embed_dim)

    def embed_documents(self, texts: list[str]) -> list[list[float]]:
        vecs = self._model.encode(
            list(texts), normalize_embeddings=True, convert_to_numpy=True
        )
        return [v.tolist() for v in vecs]

    def embed_query(self, text: str) -> list[float]:
        return self.embed_documents([text])[0]


class OpenAIEmbedder(Embedder):
    """OpenAI embeddings backend."""

    _DIMS = {
        "text-embedding-3-small": 1536,
        "text-embedding-3-large": 3072,
        "text-embedding-ada-002": 1536,
    }

    def __init__(self, model_name: str, api_key: str) -> None:
        from openai import OpenAI

        self._client = OpenAI(
            api_key=api_key,
            timeout=float(settings.openai_request_timeout_seconds),
        )
        self._model = model_name
        self.embed_dim = self._DIMS.get(model_name, 1536)
        logger.info("OpenAIEmbedder using %s (dim=%d)", model_name, self.embed_dim)

    def embed_documents(self, texts: list[str]) -> list[list[float]]:
        from backend.llm.openai_client import pipeline_timeout

        resp = self._client.embeddings.create(
            model=self._model,
            input=list(texts),
            timeout=pipeline_timeout(),
        )
        return [d.embedding for d in resp.data]

    def embed_query(self, text: str) -> list[float]:
        return self.embed_documents([text])[0]


class FakeEmbedder(Embedder):
    """Deterministic hash-based embedder for offline tests."""

    def __init__(self, dim: int = 384) -> None:
        self.embed_dim = dim

    def _vec(self, text: str) -> list[float]:
        out: list[float] = []
        seed = (text or "").encode("utf-8")
        counter = 0
        while len(out) < self.embed_dim:
            h = hashlib.sha256(seed + struct.pack(">I", counter)).digest()
            for i in range(0, len(h), 4):
                if len(out) >= self.embed_dim:
                    break
                (val,) = struct.unpack(">I", h[i : i + 4])
                out.append((val / 0xFFFFFFFF) * 2.0 - 1.0)
            counter += 1
        norm = sum(x * x for x in out) ** 0.5 or 1.0
        return [x / norm for x in out]

    def embed_documents(self, texts: list[str]) -> list[list[float]]:
        return [self._vec(t) for t in texts]

    def embed_query(self, text: str) -> list[float]:
        return self._vec(text)


_instance: Embedder | None = None


def get_embedder() -> Embedder:
    """Return the cached embedder singleton selected by configuration."""
    global _instance
    if _instance is not None:
        return _instance

    provider = (settings.embedding_provider or "local").lower()
    try:
        if provider == "openai" and settings.openai_api_key:
            _instance = OpenAIEmbedder(settings.openai_embedding_model, settings.openai_api_key)
        else:
            _instance = LocalEmbedder(settings.local_embedding_model)
    except Exception as exc:  # noqa: BLE001 - dependency/model load failures → fake
        logger.warning("Embedder init failed (%s); using FakeEmbedder (TEST ONLY).", exc)
        _instance = FakeEmbedder()
    return _instance


def reset_embedder() -> None:
    """Reset the cached embedder (tests / config reloads)."""
    global _instance
    _instance = None