check_source_accessible narrowed its handler to ValueError, but Path.exists re-raises errno values outside its ignored set (EACCES, ENAMETOOLONG). Those were swallowed before and now escaped into the rebuild sweep the guard exists to protect. Catch OSError too. Restore the arity guard in _common_path_prefix: without it an empty list raises from min() and a single label yields a prefix covering the whole path. Two tests would have hung rather than failed on regression (the vacuum skip and the protected-wait cancellation); both are now bounded. The import vacuum test raced against the done-callback that discards the task, and now spies on the call instead, with a negative control. Replace assertions that could not fail: blank-query search against an empty corpus, a batch flush counted against an empty table, a picture description asserting its own input state, and an FS scheme check with nothing on disk to resolve. The get_model matrix asserted only the returned type across 26 cases and now pins the per-provider settings. The three batching tests now count flushes, which revealed embed-only writes through chunks_table.add rather than _flush_rebuild_batch.
197 lines
6.3 KiB
Python
197 lines
6.3 KiB
Python
import pytest
|
|
|
|
from haiku.rag.config import (
|
|
AppConfig,
|
|
EmbeddingModelConfig,
|
|
EmbeddingsConfig,
|
|
OllamaConfig,
|
|
ProvidersConfig,
|
|
)
|
|
from haiku.rag.embeddings import get_embedder
|
|
|
|
|
|
def test_ollama_embedder_uses_config():
|
|
"""Test that Ollama embedder uses the config passed to get_embedder."""
|
|
custom_config = AppConfig(
|
|
embeddings=EmbeddingsConfig(
|
|
model=EmbeddingModelConfig(
|
|
provider="ollama", name="custom-model", vector_dim=512
|
|
),
|
|
),
|
|
providers=ProvidersConfig(
|
|
ollama=OllamaConfig(base_url="http://custom-ollama:8080"),
|
|
),
|
|
)
|
|
|
|
embedder = get_embedder(custom_config)
|
|
|
|
assert embedder._vector_dim == 512
|
|
|
|
|
|
def test_openai_embedder_with_base_url():
|
|
"""Test that OpenAI embedder uses custom base_url for vLLM/LM Studio."""
|
|
custom_config = AppConfig(
|
|
embeddings=EmbeddingsConfig(
|
|
model=EmbeddingModelConfig(
|
|
provider="openai",
|
|
name="some-local-model",
|
|
vector_dim=768,
|
|
base_url="http://localhost:8000/v1",
|
|
),
|
|
),
|
|
)
|
|
|
|
embedder = get_embedder(custom_config)
|
|
|
|
assert embedder._vector_dim == 768
|
|
|
|
|
|
def test_sentence_transformers_embedder_uses_config():
|
|
"""Test that SentenceTransformers embedder uses the config."""
|
|
custom_config = AppConfig(
|
|
embeddings=EmbeddingsConfig(
|
|
model=EmbeddingModelConfig(
|
|
provider="sentence-transformers",
|
|
name="all-MiniLM-L6-v2",
|
|
vector_dim=384,
|
|
),
|
|
),
|
|
)
|
|
|
|
embedder = get_embedder(custom_config)
|
|
|
|
assert embedder._vector_dim == 384
|
|
|
|
|
|
def test_unsupported_provider_raises():
|
|
"""Test that unsupported provider raises ValueError."""
|
|
custom_config = AppConfig(
|
|
embeddings=EmbeddingsConfig(
|
|
model=EmbeddingModelConfig(
|
|
provider="unsupported-provider", name="model", vector_dim=512
|
|
),
|
|
),
|
|
)
|
|
|
|
with pytest.raises(ValueError, match="Unsupported embedding provider"):
|
|
get_embedder(custom_config)
|
|
|
|
|
|
def test_ollama_embedder_appends_v1_when_missing():
|
|
"""Per-model base_url without /v1 should get it appended for Ollama."""
|
|
config = AppConfig(
|
|
embeddings=EmbeddingsConfig(
|
|
model=EmbeddingModelConfig(
|
|
provider="ollama",
|
|
name="qwen3-embedding:4b",
|
|
vector_dim=2560,
|
|
base_url="http://my-ollama:11434",
|
|
),
|
|
),
|
|
)
|
|
embedder = get_embedder(config)
|
|
pa_model = embedder._embedder._model # type: ignore[union-attr] # ty: ignore[unresolved-attribute]
|
|
assert str(pa_model.base_url).rstrip("/").endswith("/v1") # type: ignore[union-attr] # ty: ignore[unresolved-attribute]
|
|
|
|
|
|
def test_ollama_embedder_does_not_double_append_v1():
|
|
"""If the user already includes /v1 we leave it alone."""
|
|
config = AppConfig(
|
|
embeddings=EmbeddingsConfig(
|
|
model=EmbeddingModelConfig(
|
|
provider="ollama",
|
|
name="qwen3-embedding:4b",
|
|
vector_dim=2560,
|
|
base_url="http://my-ollama:11434/v1",
|
|
),
|
|
),
|
|
)
|
|
embedder = get_embedder(config)
|
|
pa_model = embedder._embedder._model # type: ignore[union-attr] # ty: ignore[unresolved-attribute]
|
|
url = str(pa_model.base_url).rstrip("/") # type: ignore[union-attr] # ty: ignore[unresolved-attribute]
|
|
assert url.endswith("/v1")
|
|
assert not url.endswith("/v1/v1")
|
|
|
|
|
|
def test_vllm_embedder_appends_v1_when_missing():
|
|
"""vLLM's chat-completions endpoint also lives under /v1. A user who
|
|
forgets the suffix would otherwise POST to <host>/embeddings and get
|
|
a 404 — match the Ollama behavior and append it."""
|
|
config = AppConfig(
|
|
embeddings=EmbeddingsConfig(
|
|
model=EmbeddingModelConfig(
|
|
provider="vllm",
|
|
name="Qwen/Qwen3-VL-Embedding-8B",
|
|
vector_dim=4096,
|
|
base_url="http://my-vllm:8000",
|
|
),
|
|
),
|
|
)
|
|
embedder = get_embedder(config)
|
|
base_url = embedder._base_url # type: ignore[attr-defined] # ty: ignore[unresolved-attribute]
|
|
assert base_url.endswith("/v1")
|
|
|
|
|
|
def test_vllm_embedder_does_not_double_append_v1():
|
|
"""If the user already includes /v1 we leave it alone."""
|
|
config = AppConfig(
|
|
embeddings=EmbeddingsConfig(
|
|
model=EmbeddingModelConfig(
|
|
provider="vllm",
|
|
name="Qwen/Qwen3-VL-Embedding-8B",
|
|
vector_dim=4096,
|
|
base_url="http://my-vllm:8000/v1",
|
|
),
|
|
),
|
|
)
|
|
embedder = get_embedder(config)
|
|
base_url = embedder._base_url.rstrip("/") # type: ignore[attr-defined] # ty: ignore[unresolved-attribute]
|
|
assert base_url.endswith("/v1")
|
|
assert not base_url.endswith("/v1/v1")
|
|
|
|
|
|
def test_vector_dim_property_reports_configured_dimension():
|
|
from haiku.rag.embeddings import EmbedderWrapper
|
|
|
|
assert EmbedderWrapper(embedder=None, vector_dim=512).vector_dim == 512
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"provider,env_var",
|
|
[("voyageai", "VOYAGE_API_KEY"), ("cohere", "CO_API_KEY")],
|
|
)
|
|
def test_saas_providers_are_wired_without_a_request(monkeypatch, provider, env_var):
|
|
"""Construction wires the SDK and reports the configured dimension."""
|
|
monkeypatch.setenv(env_var, "test-key")
|
|
config = AppConfig(
|
|
embeddings=EmbeddingsConfig(
|
|
model=EmbeddingModelConfig(
|
|
provider=provider, name="some-model", vector_dim=1024
|
|
),
|
|
),
|
|
)
|
|
|
|
embedder = get_embedder(config)
|
|
|
|
assert embedder.vector_dim == 1024
|
|
assert embedder.supports_images is False
|
|
# The provider and model reach the underlying pydantic-ai embedder.
|
|
assert embedder._embedder._model == f"{provider}:some-model" # ty: ignore[unresolved-attribute]
|
|
|
|
|
|
def test_cohere_floats_rejects_missing_embeddings():
|
|
from types import SimpleNamespace
|
|
|
|
from haiku.rag.embeddings.cohere import _floats
|
|
|
|
result = SimpleNamespace(embeddings=SimpleNamespace(float_=None))
|
|
|
|
with pytest.raises(ValueError, match="no float embeddings"):
|
|
_floats(result)
|
|
|
|
|
|
def test_voyageai_to_pil_rejects_unsupported_type():
|
|
from haiku.rag.embeddings.voyageai import _to_pil
|
|
|
|
with pytest.raises(TypeError, match="Unsupported image type"):
|
|
_to_pil("not an image") # ty: ignore[invalid-argument-type]
|