Pass config to embedders

This commit is contained in:
Yiorgis Gozadinos 2025-10-28 09:15:08 +02:00
parent 54fa3de2da
commit e11a5128f9
No known key found for this signature in database
5 changed files with 109 additions and 10 deletions

View file

@ -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}")

View file

@ -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]: ...

View file

@ -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 []

View file

@ -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 []

View file

@ -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