diff --git a/haiku_rag_slim/haiku/rag/client.py b/haiku_rag_slim/haiku/rag/client.py index 241c569a..338182f9 100644 --- a/haiku_rag_slim/haiku/rag/client.py +++ b/haiku_rag_slim/haiku/rag/client.py @@ -1407,7 +1407,7 @@ class HaikuRAG: # Generate new embeddings using contextualize for consistency texts = contextualize(chunks) - embeddings = await self.chunk_repository.embedder.embed(texts) + embeddings = await self.chunk_repository.embedder.embed_documents(texts) # Build updated records for chunk, embedding in zip(chunks, embeddings): diff --git a/tests/test_embedder.py b/tests/test_embedder.py index 4d5d0bfa..2046212a 100644 --- a/tests/test_embedder.py +++ b/tests/test_embedder.py @@ -3,20 +3,16 @@ import os import numpy as np import pytest -from haiku.rag.config import Config -from haiku.rag.embeddings import contextualize, embed_chunks -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.config import AppConfig, EmbeddingModelConfig, EmbeddingsConfig +from haiku.rag.embeddings import contextualize, embed_chunks, get_embedder from haiku.rag.store.models.chunk import Chunk OPENAI_AVAILABLE = bool(os.getenv("OPENAI_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): + """Calculate cosine similarity between embeddings and a test embedding.""" return [ np.dot(embedding, test_embedding) / (np.linalg.norm(embedding) * np.linalg.norm(test_embedding)) @@ -26,35 +22,41 @@ def similarities(embeddings, test_embedding): @pytest.mark.asyncio 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 = [ "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) + # Test batch embedding (documents) + embeddings = await embedder.embed_documents(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 query embedding 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) assert max(sims) == sims[2] 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) assert max(sims) == sims[0] 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) assert max(sims) == sims[1] @@ -62,35 +64,41 @@ async def test_ollama_embedder(): @pytest.mark.asyncio @pytest.mark.skipif(not OPENAI_AVAILABLE, reason="OpenAI API key not available") 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 = [ "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) + # Test batch embedding (documents) + embeddings = await embedder.embed_documents(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 query embedding 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) assert max(sims) == sims[2] 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) assert max(sims) == sims[0] 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) assert max(sims) == sims[1] @@ -98,38 +106,42 @@ async def test_openai_embedder(): @pytest.mark.asyncio @pytest.mark.skipif(not VOYAGEAI_AVAILABLE, reason="VoyageAI API key not available") async def test_voyageai_embedder(): + """Test VoyageAI embedder.""" try: - from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder - - embedder = VoyageAIEmbedder("voyage-3.5", 1024) + config = AppConfig( + embeddings=EmbeddingsConfig( + model=EmbeddingModelConfig( + provider="voyageai", name="voyage-3.5", vector_dim=1024 + ) + ) + ) + embedder = get_embedder(config) 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) + # Test batch embedding (documents) + embeddings = await embedder.embed_documents(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 query embedding 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) assert max(sims) == sims[2] 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) assert max(sims) == sims[0] 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) assert max(sims) == sims[1] @@ -137,44 +149,6 @@ async def test_voyageai_embedder(): 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(): """Test that contextualize prepends headings to chunk content.""" chunks = [ diff --git a/tests/test_embedder_config.py b/tests/test_embedder_config.py index a6e431d4..7dcbc6df 100644 --- a/tests/test_embedder_config.py +++ b/tests/test_embedder_config.py @@ -4,16 +4,14 @@ from haiku.rag.config import ( AppConfig, EmbeddingModelConfig, EmbeddingsConfig, - LMStudioConfig, OllamaConfig, ProvidersConfig, - VLLMConfig, ) from haiku.rag.embeddings import get_embedder -def test_embedder_uses_config_from_get_embedder(): - """Test that embedders use the config passed to 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( @@ -22,41 +20,16 @@ def test_embedder_uses_config_from_get_embedder(): ), providers=ProvidersConfig( ollama=OllamaConfig(base_url="http://custom-ollama:8080"), - vllm=VLLMConfig(embeddings_base_url="http://custom-vllm:9000"), ), ) embedder = get_embedder(custom_config) - assert embedder._model == "custom-model" 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(): - """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( embeddings=EmbeddingsConfig( model=EmbeddingModelConfig( @@ -67,48 +40,68 @@ def test_openai_embedder_uses_config(): embedder = get_embedder(custom_config) - assert embedder._model == "text-embedding-3-large" assert embedder._vector_dim == 3072 - assert embedder._config == custom_config -def test_lm_studio_embedder_uses_config(): - """Test that lm_studio embedder uses the config passed to get_embedder.""" +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="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) - assert embedder._model == "custom-lm-studio-model" assert embedder._vector_dim == 1024 - assert ( - embedder._config.providers.lm_studio.base_url == "http://custom-lmstudio:5678" - ) -@pytest.mark.skipif( - True, reason="VoyageAI is an optional dependency, may not be installed" -) -def test_voyageai_embedder_uses_config(): - """Test that voyageai embedder uses the config passed to get_embedder.""" +def test_sentence_transformers_embedder_uses_config(): + """Test that SentenceTransformers embedder uses the config.""" custom_config = AppConfig( embeddings=EmbeddingsConfig( 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) - assert embedder._model == "voyage-large-2" - assert embedder._vector_dim == 1536 - assert embedder._config == 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) diff --git a/tests/test_utils.py b/tests/test_utils.py index f81b76b4..a7edf512 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -270,20 +270,6 @@ def test_get_model_bedrock_with_thinking(): 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(): """Test get_model returns string format for unknown providers.""" model_config = ModelConfig(provider="mistral", name="mistral-large-latest")