diff --git a/src/haiku/rag/embeddings/__init__.py b/src/haiku/rag/embeddings/__init__.py index 17e85a3e..c8e00fce 100644 --- a/src/haiku/rag/embeddings/__init__.py +++ b/src/haiku/rag/embeddings/__init__.py @@ -15,7 +15,9 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase: """ if config.embeddings.provider == "ollama": - return OllamaEmbedder(config.embeddings.model, config.embeddings.vector_dim) + return OllamaEmbedder( + config.embeddings.model, config.embeddings.vector_dim, config + ) if config.embeddings.provider == "voyageai": try: @@ -26,16 +28,22 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase: "Please install haiku.rag with the 'voyageai' extra: " "uv pip install haiku.rag[voyageai]" ) - return VoyageAIEmbedder(config.embeddings.model, config.embeddings.vector_dim) + return VoyageAIEmbedder( + config.embeddings.model, config.embeddings.vector_dim, config + ) if config.embeddings.provider == "openai": from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder - return OpenAIEmbedder(config.embeddings.model, config.embeddings.vector_dim) + return OpenAIEmbedder( + config.embeddings.model, config.embeddings.vector_dim, config + ) if config.embeddings.provider == "vllm": from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder - return VllmEmbedder(config.embeddings.model, config.embeddings.vector_dim) + return VllmEmbedder( + config.embeddings.model, config.embeddings.vector_dim, config + ) raise ValueError(f"Unsupported embedding provider: {config.embeddings.provider}") diff --git a/src/haiku/rag/embeddings/base.py b/src/haiku/rag/embeddings/base.py index fe04e811..bcd80f91 100644 --- a/src/haiku/rag/embeddings/base.py +++ b/src/haiku/rag/embeddings/base.py @@ -1,15 +1,17 @@ from typing import overload -from haiku.rag.config import Config +from haiku.rag.config import AppConfig, Config class EmbedderBase: _model: str = Config.embeddings.model _vector_dim: int = Config.embeddings.vector_dim + _config: AppConfig = Config - def __init__(self, model: str, vector_dim: int): + def __init__(self, model: str, vector_dim: int, config: AppConfig = Config): self._model = model self._vector_dim = vector_dim + self._config = config @overload async def embed(self, text: str) -> list[float]: ... diff --git a/src/haiku/rag/embeddings/ollama.py b/src/haiku/rag/embeddings/ollama.py index e9fbe8a5..9fccccc9 100644 --- a/src/haiku/rag/embeddings/ollama.py +++ b/src/haiku/rag/embeddings/ollama.py @@ -2,7 +2,6 @@ from typing import overload from openai import AsyncOpenAI -from haiku.rag.config import Config from haiku.rag.embeddings.base import EmbedderBase @@ -15,7 +14,7 @@ class Embedder(EmbedderBase): async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]: client = AsyncOpenAI( - base_url=f"{Config.providers.ollama.base_url}/v1", api_key="dummy" + base_url=f"{self._config.providers.ollama.base_url}/v1", api_key="dummy" ) if not text: return [] diff --git a/src/haiku/rag/embeddings/vllm.py b/src/haiku/rag/embeddings/vllm.py index 4a55bf7d..fb35e43e 100644 --- a/src/haiku/rag/embeddings/vllm.py +++ b/src/haiku/rag/embeddings/vllm.py @@ -2,7 +2,6 @@ from typing import overload from openai import AsyncOpenAI -from haiku.rag.config import Config from haiku.rag.embeddings.base import EmbedderBase @@ -15,7 +14,8 @@ class Embedder(EmbedderBase): async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]: client = AsyncOpenAI( - base_url=f"{Config.providers.vllm.embeddings_base_url}/v1", api_key="dummy" + base_url=f"{self._config.providers.vllm.embeddings_base_url}/v1", + api_key="dummy", ) if not text: return [] diff --git a/tests/test_embedder_config.py b/tests/test_embedder_config.py new file mode 100644 index 00000000..a386d4b3 --- /dev/null +++ b/tests/test_embedder_config.py @@ -0,0 +1,90 @@ +import pytest + +from haiku.rag.config import ( + AppConfig, + EmbeddingsConfig, + 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.""" + custom_config = AppConfig( + embeddings=EmbeddingsConfig( + provider="ollama", + model="custom-model", + vector_dim=512, + ), + 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( + provider="vllm", + model="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.""" + custom_config = AppConfig( + embeddings=EmbeddingsConfig( + provider="openai", + model="text-embedding-3-large", + vector_dim=3072, + ), + ) + + embedder = get_embedder(custom_config) + + assert embedder._model == "text-embedding-3-large" + assert embedder._vector_dim == 3072 + assert embedder._config == custom_config + + +@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.""" + custom_config = AppConfig( + embeddings=EmbeddingsConfig( + provider="voyageai", + model="voyage-large-2", + vector_dim=1536, + ), + ) + + embedder = get_embedder(custom_config) + + assert embedder._model == "voyage-large-2" + assert embedder._vector_dim == 1536 + assert embedder._config == custom_config