diff --git a/src/haiku/rag/embeddings/__init__.py b/src/haiku/rag/embeddings/__init__.py index bb371bd1..c463e2ee 100644 --- a/src/haiku/rag/embeddings/__init__.py +++ b/src/haiku/rag/embeddings/__init__.py @@ -17,20 +17,14 @@ def get_embedder() -> EmbedderBase: except ImportError: raise ImportError( "VoyageAI embedder requires the 'voyageai' package. " - "Please install haiku.rag with the 'voyageai' extra:" + "Please install haiku.rag with the 'voyageai' extra: " "uv pip install haiku.rag[voyageai]" ) return VoyageAIEmbedder(Config.EMBEDDINGS_MODEL, Config.EMBEDDINGS_VECTOR_DIM) if Config.EMBEDDINGS_PROVIDER == "openai": - try: - from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder - except ImportError: - raise ImportError( - "OpenAI embedder requires the 'openai' package. " - "Please install haiku.rag with the 'openai' extra:" - "uv pip install haiku.rag[openai]" - ) + from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder + return OpenAIEmbedder(Config.EMBEDDINGS_MODEL, Config.EMBEDDINGS_VECTOR_DIM) raise ValueError(f"Unsupported embedding provider: {Config.EMBEDDINGS_PROVIDER}") diff --git a/src/haiku/rag/embeddings/openai.py b/src/haiku/rag/embeddings/openai.py index 818f0e5b..485c97fe 100644 --- a/src/haiku/rag/embeddings/openai.py +++ b/src/haiku/rag/embeddings/openai.py @@ -1,16 +1,13 @@ -try: - from openai import AsyncOpenAI +from openai import AsyncOpenAI - from haiku.rag.embeddings.base import EmbedderBase +from haiku.rag.embeddings.base import EmbedderBase - class Embedder(EmbedderBase): - async def embed(self, text: str) -> list[float]: - client = AsyncOpenAI() - response = await client.embeddings.create( - model=self._model, - input=text, - ) - return response.data[0].embedding -except ImportError: - pass +class Embedder(EmbedderBase): + async def embed(self, text: str) -> list[float]: + client = AsyncOpenAI() + response = await client.embeddings.create( + model=self._model, + input=text, + ) + return response.data[0].embedding diff --git a/tests/test_embedder.py b/tests/test_embedder.py index f6421631..7d227b1c 100644 --- a/tests/test_embedder.py +++ b/tests/test_embedder.py @@ -1,19 +1,26 @@ import numpy as np import pytest -from haiku.rag.embeddings import get_embedder +from haiku.rag.config import Config +from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder +from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder + +OPENAI_AVAILABLE = bool(Config.OPENAI_API_KEY) +VOYAGEAI_AVAILABLE = bool(Config.VOYAGE_API_KEY) + + +# Calculate cosine similarity +def similarities(embeddings, test_embedding): + return [ + np.dot(embedding, test_embedding) + / (np.linalg.norm(embedding) * np.linalg.norm(test_embedding)) + for embedding in embeddings + ] @pytest.mark.asyncio -async def test_embedder(): - embedder = get_embedder() - embedding = await embedder.embed("hello world") - assert len(embedding) == embedder._vector_dim - - -@pytest.mark.asyncio -async def test_similarity(): - embedder = get_embedder() +async def test_ollama_embedder(): + embedder = OllamaEmbedder("mxbai-embed-large", 1024) phrases = [ "I enjoy eating great food.", "Python is my favorite programming language.", @@ -21,14 +28,6 @@ async def test_similarity(): ] embeddings = [np.array(await embedder.embed(phrase)) for phrase in phrases] - # Calculate cosine similarity - def similarities(embeddings, test_embedding): - return [ - np.dot(embedding, test_embedding) - / (np.linalg.norm(embedding) * np.linalg.norm(test_embedding)) - for embedding in embeddings - ] - test_phrase = "I am going for a camping trip." test_embedding = await embedder.embed(test_phrase) @@ -49,80 +48,66 @@ async def test_similarity(): @pytest.mark.asyncio -async def test_openai_embedder(monkeypatch): - monkeypatch.setenv("EMBEDDINGS_PROVIDER", "openai") - monkeypatch.setenv("EMBEDDINGS_MODEL", "text-embedding-3-small") +@pytest.mark.skipif(not OPENAI_AVAILABLE, reason="OpenAI API key not available") +async def test_openai_embedder(): + embedder = OpenAIEmbedder("text-embedding-3-small", 1536) + phrases = [ + "I enjoy eating great food.", + "Python is my favorite programming language.", + "I love to travel and see new places.", + ] + embeddings = [np.array(await embedder.embed(phrase)) for phrase in phrases] - try: - from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder + test_phrase = "I am going for a camping trip." + test_embedding = await embedder.embed(test_phrase) - embedder = OpenAIEmbedder("text-embedding-3-small", 1536) + sims = similarities(embeddings, test_embedding) + assert max(sims) == sims[2] - # Mock the OpenAI client - class MockEmbeddingData: - def __init__(self, embedding): - self.embedding = embedding + test_phrase = "When is dinner ready?" + test_embedding = await embedder.embed(test_phrase) - class MockResponse: - def __init__(self, embedding): - self.data = [MockEmbeddingData(embedding)] + sims = similarities(embeddings, test_embedding) + assert max(sims) == sims[0] - class MockAsyncOpenAI: - class MockEmbeddings: - async def create(self, model, input): - return MockResponse([0.1] * 1536) + test_phrase = "I work as a software developer." + test_embedding = await embedder.embed(test_phrase) - def __init__(self): - self.embeddings = self.MockEmbeddings() - - # Patch the AsyncOpenAI import - import haiku.rag.embeddings.openai - - original_client = haiku.rag.embeddings.openai.AsyncOpenAI - haiku.rag.embeddings.openai.AsyncOpenAI = MockAsyncOpenAI - - try: - embedding = await embedder.embed("test text") - assert len(embedding) == 1536 - assert all(isinstance(x, float) for x in embedding) - finally: - haiku.rag.embeddings.openai.AsyncOpenAI = original_client - - except ImportError: - pytest.skip("OpenAI package not installed") + sims = similarities(embeddings, test_embedding) + assert max(sims) == sims[1] @pytest.mark.asyncio -async def test_voyageai_embedder(monkeypatch): - monkeypatch.setenv("EMBEDDINGS_PROVIDER", "voyageai") - monkeypatch.setenv("EMBEDDINGS_MODEL", "voyage-3.5") - +@pytest.mark.skipif(not VOYAGEAI_AVAILABLE, reason="VoyageAI API key not available") +async def test_voyageai_embedder(): try: from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder embedder = VoyageAIEmbedder("voyage-3.5", 1024) + phrases = [ + "I enjoy eating great food.", + "Python is my favorite programming language.", + "I love to travel and see new places.", + ] + embeddings = [np.array(await embedder.embed(phrase)) for phrase in phrases] - # Mock the VoyageAI client - class MockEmbeddings: - def __init__(self, embeddings): - self.embeddings = embeddings + test_phrase = "I am going for a camping trip." + test_embedding = await embedder.embed(test_phrase) - class MockClient: - def embed(self, texts, model, output_dtype): - return MockEmbeddings([[0.1] * 1024]) + sims = similarities(embeddings, test_embedding) + assert max(sims) == sims[2] - # Patch the Client import - import haiku.rag.embeddings.voyageai + test_phrase = "When is dinner ready?" + test_embedding = await embedder.embed(test_phrase) - original_client = haiku.rag.embeddings.voyageai.Client - haiku.rag.embeddings.voyageai.Client = MockClient + sims = similarities(embeddings, test_embedding) + assert max(sims) == sims[0] - try: - embedding = await embedder.embed("test text") - assert len(embedding) == 1024 - assert all(isinstance(x, float) for x in embedding) - finally: - haiku.rag.embeddings.voyageai.Client = original_client + 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] except ImportError: pytest.skip("VoyageAI package not installed") diff --git a/uv.lock b/uv.lock index 8c270ec3..b267b317 100644 --- a/uv.lock +++ b/uv.lock @@ -1027,18 +1027,9 @@ dependencies = [ ] [package.optional-dependencies] -anthropic = [ - { name = "anthropic" }, -] -cohere = [ - { name = "cohere" }, -] mxbai = [ { name = "mxbai-rerank" }, ] -openai = [ - { name = "openai" }, -] voyageai = [ { name = "voyageai" }, ] @@ -1058,14 +1049,11 @@ dev = [ [package.metadata] requires-dist = [ - { name = "anthropic", marker = "extra == 'anthropic'", specifier = ">=0.56.0" }, - { name = "cohere", marker = "extra == 'cohere'", specifier = ">=5.16.1" }, { name = "docling", specifier = ">=2.15.0" }, { name = "fastmcp", specifier = ">=2.8.1" }, { name = "httpx", specifier = ">=0.28.1" }, { name = "mxbai-rerank", marker = "extra == 'mxbai'", specifier = ">=0.1.6" }, { name = "ollama", specifier = ">=0.5.3" }, - { name = "openai", marker = "extra == 'openai'", specifier = ">=1.0.0" }, { name = "pydantic", specifier = ">=2.11.7" }, { name = "pydantic-ai", specifier = ">=0.7.2" }, { name = "python-dotenv", specifier = ">=1.1.0" }, @@ -1076,7 +1064,7 @@ requires-dist = [ { name = "voyageai", marker = "extra == 'voyageai'", specifier = ">=0.3.2" }, { name = "watchfiles", specifier = ">=1.1.0" }, ] -provides-extras = ["voyageai", "openai", "anthropic", "cohere", "mxbai"] +provides-extras = ["voyageai", "mxbai"] [package.metadata.requires-dev] dev = [