Update tests and client for explicit embed_query/embed_documents API

This commit is contained in:
Yiorgis Gozadinos 2025-12-26 11:26:56 +02:00
parent 132b8a36bc
commit 3852e961b9
No known key found for this signature in database
4 changed files with 93 additions and 140 deletions

View file

@ -1407,7 +1407,7 @@ class HaikuRAG:
# Generate new embeddings using contextualize for consistency # Generate new embeddings using contextualize for consistency
texts = contextualize(chunks) texts = contextualize(chunks)
embeddings = await self.chunk_repository.embedder.embed(texts) embeddings = await self.chunk_repository.embedder.embed_documents(texts)
# Build updated records # Build updated records
for chunk, embedding in zip(chunks, embeddings): for chunk, embedding in zip(chunks, embeddings):

View file

@ -3,20 +3,16 @@ import os
import numpy as np import numpy as np
import pytest import pytest
from haiku.rag.config import Config from haiku.rag.config import AppConfig, EmbeddingModelConfig, EmbeddingsConfig
from haiku.rag.embeddings import contextualize, embed_chunks from haiku.rag.embeddings import contextualize, embed_chunks, get_embedder
from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
from haiku.rag.embeddings.vllm import Embedder as VLLMEmbedder
from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.chunk import Chunk
OPENAI_AVAILABLE = bool(os.getenv("OPENAI_API_KEY")) OPENAI_AVAILABLE = bool(os.getenv("OPENAI_API_KEY"))
VOYAGEAI_AVAILABLE = bool(os.getenv("VOYAGE_API_KEY")) VOYAGEAI_AVAILABLE = bool(os.getenv("VOYAGE_API_KEY"))
VLLM_EMBEDDINGS_AVAILABLE = bool(Config.providers.vllm.embeddings_base_url)
# Calculate cosine similarity
def similarities(embeddings, test_embedding): def similarities(embeddings, test_embedding):
"""Calculate cosine similarity between embeddings and a test embedding."""
return [ return [
np.dot(embedding, test_embedding) np.dot(embedding, test_embedding)
/ (np.linalg.norm(embedding) * np.linalg.norm(test_embedding)) / (np.linalg.norm(embedding) * np.linalg.norm(test_embedding))
@ -26,35 +22,41 @@ def similarities(embeddings, test_embedding):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ollama_embedder(): async def test_ollama_embedder():
embedder = OllamaEmbedder("mxbai-embed-large", 1024) """Test Ollama embedder via pydantic-ai."""
config = AppConfig(
embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig(
provider="ollama", name="mxbai-embed-large", vector_dim=1024
)
)
)
embedder = get_embedder(config)
phrases = [ phrases = [
"I enjoy eating great food.", "I enjoy eating great food.",
"Python is my favorite programming language.", "Python is my favorite programming language.",
"I love to travel and see new places.", "I love to travel and see new places.",
] ]
# Test batch embedding # Test batch embedding (documents)
embeddings = await embedder.embed(phrases) embeddings = await embedder.embed_documents(phrases)
assert isinstance(embeddings, list) assert isinstance(embeddings, list)
assert len(embeddings) == 3 assert len(embeddings) == 3
assert all(isinstance(emb, list) for emb in embeddings) assert all(isinstance(emb, list) for emb in embeddings)
embeddings = [np.array(emb) for emb in embeddings] embeddings = [np.array(emb) for emb in embeddings]
# Test query embedding
test_phrase = "I am going for a camping trip." test_phrase = "I am going for a camping trip."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed_query(test_phrase)
sims = similarities(embeddings, test_embedding) sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[2] assert max(sims) == sims[2]
test_phrase = "When is dinner ready?" test_phrase = "When is dinner ready?"
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed_query(test_phrase)
sims = similarities(embeddings, test_embedding) sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[0] assert max(sims) == sims[0]
test_phrase = "I work as a software developer." test_phrase = "I work as a software developer."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed_query(test_phrase)
sims = similarities(embeddings, test_embedding) sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[1] assert max(sims) == sims[1]
@ -62,35 +64,41 @@ async def test_ollama_embedder():
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.skipif(not OPENAI_AVAILABLE, reason="OpenAI API key not available") @pytest.mark.skipif(not OPENAI_AVAILABLE, reason="OpenAI API key not available")
async def test_openai_embedder(): async def test_openai_embedder():
embedder = OpenAIEmbedder("text-embedding-3-small", 1536) """Test OpenAI embedder via pydantic-ai."""
config = AppConfig(
embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig(
provider="openai", name="text-embedding-3-small", vector_dim=1536
)
)
)
embedder = get_embedder(config)
phrases = [ phrases = [
"I enjoy eating great food.", "I enjoy eating great food.",
"Python is my favorite programming language.", "Python is my favorite programming language.",
"I love to travel and see new places.", "I love to travel and see new places.",
] ]
# Test batch embedding # Test batch embedding (documents)
embeddings = await embedder.embed(phrases) embeddings = await embedder.embed_documents(phrases)
assert isinstance(embeddings, list) assert isinstance(embeddings, list)
assert len(embeddings) == 3 assert len(embeddings) == 3
assert all(isinstance(emb, list) for emb in embeddings) assert all(isinstance(emb, list) for emb in embeddings)
embeddings = [np.array(emb) for emb in embeddings] embeddings = [np.array(emb) for emb in embeddings]
# Test query embedding
test_phrase = "I am going for a camping trip." test_phrase = "I am going for a camping trip."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed_query(test_phrase)
sims = similarities(embeddings, test_embedding) sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[2] assert max(sims) == sims[2]
test_phrase = "When is dinner ready?" test_phrase = "When is dinner ready?"
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed_query(test_phrase)
sims = similarities(embeddings, test_embedding) sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[0] assert max(sims) == sims[0]
test_phrase = "I work as a software developer." test_phrase = "I work as a software developer."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed_query(test_phrase)
sims = similarities(embeddings, test_embedding) sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[1] assert max(sims) == sims[1]
@ -98,38 +106,42 @@ async def test_openai_embedder():
@pytest.mark.asyncio @pytest.mark.asyncio
@pytest.mark.skipif(not VOYAGEAI_AVAILABLE, reason="VoyageAI API key not available") @pytest.mark.skipif(not VOYAGEAI_AVAILABLE, reason="VoyageAI API key not available")
async def test_voyageai_embedder(): async def test_voyageai_embedder():
"""Test VoyageAI embedder."""
try: try:
from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder config = AppConfig(
embeddings=EmbeddingsConfig(
embedder = VoyageAIEmbedder("voyage-3.5", 1024) model=EmbeddingModelConfig(
provider="voyageai", name="voyage-3.5", vector_dim=1024
)
)
)
embedder = get_embedder(config)
phrases = [ phrases = [
"I enjoy eating great food.", "I enjoy eating great food.",
"Python is my favorite programming language.", "Python is my favorite programming language.",
"I love to travel and see new places.", "I love to travel and see new places.",
] ]
# Test batch embedding # Test batch embedding (documents)
embeddings = await embedder.embed(phrases) embeddings = await embedder.embed_documents(phrases)
assert isinstance(embeddings, list) assert isinstance(embeddings, list)
assert len(embeddings) == 3 assert len(embeddings) == 3
assert all(isinstance(emb, list) for emb in embeddings) assert all(isinstance(emb, list) for emb in embeddings)
embeddings = [np.array(emb) for emb in embeddings] embeddings = [np.array(emb) for emb in embeddings]
# Test query embedding
test_phrase = "I am going for a camping trip." test_phrase = "I am going for a camping trip."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed_query(test_phrase)
sims = similarities(embeddings, test_embedding) sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[2] assert max(sims) == sims[2]
test_phrase = "When is dinner ready?" test_phrase = "When is dinner ready?"
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed_query(test_phrase)
sims = similarities(embeddings, test_embedding) sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[0] assert max(sims) == sims[0]
test_phrase = "I work as a software developer." test_phrase = "I work as a software developer."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed_query(test_phrase)
sims = similarities(embeddings, test_embedding) sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[1] assert max(sims) == sims[1]
@ -137,44 +149,6 @@ async def test_voyageai_embedder():
pytest.skip("VoyageAI package not installed") pytest.skip("VoyageAI package not installed")
@pytest.mark.asyncio
@pytest.mark.skipif(
not VLLM_EMBEDDINGS_AVAILABLE, reason="vLLM embeddings server not configured"
)
async def test_vllm_embedder():
embedder = VLLMEmbedder("mixedbread-ai/mxbai-embed-large-v1", 512)
phrases = [
"I enjoy eating great food.",
"Python is my favorite programming language.",
"I love to travel and see new places.",
]
# Test batch embedding
embeddings = await embedder.embed(phrases)
assert isinstance(embeddings, list)
assert len(embeddings) == 3
assert all(isinstance(emb, list) for emb in embeddings)
embeddings = [np.array(emb) for emb in embeddings]
test_phrase = "I am going for a camping trip."
test_embedding = await embedder.embed(test_phrase)
sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[2]
test_phrase = "When is dinner ready?"
test_embedding = await embedder.embed(test_phrase)
sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[0]
test_phrase = "I work as a software developer."
test_embedding = await embedder.embed(test_phrase)
sims = similarities(embeddings, test_embedding)
assert max(sims) == sims[1]
def test_contextualize_with_headings(): def test_contextualize_with_headings():
"""Test that contextualize prepends headings to chunk content.""" """Test that contextualize prepends headings to chunk content."""
chunks = [ chunks = [

View file

@ -4,16 +4,14 @@ from haiku.rag.config import (
AppConfig, AppConfig,
EmbeddingModelConfig, EmbeddingModelConfig,
EmbeddingsConfig, EmbeddingsConfig,
LMStudioConfig,
OllamaConfig, OllamaConfig,
ProvidersConfig, ProvidersConfig,
VLLMConfig,
) )
from haiku.rag.embeddings import get_embedder from haiku.rag.embeddings import get_embedder
def test_embedder_uses_config_from_get_embedder(): def test_ollama_embedder_uses_config():
"""Test that embedders use the config passed to get_embedder.""" """Test that Ollama embedder uses the config passed to get_embedder."""
custom_config = AppConfig( custom_config = AppConfig(
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig( model=EmbeddingModelConfig(
@ -22,41 +20,16 @@ def test_embedder_uses_config_from_get_embedder():
), ),
providers=ProvidersConfig( providers=ProvidersConfig(
ollama=OllamaConfig(base_url="http://custom-ollama:8080"), ollama=OllamaConfig(base_url="http://custom-ollama:8080"),
vllm=VLLMConfig(embeddings_base_url="http://custom-vllm:9000"),
), ),
) )
embedder = get_embedder(custom_config) embedder = get_embedder(custom_config)
assert embedder._model == "custom-model"
assert embedder._vector_dim == 512 assert embedder._vector_dim == 512
assert embedder._config.providers.ollama.base_url == "http://custom-ollama:8080"
def test_vllm_embedder_uses_config():
"""Test that vllm embedder uses the config passed to get_embedder."""
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig(
provider="vllm", name="custom-vllm-model", vector_dim=768
),
),
providers=ProvidersConfig(
vllm=VLLMConfig(embeddings_base_url="http://custom-vllm:9001"),
),
)
embedder = get_embedder(custom_config)
assert embedder._model == "custom-vllm-model"
assert embedder._vector_dim == 768
assert (
embedder._config.providers.vllm.embeddings_base_url == "http://custom-vllm:9001"
)
def test_openai_embedder_uses_config(): def test_openai_embedder_uses_config():
"""Test that openai embedder uses the config passed to get_embedder.""" """Test that OpenAI embedder uses the config passed to get_embedder."""
custom_config = AppConfig( custom_config = AppConfig(
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig( model=EmbeddingModelConfig(
@ -67,48 +40,68 @@ def test_openai_embedder_uses_config():
embedder = get_embedder(custom_config) embedder = get_embedder(custom_config)
assert embedder._model == "text-embedding-3-large"
assert embedder._vector_dim == 3072 assert embedder._vector_dim == 3072
assert embedder._config == custom_config
def test_lm_studio_embedder_uses_config(): def test_openai_embedder_with_base_url():
"""Test that lm_studio embedder uses the config passed to get_embedder.""" """Test that OpenAI embedder uses custom base_url for vLLM/LM Studio."""
custom_config = AppConfig( custom_config = AppConfig(
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig( model=EmbeddingModelConfig(
provider="lm_studio", name="custom-lm-studio-model", vector_dim=1024 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_cohere_embedder_uses_config():
"""Test that Cohere embedder uses the config passed to get_embedder."""
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig(
provider="cohere", name="embed-v4.0", vector_dim=1024
), ),
), ),
providers=ProvidersConfig(
lm_studio=LMStudioConfig(base_url="http://custom-lmstudio:5678"),
),
) )
embedder = get_embedder(custom_config) embedder = get_embedder(custom_config)
assert embedder._model == "custom-lm-studio-model"
assert embedder._vector_dim == 1024 assert embedder._vector_dim == 1024
assert (
embedder._config.providers.lm_studio.base_url == "http://custom-lmstudio:5678"
)
@pytest.mark.skipif( def test_sentence_transformers_embedder_uses_config():
True, reason="VoyageAI is an optional dependency, may not be installed" """Test that SentenceTransformers embedder uses the config."""
)
def test_voyageai_embedder_uses_config():
"""Test that voyageai embedder uses the config passed to get_embedder."""
custom_config = AppConfig( custom_config = AppConfig(
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig( model=EmbeddingModelConfig(
provider="voyageai", name="voyage-large-2", vector_dim=1536 provider="sentence-transformers",
name="all-MiniLM-L6-v2",
vector_dim=384,
), ),
), ),
) )
embedder = get_embedder(custom_config) embedder = get_embedder(custom_config)
assert embedder._model == "voyage-large-2" assert embedder._vector_dim == 384
assert embedder._vector_dim == 1536
assert embedder._config == custom_config
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)

View file

@ -270,20 +270,6 @@ def test_get_model_bedrock_with_thinking():
assert isinstance(result, BedrockConverseModel) assert isinstance(result, BedrockConverseModel)
def test_get_model_vllm():
"""Test get_model returns OpenAIChatModel for vLLM."""
model_config = ModelConfig(provider="vllm", name="Qwen/Qwen3-4B")
result = get_model(model_config)
assert isinstance(result, OpenAIChatModel)
def test_get_model_vllm_with_thinking():
"""Test get_model configures thinking for gpt-oss on vLLM."""
model_config = ModelConfig(provider="vllm", name="gpt-oss", enable_thinking=False)
result = get_model(model_config)
assert isinstance(result, OpenAIChatModel)
def test_get_model_unknown_provider(): def test_get_model_unknown_provider():
"""Test get_model returns string format for unknown providers.""" """Test get_model returns string format for unknown providers."""
model_config = ModelConfig(provider="mistral", name="mistral-large-latest") model_config = ModelConfig(provider="mistral", name="mistral-large-latest")