Files
ai-rpg/app/core/embeddings.py
2026-06-20 19:13:05 +03:00

154 lines
4.4 KiB
Python

"""Embedders for RAG.
Two implementations:
- `HashEmbedder`: offline, deterministic bag-of-words + hash projection. Used for dev/test.
- `OpenAIEmbedder`: calls an OpenAI-compatible embeddings API at runtime.
The active embedder is chosen via `settings.embeddings.provider`.
"""
from __future__ import annotations
import hashlib
import math
import re
from collections import Counter
from typing import Protocol, runtime_checkable
import httpx
from app.core.logging import get_logger
_logger = get_logger(__name__)
_WORD_RE = re.compile(r"\w+", re.UNICODE)
def _tokenize(text: str) -> list[str]:
return [w.lower() for w in _WORD_RE.findall(text)]
@runtime_checkable
class Embedder(Protocol):
async def embed(self, texts: list[str]) -> list[list[float]]: ...
@property
def dimension(self) -> int: ...
class HashEmbedder:
"""Offline bag-of-words embedder with hash projection.
Not semantically meaningful, but deterministic and fast — sufficient for
integration tests and local dev. Cosine similarity is non-zero only when
texts share tokens.
"""
def __init__(self, dimension: int = 256):
if dimension <= 0:
raise ValueError("dimension must be positive")
self._dim = dimension
@property
def dimension(self) -> int:
return self._dim
async def embed(self, texts: list[str]) -> list[list[float]]:
out: list[list[float]] = []
for text in texts:
out.append(self._hash_project(text))
return out
def _hash_project(self, text: str) -> list[float]:
vec = [0.0] * self._dim
tokens = _tokenize(text)
if not tokens:
return vec
counts = Counter(tokens)
for token, count in counts.items():
h = hashlib.md5(token.encode("utf-8")).digest()
# Use first 4 bytes for index, next 4 bytes for sign
idx = int.from_bytes(h[:4], "little") % self._dim
sign = 1.0 if (h[4] & 1) == 0 else -1.0
vec[idx] += sign * math.sqrt(count)
# L2 normalize
norm = math.sqrt(sum(v * v for v in vec))
if norm > 0:
vec = [v / norm for v in vec]
return vec
class OpenAIEmbedder:
"""OpenAI-compatible embeddings API client."""
def __init__(
self,
api_url: str,
api_key: str,
model: str,
dimension: int,
timeout: float = 30.0,
batch_size: int = 32,
):
self._api_url = api_url.rstrip("/")
self._api_key = api_key
self._model = model
self._dim = dimension
self._timeout = timeout
self._batch_size = batch_size
@property
def dimension(self) -> int:
return self._dim
async def embed(self, texts: list[str]) -> list[list[float]]:
if not texts:
return []
out: list[list[float]] = []
async with httpx.AsyncClient(timeout=self._timeout) as client:
for i in range(0, len(texts), self._batch_size):
batch = texts[i : i + self._batch_size]
resp = await client.post(
f"{self._api_url}/embeddings",
headers={"Authorization": f"Bearer {self._api_key}"},
json={"model": self._model, "input": batch},
)
resp.raise_for_status()
data = resp.json()
# Sort by index to preserve order
sorted_data = sorted(data["data"], key=lambda x: x["index"])
out.extend(d["embedding"] for d in sorted_data)
return out
async def probe_dimension(self, sample_text: str = "hello world") -> int:
"""Make a single embedding call and return the dimension of the result.
Useful for the "auto-probe dimension" admin button.
"""
result = await self.embed([sample_text])
if not result:
raise RuntimeError("Empty embeddings response")
return len(result[0])
def build_hash_embedder(dimension: int) -> HashEmbedder:
return HashEmbedder(dimension=dimension)
def build_openai_embedder(
api_url: str,
api_key: str,
model: str,
dimension: int,
timeout: float = 30.0,
batch_size: int = 32,
) -> OpenAIEmbedder:
return OpenAIEmbedder(
api_url=api_url,
api_key=api_key,
model=model,
dimension=dimension,
timeout=timeout,
batch_size=batch_size,
)