Do not use mocks in embeddings test

This commit is contained in:
Yiorgis Gozadinos 2025-08-15 22:01:57 +02:00
parent 4e2e6e490e
commit 79841f32ca
No known key found for this signature in database
4 changed files with 73 additions and 109 deletions

View file

@ -17,20 +17,14 @@ def get_embedder() -> EmbedderBase:
except ImportError: except ImportError:
raise ImportError( raise ImportError(
"VoyageAI embedder requires the 'voyageai' package. " "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]" "uv pip install haiku.rag[voyageai]"
) )
return VoyageAIEmbedder(Config.EMBEDDINGS_MODEL, Config.EMBEDDINGS_VECTOR_DIM) return VoyageAIEmbedder(Config.EMBEDDINGS_MODEL, Config.EMBEDDINGS_VECTOR_DIM)
if Config.EMBEDDINGS_PROVIDER == "openai": if Config.EMBEDDINGS_PROVIDER == "openai":
try: from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
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]"
)
return OpenAIEmbedder(Config.EMBEDDINGS_MODEL, Config.EMBEDDINGS_VECTOR_DIM) return OpenAIEmbedder(Config.EMBEDDINGS_MODEL, Config.EMBEDDINGS_VECTOR_DIM)
raise ValueError(f"Unsupported embedding provider: {Config.EMBEDDINGS_PROVIDER}") raise ValueError(f"Unsupported embedding provider: {Config.EMBEDDINGS_PROVIDER}")

View file

@ -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: class Embedder(EmbedderBase):
pass 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

View file

@ -1,19 +1,26 @@
import numpy as np import numpy as np
import pytest 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 @pytest.mark.asyncio
async def test_embedder(): async def test_ollama_embedder():
embedder = get_embedder() embedder = OllamaEmbedder("mxbai-embed-large", 1024)
embedding = await embedder.embed("hello world")
assert len(embedding) == embedder._vector_dim
@pytest.mark.asyncio
async def test_similarity():
embedder = get_embedder()
phrases = [ phrases = [
"I enjoy eating great food.", "I enjoy eating great food.",
"Python is my favorite programming language.", "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] 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_phrase = "I am going for a camping trip."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed(test_phrase)
@ -49,80 +48,66 @@ async def test_similarity():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_openai_embedder(monkeypatch): @pytest.mark.skipif(not OPENAI_AVAILABLE, reason="OpenAI API key not available")
monkeypatch.setenv("EMBEDDINGS_PROVIDER", "openai") async def test_openai_embedder():
monkeypatch.setenv("EMBEDDINGS_MODEL", "text-embedding-3-small") 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: test_phrase = "I am going for a camping trip."
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder 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 test_phrase = "When is dinner ready?"
class MockEmbeddingData: test_embedding = await embedder.embed(test_phrase)
def __init__(self, embedding):
self.embedding = embedding
class MockResponse: sims = similarities(embeddings, test_embedding)
def __init__(self, embedding): assert max(sims) == sims[0]
self.data = [MockEmbeddingData(embedding)]
class MockAsyncOpenAI: test_phrase = "I work as a software developer."
class MockEmbeddings: test_embedding = await embedder.embed(test_phrase)
async def create(self, model, input):
return MockResponse([0.1] * 1536)
def __init__(self): sims = similarities(embeddings, test_embedding)
self.embeddings = self.MockEmbeddings() assert max(sims) == sims[1]
# 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")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_voyageai_embedder(monkeypatch): @pytest.mark.skipif(not VOYAGEAI_AVAILABLE, reason="VoyageAI API key not available")
monkeypatch.setenv("EMBEDDINGS_PROVIDER", "voyageai") async def test_voyageai_embedder():
monkeypatch.setenv("EMBEDDINGS_MODEL", "voyage-3.5")
try: try:
from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder
embedder = VoyageAIEmbedder("voyage-3.5", 1024) 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 test_phrase = "I am going for a camping trip."
class MockEmbeddings: test_embedding = await embedder.embed(test_phrase)
def __init__(self, embeddings):
self.embeddings = embeddings
class MockClient: sims = similarities(embeddings, test_embedding)
def embed(self, texts, model, output_dtype): assert max(sims) == sims[2]
return MockEmbeddings([[0.1] * 1024])
# Patch the Client import test_phrase = "When is dinner ready?"
import haiku.rag.embeddings.voyageai test_embedding = await embedder.embed(test_phrase)
original_client = haiku.rag.embeddings.voyageai.Client sims = similarities(embeddings, test_embedding)
haiku.rag.embeddings.voyageai.Client = MockClient assert max(sims) == sims[0]
try: test_phrase = "I work as a software developer."
embedding = await embedder.embed("test text") test_embedding = await embedder.embed(test_phrase)
assert len(embedding) == 1024
assert all(isinstance(x, float) for x in embedding) sims = similarities(embeddings, test_embedding)
finally: assert max(sims) == sims[1]
haiku.rag.embeddings.voyageai.Client = original_client
except ImportError: except ImportError:
pytest.skip("VoyageAI package not installed") pytest.skip("VoyageAI package not installed")

14
uv.lock
View file

@ -1027,18 +1027,9 @@ dependencies = [
] ]
[package.optional-dependencies] [package.optional-dependencies]
anthropic = [
{ name = "anthropic" },
]
cohere = [
{ name = "cohere" },
]
mxbai = [ mxbai = [
{ name = "mxbai-rerank" }, { name = "mxbai-rerank" },
] ]
openai = [
{ name = "openai" },
]
voyageai = [ voyageai = [
{ name = "voyageai" }, { name = "voyageai" },
] ]
@ -1058,14 +1049,11 @@ dev = [
[package.metadata] [package.metadata]
requires-dist = [ 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 = "docling", specifier = ">=2.15.0" },
{ name = "fastmcp", specifier = ">=2.8.1" }, { name = "fastmcp", specifier = ">=2.8.1" },
{ name = "httpx", specifier = ">=0.28.1" }, { name = "httpx", specifier = ">=0.28.1" },
{ name = "mxbai-rerank", marker = "extra == 'mxbai'", specifier = ">=0.1.6" }, { name = "mxbai-rerank", marker = "extra == 'mxbai'", specifier = ">=0.1.6" },
{ name = "ollama", specifier = ">=0.5.3" }, { name = "ollama", specifier = ">=0.5.3" },
{ name = "openai", marker = "extra == 'openai'", specifier = ">=1.0.0" },
{ name = "pydantic", specifier = ">=2.11.7" }, { name = "pydantic", specifier = ">=2.11.7" },
{ name = "pydantic-ai", specifier = ">=0.7.2" }, { name = "pydantic-ai", specifier = ">=0.7.2" },
{ name = "python-dotenv", specifier = ">=1.1.0" }, { name = "python-dotenv", specifier = ">=1.1.0" },
@ -1076,7 +1064,7 @@ requires-dist = [
{ name = "voyageai", marker = "extra == 'voyageai'", specifier = ">=0.3.2" }, { name = "voyageai", marker = "extra == 'voyageai'", specifier = ">=0.3.2" },
{ name = "watchfiles", specifier = ">=1.1.0" }, { name = "watchfiles", specifier = ">=1.1.0" },
] ]
provides-extras = ["voyageai", "openai", "anthropic", "cohere", "mxbai"] provides-extras = ["voyageai", "mxbai"]
[package.metadata.requires-dev] [package.metadata.requires-dev]
dev = [ dev = [