Files
Project-Thoth/processors/rag_ingestion/embeddings.py
T

137 lines
4.6 KiB
Python

"""Explicit provider boundary for locally generated text embeddings."""
import json
import math
from typing import Protocol
from urllib.error import HTTPError, URLError
from urllib.request import Request, urlopen
from .config import EmbeddingConfig
from .models import Chunk, Embedding
class EmbeddingConnectionError(RuntimeError):
"""The configured Oracle embedding service could not be reached."""
class EmbeddingRequestError(RuntimeError):
"""Oracle was reached but rejected or malformed the embedding request."""
class EmbeddingProvider(Protocol):
"""The narrow text-in, vector-out boundary used by Project Thoth."""
model_name: str
provider_name: str
expected_dimension: int | None
def embed(self, text: str) -> list[float]:
"""Return one numerical vector for the supplied text."""
class OllamaEmbeddingProvider:
"""Generate embeddings through Oracle's Ollama-compatible HTTP API."""
provider_name = "Ollama"
def __init__(
self,
model_name: str,
base_url: str,
expected_dimension: int | None,
timeout_seconds: float = 60.0,
) -> None:
self.model_name = model_name
self.base_url = base_url.rstrip("/")
self.expected_dimension = expected_dimension
self.timeout_seconds = timeout_seconds
def embed(self, text: str) -> list[float]:
"""Send exactly ``text`` to Ollama and return its single vector."""
payload = json.dumps(
{
"model": self.model_name,
"input": text,
# An oversized input should fail visibly rather than be changed
# without the developer knowing what the model received.
"truncate": False,
}
).encode("utf-8")
request = Request(
f"{self.base_url}/api/embed",
data=payload,
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urlopen(request, timeout=self.timeout_seconds) as response:
result = json.load(response)
except HTTPError as error:
details = error.read().decode("utf-8", errors="replace")
raise EmbeddingRequestError(
f"Ollama embedding request failed with HTTP {error.code}: {details}"
) from error
except URLError as error:
raise EmbeddingConnectionError(
f"Cannot reach Ollama embedding service at {self.base_url}: "
f"{error.reason}"
) from error
embeddings = result.get("embeddings")
if not isinstance(embeddings, list) or len(embeddings) != 1:
raise EmbeddingRequestError(
"Ollama response did not contain exactly one embedding"
)
vector = embeddings[0]
if not isinstance(vector, list) or not vector:
raise EmbeddingRequestError(
"Ollama returned an empty or invalid embedding vector"
)
if not all(isinstance(value, (int, float)) for value in vector):
raise EmbeddingRequestError(
"Ollama embedding vector contains non-numerical values"
)
return [float(value) for value in vector]
def embed_chunk(chunk: Chunk, provider: EmbeddingProvider) -> Embedding:
"""Embed only ``chunk.text`` and retain its source linkage."""
vector = provider.embed(chunk.text)
if (
provider.expected_dimension is not None
and len(vector) != provider.expected_dimension
):
raise ValueError(
f"model {provider.model_name} returned {len(vector)} dimensions; "
f"expected {provider.expected_dimension}"
)
if not all(math.isfinite(value) for value in vector):
raise ValueError("embedding vector contains a non-finite value")
return Embedding(
chunk_id=chunk.chunk_id,
document_id=chunk.document_id,
embedding_model=provider.model_name,
embedding_provider=provider.provider_name,
embedding_dimension=len(vector),
vector=tuple(vector),
)
def provider_from_config(config: EmbeddingConfig) -> OllamaEmbeddingProvider:
"""Construct the provider only from validated centralized configuration."""
validated = config.require_live_embedding()
assert validated.oracle_base_url is not None
assert validated.embedding_model is not None
return OllamaEmbeddingProvider(
model_name=validated.embedding_model,
base_url=validated.oracle_base_url,
expected_dimension=validated.embedding_dimension,
)