feat: use internal ollama embeddings
This commit is contained in:
@@ -1,16 +1,16 @@
|
||||
import pytest
|
||||
|
||||
from tht.adapters.evidence import FilesystemEvidenceSource, HttpManifestEvidenceSource
|
||||
from tht.adapters.factory import build_evidence_sources
|
||||
from tht.config import (
|
||||
ConfigError,
|
||||
PgvectorDirectConfig,
|
||||
PostgresDwhConfig,
|
||||
ThothRestDwhConfig,
|
||||
ThothVectorHttpConfig,
|
||||
workspace_id_for_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):
|
||||
@@ -213,6 +213,95 @@ embeddings: {base_url: http://ollama:11434, dim: 768}
|
||||
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):
|
||||
common = """
|
||||
dwh:
|
||||
|
||||
@@ -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"])
|
||||
Reference in New Issue
Block a user