79 lines
3.0 KiB
Python
79 lines
3.0 KiB
Python
import math
|
|
|
|
import requests
|
|
|
|
from tht.config import EmbeddingsConfig
|
|
|
|
|
|
class EmbeddingsError(Exception):
|
|
pass
|
|
|
|
|
|
class OllamaInternalEmbeddings:
|
|
"""Client embeddings for the installation-owned internal Ollama endpoint."""
|
|
|
|
def __init__(self, cfg: EmbeddingsConfig, *, session: requests.Session | None = None):
|
|
self.cfg = cfg
|
|
self._session = session or requests.Session()
|
|
self.base_url = self.cfg.base_url.rstrip("/")
|
|
self.model = self.cfg.model
|
|
self.dim = self.cfg.dim
|
|
self.timeout = (self.cfg.connect_timeout, self.cfg.timeout)
|
|
self.batch_size = self.cfg.batch_size
|
|
|
|
def _post(self, texts: list[str]) -> list[list[float]]:
|
|
try:
|
|
response = self._session.post(
|
|
f"{self.base_url}/api/embed",
|
|
json={"model": self.model, "input": texts},
|
|
timeout=self.timeout,
|
|
)
|
|
response.raise_for_status()
|
|
except requests.RequestException as exc:
|
|
raise EmbeddingsError(
|
|
f"internal Ollama embeddings request failed for model {self.model}"
|
|
) from exc
|
|
try:
|
|
payload = response.json()
|
|
except ValueError as exc:
|
|
raise EmbeddingsError("internal Ollama returned an invalid JSON response") from exc
|
|
if not isinstance(payload, dict):
|
|
raise EmbeddingsError("internal Ollama returned a non-object response payload")
|
|
embeddings = payload.get("embeddings")
|
|
if not isinstance(embeddings, list):
|
|
raise EmbeddingsError("internal Ollama response is missing embeddings")
|
|
if len(embeddings) != len(texts):
|
|
raise EmbeddingsError(
|
|
f"unexpected embedding count: {len(embeddings)} != {len(texts)}"
|
|
)
|
|
validated: list[list[float]] = []
|
|
for vector in embeddings:
|
|
if not isinstance(vector, list):
|
|
raise EmbeddingsError("internal Ollama returned a non-vector embedding")
|
|
if len(vector) != self.dim:
|
|
raise EmbeddingsError(
|
|
f"unexpected embedding dimension: {len(vector)} != {self.dim}"
|
|
)
|
|
cleaned: list[float] = []
|
|
for value in vector:
|
|
if not isinstance(value, (int, float)) or not math.isfinite(value):
|
|
raise EmbeddingsError("internal Ollama returned a non-finite embedding value")
|
|
cleaned.append(float(value))
|
|
validated.append(cleaned)
|
|
return validated
|
|
|
|
def embed(self, texts: list[str]) -> list[list[float]]:
|
|
out: list[list[float]] = []
|
|
for i in range(0, len(texts), self.batch_size):
|
|
out.extend(self._post(texts[i : i + self.batch_size]))
|
|
return out
|
|
|
|
def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
|
return self.embed(texts)
|
|
|
|
def embed_query(self, text: str) -> list[float]:
|
|
return self.embed([text])[0]
|
|
|
|
|
|
OllamaEmbeddings = OllamaInternalEmbeddings
|