Files

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