feat: use internal ollama embeddings

This commit is contained in:
2026-08-08 17:29:36 +02:00
parent 9f104171b6
commit 2911e008d1
7 changed files with 455 additions and 92 deletions
@@ -0,0 +1,92 @@
# Task 4 Report — Narrow harness embedding configuration to internal Ollama
## Status
Implemented on 2026-08-08 in `/Users/mp/projects/ThothII/.worktrees/git-workspace-registry`.
## RED evidence
Command:
```bash
cd harness
./.venv/bin/pytest tests/test_internal_embeddings.py tests/test_config_resources.py -q
```
Observed before implementation:
- exit code `1`
- `10 failed, 10 passed`
- failures proved the missing `OllamaInternalEmbeddings` client and missing internal-only config validation
Representative failures:
- `ImportError: cannot import name 'OllamaInternalEmbeddings'`
- `AttributeError: 'EmbeddingsConfig' object has no attribute 'provider'`
- config tests `DID NOT RAISE ConfigError` for external provider, API key, and non-private base URL
## GREEN evidence
Focused behavior suite:
```bash
cd harness
./.venv/bin/pytest tests/test_internal_embeddings.py tests/test_config_resources.py -q
```
- exit code `0`
- `20 passed`
Relevant harness verification:
```bash
cd harness
./.venv/bin/pytest tests/test_internal_embeddings.py tests/test_config_resources.py tests/test_ollama_ensure.py -q
```
- exit code `0`
- `36 passed, 2 warnings`
Changed-file lint:
```bash
cd harness
./.venv/bin/ruff check tht/config.py tht/config_compat.py tht/vectorstore/embeddings.py tht/cli/ollama_cmd.py tests/test_config_resources.py tests/test_internal_embeddings.py
```
- exit code `0`
- `All checks passed!`
Patch hygiene:
```bash
git diff --check
```
- exit code `0`
## What changed
- translated schema-v3 `resources.embeddings` into the harness-compatible embedding config view
- validated the internal embedding contract only for that runtime-owned `resources.embeddings` path:
- provider must be `ollama_internal`
- model must be `qwen3-embedding:0.6b`
- dimensions must be `1024`
- base URL must be `http://embedding:11434` or loopback HTTP on port `11434`
- extra fields like `api_key` are rejected
- replaced the active embed client with `OllamaInternalEmbeddings`, using one bounded `/api/embed` request per batch
- removed task/query prefix rewriting from the active embedding path
- validated response count, vector dimension, and finite numeric values before returning embeddings
- kept `tht ollama ensure --json` stdout pristine while warming through the internal client
## Self-review
- kept changes inside the brief-listed files
- preserved DWH and session-persistence behavior
- preserved the legacy `OllamaEmbeddings` import path as an alias to avoid unrelated call-site churn
## Concerns
- the focused harness verification still emits two pre-existing warnings:
- `DeprecationWarning` from `testcontainers.postgres`
- `FutureWarning` because `resources` currently flows through the legacy config translation path
+92 -3
View File
@@ -1,16 +1,16 @@
import pytest import pytest
from tht.adapters.evidence import FilesystemEvidenceSource, HttpManifestEvidenceSource
from tht.adapters.factory import build_evidence_sources
from tht.config import ( from tht.config import (
ConfigError, ConfigError,
PgvectorDirectConfig, PgvectorDirectConfig,
PostgresDwhConfig, PostgresDwhConfig,
ThothRestDwhConfig, ThothRestDwhConfig,
ThothVectorHttpConfig, ThothVectorHttpConfig,
workspace_id_for_config,
load_config, load_config,
workspace_id_for_config,
) )
from tht.adapters.evidence import FilesystemEvidenceSource, HttpManifestEvidenceSource
from tht.adapters.factory import build_evidence_sources
def test_direct_vector_passwords_load_from_file_references(monkeypatch, tmp_path): def test_direct_vector_passwords_load_from_file_references(monkeypatch, tmp_path):
@@ -213,6 +213,95 @@ embeddings: {base_url: http://ollama:11434, dim: 768}
assert cfg.vectors.writer.api_key == "writer" assert cfg.vectors.writer.api_key == "writer"
def test_accepts_only_internal_ollama_embedding_contract(tmp_path):
workspace = tmp_path / "workspace.yaml"
workspace.write_text(
"""
dwh:
type: postgres_direct
connection: {database: analytics, schema: mart, user: reader, password: secret}
resources:
embeddings:
provider: ollama_internal
base_url: http://embedding:11434
model: qwen3-embedding:0.6b
dimensions: 1024
"""
)
cfg = load_config(workspace)
assert cfg.embeddings.provider == "ollama_internal"
assert cfg.embeddings.base_url == "http://embedding:11434"
assert cfg.embeddings.model == "qwen3-embedding:0.6b"
assert cfg.embeddings.dim == 1024
@pytest.mark.parametrize(
("snippet", "pattern"),
[
(
"""
resources:
embeddings:
provider: openai_compatible
base_url: http://embedding:11434
model: qwen3-embedding:0.6b
dimensions: 1024
""",
"ollama_internal|provider",
),
(
"""
resources:
embeddings:
provider: ollama_internal
base_url: http://embedding:11434
model: qwen3-embedding:0.6b
dimensions: 1024
api_key: secret
""",
"api_key|extra",
),
(
"""
resources:
embeddings:
provider: ollama_internal
base_url: https://embedding:11434
model: qwen3-embedding:0.6b
dimensions: 1024
""",
"base_url|internal|private|host",
),
(
"""
resources:
embeddings:
provider: ollama_internal
base_url: http://example.com:11434
model: qwen3-embedding:0.6b
dimensions: 1024
""",
"base_url|internal|private|host",
),
],
)
def test_rejects_external_embedding_configuration(tmp_path, snippet, pattern):
workspace = tmp_path / "workspace.yaml"
workspace.write_text(
"""
dwh:
type: postgres_direct
connection: {database: analytics, schema: mart, user: reader, password: secret}
"""
+ snippet
)
with pytest.raises(ConfigError, match=pattern):
load_config(workspace)
def test_builds_typed_evidence_sources_and_keeps_legacy_compatible(tmp_path): def test_builds_typed_evidence_sources_and_keeps_legacy_compatible(tmp_path):
common = """ common = """
dwh: dwh:
+139
View File
@@ -0,0 +1,139 @@
import math
import pytest
from tht.config import EmbeddingsConfig
from tht.vectorstore.embeddings import EmbeddingsError
class _Response:
def __init__(self, payload, status_code=200):
self._payload = payload
self.status_code = status_code
def raise_for_status(self):
if self.status_code >= 400:
raise RuntimeError(f"http {self.status_code}")
def json(self):
return self._payload
class _Session:
def __init__(self, responses):
self._responses = list(responses)
self.calls = []
def post(self, url, json, timeout):
self.calls.append({"url": url, "json": json, "timeout": timeout})
if not self._responses:
raise AssertionError("unexpected extra request")
return self._responses.pop(0)
def _vector(value: float, *, dim: int = 1024):
return [value] * dim
def test_internal_embeddings_posts_model_and_batch_input_without_prefixes():
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
session = _Session([_Response({"embeddings": [_vector(1.0), _vector(2.0)]})])
embedder = OllamaInternalEmbeddings(
EmbeddingsConfig(
provider="ollama_internal",
base_url="http://embedding:11434",
model="qwen3-embedding:0.6b",
dim=1024,
batch_size=2,
timeout=9,
connect_timeout=4,
),
session=session,
)
vectors = embedder.embed(["alpha", "beta"])
assert vectors == [_vector(1.0), _vector(2.0)]
assert session.calls == [{
"url": "http://embedding:11434/api/embed",
"json": {"model": "qwen3-embedding:0.6b", "input": ["alpha", "beta"]},
"timeout": (4, 9),
}]
def test_internal_embeddings_batching_returns_1024d_vectors():
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
session = _Session([
_Response({"embeddings": [_vector(1.0), _vector(2.0)]}),
_Response({"embeddings": [_vector(3.0)]}),
])
embedder = OllamaInternalEmbeddings(
EmbeddingsConfig(
provider="ollama_internal",
base_url="http://embedding:11434",
model="qwen3-embedding:0.6b",
dim=1024,
batch_size=2,
),
session=session,
)
vectors = embedder.embed(["one", "two", "three"])
assert [len(vector) for vector in vectors] == [1024, 1024, 1024]
assert [vector[0] for vector in vectors] == [1.0, 2.0, 3.0]
def test_internal_embeddings_reject_count_mismatch():
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
embedder = OllamaInternalEmbeddings(
EmbeddingsConfig(
provider="ollama_internal",
base_url="http://embedding:11434",
model="qwen3-embedding:0.6b",
dim=1024,
),
session=_Session([_Response({"embeddings": [_vector(1.0)]})]),
)
with pytest.raises(EmbeddingsError, match="count|numero"):
embedder.embed(["alpha", "beta"])
def test_internal_embeddings_reject_dimension_mismatch():
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
embedder = OllamaInternalEmbeddings(
EmbeddingsConfig(
provider="ollama_internal",
base_url="http://embedding:11434",
model="qwen3-embedding:0.6b",
dim=1024,
),
session=_Session([_Response({"embeddings": [[1.0] * 8]})]),
)
with pytest.raises(EmbeddingsError, match="dimensione|dimension"):
embedder.embed(["alpha"])
def test_internal_embeddings_reject_non_finite_values():
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
bad = _vector(0.0)
bad[10] = math.nan
embedder = OllamaInternalEmbeddings(
EmbeddingsConfig(
provider="ollama_internal",
base_url="http://embedding:11434",
model="qwen3-embedding:0.6b",
dim=1024,
),
session=_Session([_Response({"embeddings": [bad]})]),
)
with pytest.raises(EmbeddingsError, match="finite|finit"):
embedder.embed(["alpha"])
+4 -4
View File
@@ -39,16 +39,16 @@ def _installed_models(base_url: str, timeout: float = 5.0) -> set[str]:
def _start(start_cmd: list[str]) -> None: def _start(start_cmd: list[str]) -> None:
# Detached so the server outlives this short-lived CLI process. # Detached so the server outlives this short-lived CLI process.
subprocess.Popen( # noqa: S603 subprocess.Popen(
start_cmd, start_new_session=True, start_cmd, start_new_session=True,
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
) )
def _warm(cfg) -> None: def _warm(cfg) -> None:
from tht.vectorstore.embeddings import OllamaEmbeddings from tht.vectorstore.embeddings import OllamaInternalEmbeddings
OllamaEmbeddings(cfg.embeddings).embed_query("ping") OllamaInternalEmbeddings(cfg.embeddings).embed_query("ping")
def _model_present(installed: set[str], model: str) -> bool: def _model_present(installed: set[str], model: str) -> bool:
@@ -104,7 +104,7 @@ def ensure_ollama(
if probe(base_url): if probe(base_url):
up = True up = True
break break
except Exception: # noqa: BLE001 - transient during startup; keep polling except Exception: # noqa: BLE001,S112 - transient during startup; keep polling
continue continue
if not up: if not up:
return {"ok": False, "stage": "server", return {"ok": False, "stage": "server",
+69 -2
View File
@@ -1,11 +1,13 @@
import os import os
import re import re
import warnings import warnings
from ipaddress import ip_address
from pathlib import Path from pathlib import Path
from typing import Annotated, Any, Literal from typing import Annotated, Any, Literal
from urllib.parse import urlparse
import yaml import yaml
from pydantic import BaseModel, Field, PrivateAttr, SecretStr, model_validator, ValidationError from pydantic import BaseModel, Field, PrivateAttr, SecretStr, ValidationError, model_validator
from tht.config_compat import translate_legacy_config from tht.config_compat import translate_legacy_config
@@ -276,15 +278,18 @@ class EvidenceSourcesConfig(BaseModel):
class EmbeddingsConfig(BaseModel): class EmbeddingsConfig(BaseModel):
provider: str = "ollama_internal"
base_url: str base_url: str
model: str = "nomic-embed-text-v2-moe" model: str = "nomic-embed-text-v2-moe"
dim: int = 768 dim: int = Field(default=768, alias="dimensions")
batch_size: int = 32 batch_size: int = 32
timeout: int = 30 timeout: int = 30
connect_timeout: int = 5 connect_timeout: int = 5
bin: str = "ollama" bin: str = "ollama"
start_cmd: list[str] | None = None start_cmd: list[str] | None = None
model_config = {"populate_by_name": True, "extra": "forbid"}
class VectorConfig(BaseModel): class VectorConfig(BaseModel):
max_chunk_chars: int = 4000 max_chunk_chars: int = 4000
@@ -391,6 +396,7 @@ def load_config(path: Path) -> Config:
if not isinstance(raw, dict): if not isinstance(raw, dict):
raise ConfigError(f"Configurazione non valida (atteso un mapping YAML): {path}") raise ConfigError(f"Configurazione non valida (atteso un mapping YAML): {path}")
expanded = _resolve_secret_files(_expand_env(raw)) expanded = _resolve_secret_files(_expand_env(raw))
_validate_internal_embedding_contract(expanded, path)
translated, used_legacy = translate_legacy_config(expanded) translated, used_legacy = translate_legacy_config(expanded)
_populate_legacy_views(translated) _populate_legacy_views(translated)
try: try:
@@ -451,6 +457,67 @@ def load_config(path: Path) -> Config:
return cfg return cfg
def _validate_internal_embedding_contract(raw: dict[str, Any], path: Path) -> None:
resources = raw.get("resources")
if not isinstance(resources, dict):
return
embeddings = resources.get("embeddings")
if not isinstance(embeddings, dict):
return
provider = embeddings.get("provider")
model = embeddings.get("model")
dimensions = embeddings.get("dimensions")
base_url = embeddings.get("base_url")
allowed = {"provider", "base_url", "model", "dimensions"}
unexpected = sorted(set(embeddings) - allowed)
if unexpected:
raise ConfigError(
f"Configurazione non valida in {path}:\n"
f"resources.embeddings non supporta: {', '.join(unexpected)}"
)
if provider != "ollama_internal":
raise ConfigError(
f"Configurazione non valida in {path}:\n"
"resources.embeddings.provider deve essere 'ollama_internal'"
)
if model != "qwen3-embedding:0.6b":
raise ConfigError(
f"Configurazione non valida in {path}:\n"
"resources.embeddings.model deve essere 'qwen3-embedding:0.6b'"
)
if dimensions != 1024:
raise ConfigError(
f"Configurazione non valida in {path}:\n"
"resources.embeddings.dimensions deve essere 1024"
)
if not _is_allowed_internal_embedding_url(base_url):
raise ConfigError(
f"Configurazione non valida in {path}:\n"
"resources.embeddings.base_url deve usare http://embedding:11434 "
"oppure un endpoint loopback di sviluppo su porta 11434"
)
def _is_allowed_internal_embedding_url(value: Any) -> bool:
if not isinstance(value, str):
return False
parsed = urlparse(value)
if parsed.scheme != "http" or not parsed.hostname or parsed.port != 11434:
return False
if parsed.params or parsed.query or parsed.fragment:
return False
if parsed.path not in ("", "/"):
return False
if parsed.hostname == "embedding":
return True
try:
host = ip_address(parsed.hostname)
except ValueError:
return parsed.hostname == "localhost"
return host.is_loopback
def _populate_legacy_views(raw: dict[str, Any]) -> None: def _populate_legacy_views(raw: dict[str, Any]) -> None:
"""Populate old Config attributes for command compatibility during migration.""" """Populate old Config attributes for command compatibility during migration."""
dwh = raw.get("dwh") dwh = raw.get("dwh")
+8 -1
View File
@@ -3,7 +3,6 @@ from __future__ import annotations
from copy import deepcopy from copy import deepcopy
from typing import Any from typing import Any
_LEGACY_RESOURCE_KEYS = { _LEGACY_RESOURCE_KEYS = {
"database", "database",
"rest", "rest",
@@ -11,6 +10,7 @@ _LEGACY_RESOURCE_KEYS = {
"vector_rest", "vector_rest",
"vector_write_rest", "vector_write_rest",
"paths", "paths",
"resources",
} }
@@ -26,6 +26,13 @@ def _as_mapping(value: Any) -> dict[str, Any] | None:
def translate_legacy_config(raw: dict[str, Any]) -> tuple[dict[str, Any], bool]: def translate_legacy_config(raw: dict[str, Any]) -> tuple[dict[str, Any], bool]:
"""Translate the legacy flat resource keys without validating their contents.""" """Translate the legacy flat resource keys without validating their contents."""
translated = deepcopy(raw) translated = deepcopy(raw)
resources = _as_mapping(raw.get("resources"))
if isinstance(resources, dict) and "embeddings" in resources and "embeddings" not in translated:
embedding = _as_mapping(resources.get("embeddings"))
if isinstance(embedding, dict):
translated["embeddings"] = embedding
if "dimensions" in translated["embeddings"] and "dim" not in translated["embeddings"]:
translated["embeddings"]["dim"] = translated["embeddings"].pop("dimensions")
legacy = any(key in raw for key in _LEGACY_RESOURCE_KEYS) legacy = any(key in raw for key in _LEGACY_RESOURCE_KEYS)
if not legacy: if not legacy:
return translated, False return translated, False
+49 -80
View File
@@ -1,104 +1,73 @@
import subprocess import math
import sys
import time
import requests import requests
from tht.config import EmbeddingsConfig from tht.config import EmbeddingsConfig
DOC_PREFIX = "search_document: "
QUERY_PREFIX = "search_query: "
_RESTART_WAIT = 8 # secondi di attesa dopo aver avviato Ollama
_RESTART_POLL = 1.0
class EmbeddingsError(Exception): class EmbeddingsError(Exception):
pass pass
class OllamaEmbeddings: class OllamaInternalEmbeddings:
"""Client embeddings via Ollama. Applica i prefissi di task richiesti da nomic v2: """Client embeddings for the installation-owned internal Ollama endpoint."""
ometterli degrada il retrieval in modo silenzioso."""
def __init__(self, cfg: EmbeddingsConfig): def __init__(self, cfg: EmbeddingsConfig, *, session: requests.Session | None = None):
self.cfg = cfg 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 _is_up(self) -> bool: def _post(self, texts: list[str]) -> list[list[float]]:
try: try:
r = requests.get( response = self._session.post(
f"{self.cfg.base_url.rstrip('/')}/api/tags", f"{self.base_url}/api/embed",
timeout=(self.cfg.connect_timeout, 5), json={"model": self.model, "input": texts},
timeout=self.timeout,
) )
return r.status_code == 200 response.raise_for_status()
except requests.RequestException: except requests.RequestException as exc:
return False raise EmbeddingsError(
f"internal Ollama embeddings request failed for model {self.model}"
def _try_restart(self) -> bool: ) from exc
"""Tenta di avviare Ollama e attende che sia raggiungibile.""" payload = response.json()
start_cmd = self.cfg.start_cmd if self.cfg.start_cmd is not None else [self.cfg.bin, "serve"] embeddings = payload.get("embeddings")
if not start_cmd: if not isinstance(embeddings, list):
return False raise EmbeddingsError("internal Ollama response is missing embeddings")
try: if len(embeddings) != len(texts):
subprocess.Popen( # noqa: S603 raise EmbeddingsError(
start_cmd, start_new_session=True, f"unexpected embedding count: {len(embeddings)} != {len(texts)}"
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
) )
except Exception: # noqa: BLE001 validated: list[list[float]] = []
return False for vector in embeddings:
print("[embeddings] Ollama non raggiungibile, avvio in corso…", file=sys.stderr) if not isinstance(vector, list):
deadline = time.monotonic() + _RESTART_WAIT raise EmbeddingsError("internal Ollama returned a non-vector embedding")
while time.monotonic() < deadline: if len(vector) != self.dim:
time.sleep(_RESTART_POLL) raise EmbeddingsError(
if self._is_up(): f"unexpected embedding dimension: {len(vector)} != {self.dim}"
return True
return False
def _post(self, url: str, batch: list[str]) -> requests.Response:
resp = requests.post(
url, json={"model": self.cfg.model, "input": batch},
timeout=(self.cfg.connect_timeout, self.cfg.timeout),
) )
resp.raise_for_status() cleaned: list[float] = []
return resp 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]]: def embed(self, texts: list[str]) -> list[list[float]]:
url = f"{self.cfg.base_url.rstrip('/')}/api/embed"
out: list[list[float]] = [] out: list[list[float]] = []
for i in range(0, len(texts), self.cfg.batch_size): for i in range(0, len(texts), self.batch_size):
batch = texts[i : i + self.cfg.batch_size] out.extend(self._post(texts[i : i + self.batch_size]))
try:
resp = self._post(url, batch)
except requests.ConnectionError:
if not self._try_restart():
raise EmbeddingsError(
f"Ollama non raggiungibile su {self.cfg.base_url} "
f"(modello {self.cfg.model}), avvio automatico fallito"
)
try:
resp = self._post(url, batch)
except requests.RequestException as e:
raise EmbeddingsError(
f"Ollama non raggiungibile su {self.cfg.base_url} "
f"(modello {self.cfg.model}): {e}"
) from e
except requests.RequestException as e:
raise EmbeddingsError(
f"Ollama non raggiungibile su {self.cfg.base_url} "
f"(modello {self.cfg.model}): {e}"
) from e
embeddings = resp.json().get("embeddings", [])
for v in embeddings:
if len(v) != self.cfg.dim:
raise EmbeddingsError(
f"dimensione embedding inattesa: {len(v)} != {self.cfg.dim} "
f"(modello {self.cfg.model})"
)
out.extend(embeddings)
return out return out
def embed_documents(self, texts: list[str]) -> list[list[float]]: def embed_documents(self, texts: list[str]) -> list[list[float]]:
return self._embed([DOC_PREFIX + t for t in texts]) return self.embed(texts)
def embed_query(self, text: str) -> list[float]: def embed_query(self, text: str) -> list[float]:
return self._embed([QUERY_PREFIX + text])[0] return self.embed([text])[0]
OllamaEmbeddings = OllamaInternalEmbeddings