140 lines
4.0 KiB
Python
140 lines
4.0 KiB
Python
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"])
|