137 lines
4.6 KiB
Python
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,
|
|
)
|