diff --git a/haiku_rag_slim/haiku/rag/embeddings/__init__.py b/haiku_rag_slim/haiku/rag/embeddings/__init__.py index 510cf01b..75d39f77 100644 --- a/haiku_rag_slim/haiku/rag/embeddings/__init__.py +++ b/haiku_rag_slim/haiku/rag/embeddings/__init__.py @@ -1,13 +1,40 @@ -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, overload + +from pydantic_ai.embeddings import Embedder +from pydantic_ai.embeddings.openai import OpenAIEmbeddingModel +from pydantic_ai.providers.ollama import OllamaProvider +from pydantic_ai.providers.openai import OpenAIProvider from haiku.rag.config import AppConfig, Config -from haiku.rag.embeddings.base import EmbedderBase -from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder if TYPE_CHECKING: from haiku.rag.store.models.chunk import Chunk +class EmbedderWrapper: + """Wrapper around pydantic-ai Embedder to provide simple embed() interface.""" + + def __init__(self, embedder: Embedder, vector_dim: int): + self._embedder = embedder + self._vector_dim = vector_dim + + @overload + async def embed(self, text: str) -> list[float]: ... + + @overload + async def embed(self, text: list[str]) -> list[list[float]]: ... + + async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]: + if isinstance(text, str): + result = await self._embedder.embed_query(text) + return list(result.embeddings[0]) + else: + if not text: + return [] + result = await self._embedder.embed_documents(text) + return [list(e) for e in result.embeddings] + + def contextualize(chunks: list["Chunk"]) -> list[str]: """Prepare chunk content for embedding by adding context. @@ -71,9 +98,8 @@ async def embed_chunks( ] -def get_embedder(config: AppConfig = Config) -> EmbedderBase: - """ - Factory function to get the appropriate embedder based on the configuration. +def get_embedder(config: AppConfig = Config) -> EmbedderWrapper: + """Factory function to get the appropriate embedder based on the configuration. Args: config: Configuration to use. Defaults to global Config. @@ -82,11 +108,29 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase: An embedder instance configured according to the config. """ embedding_model = config.embeddings.model + provider = embedding_model.provider + model_name = embedding_model.name + vector_dim = embedding_model.vector_dim - if embedding_model.provider == "ollama": - return OllamaEmbedder(embedding_model.name, embedding_model.vector_dim, config) + if provider == "ollama": + # Use model-level base_url if set, otherwise fall back to providers config + base_url = embedding_model.base_url or f"{config.providers.ollama.base_url}/v1" + model = OpenAIEmbeddingModel( + model_name, + provider=OllamaProvider(base_url=base_url), + ) + return EmbedderWrapper(Embedder(model), vector_dim) - if embedding_model.provider == "voyageai": + if provider == "openai": + if embedding_model.base_url: + model = OpenAIEmbeddingModel( + model_name, + provider=OpenAIProvider(base_url=embedding_model.base_url), + ) + return EmbedderWrapper(Embedder(model), vector_dim) + return EmbedderWrapper(Embedder(f"openai:{model_name}"), vector_dim) + + if provider == "voyageai": try: from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder except ImportError: @@ -95,25 +139,14 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase: "Please install haiku.rag with the 'voyageai' extra: " "uv pip install haiku.rag[voyageai]" ) - return VoyageAIEmbedder( - embedding_model.name, embedding_model.vector_dim, config + return VoyageAIEmbedder(model_name, vector_dim, config) # type: ignore[return-value] + + if provider == "cohere": + return EmbedderWrapper(Embedder(f"cohere:{model_name}"), vector_dim) + + if provider == "sentence-transformers": + return EmbedderWrapper( + Embedder(f"sentence-transformers:{model_name}"), vector_dim ) - if embedding_model.provider == "openai": - from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder - - return OpenAIEmbedder(embedding_model.name, embedding_model.vector_dim, config) - - if embedding_model.provider == "vllm": - from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder - - return VllmEmbedder(embedding_model.name, embedding_model.vector_dim, config) - - if embedding_model.provider == "lm_studio": - from haiku.rag.embeddings.lm_studio import Embedder as LMStudioEmbedder - - return LMStudioEmbedder( - embedding_model.name, embedding_model.vector_dim, config - ) - - raise ValueError(f"Unsupported embedding provider: {embedding_model.provider}") + raise ValueError(f"Unsupported embedding provider: {provider}") diff --git a/haiku_rag_slim/haiku/rag/embeddings/base.py b/haiku_rag_slim/haiku/rag/embeddings/base.py deleted file mode 100644 index 6049a840..00000000 --- a/haiku_rag_slim/haiku/rag/embeddings/base.py +++ /dev/null @@ -1,25 +0,0 @@ -from typing import overload - -from haiku.rag.config import AppConfig, Config - - -class EmbedderBase: - _model: str = Config.embeddings.model.name - _vector_dim: int = Config.embeddings.model.vector_dim - _config: AppConfig = Config - - 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]: ... - - @overload - async def embed(self, text: list[str]) -> list[list[float]]: ... - - async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]: - raise NotImplementedError( - "Embedder is an abstract class. Please implement the embed method in a subclass." - ) diff --git a/haiku_rag_slim/haiku/rag/embeddings/lm_studio.py b/haiku_rag_slim/haiku/rag/embeddings/lm_studio.py deleted file mode 100644 index 86bf825f..00000000 --- a/haiku_rag_slim/haiku/rag/embeddings/lm_studio.py +++ /dev/null @@ -1,28 +0,0 @@ -from typing import overload - -from openai import AsyncOpenAI - -from haiku.rag.embeddings.base import EmbedderBase - - -class Embedder(EmbedderBase): # pragma: no cover - @overload - async def embed(self, text: str) -> list[float]: ... - - @overload - async def embed(self, text: list[str]) -> list[list[float]]: ... - - async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]: - client = AsyncOpenAI( - base_url=f"{self._config.providers.lm_studio.base_url}/v1", api_key="dummy" - ) - if not text: - return [] - response = await client.embeddings.create( - model=self._model, - input=text, - ) - if isinstance(text, str): - return response.data[0].embedding - else: - return [item.embedding for item in response.data] diff --git a/haiku_rag_slim/haiku/rag/embeddings/ollama.py b/haiku_rag_slim/haiku/rag/embeddings/ollama.py deleted file mode 100644 index 9fccccc9..00000000 --- a/haiku_rag_slim/haiku/rag/embeddings/ollama.py +++ /dev/null @@ -1,28 +0,0 @@ -from typing import overload - -from openai import AsyncOpenAI - -from haiku.rag.embeddings.base import EmbedderBase - - -class Embedder(EmbedderBase): - @overload - async def embed(self, text: str) -> list[float]: ... - - @overload - async def embed(self, text: list[str]) -> list[list[float]]: ... - - async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]: - client = AsyncOpenAI( - base_url=f"{self._config.providers.ollama.base_url}/v1", api_key="dummy" - ) - if not text: - return [] - response = await client.embeddings.create( - model=self._model, - input=text, - ) - if isinstance(text, str): - return response.data[0].embedding - else: - return [item.embedding for item in response.data] diff --git a/haiku_rag_slim/haiku/rag/embeddings/openai.py b/haiku_rag_slim/haiku/rag/embeddings/openai.py deleted file mode 100644 index 7660a2e9..00000000 --- a/haiku_rag_slim/haiku/rag/embeddings/openai.py +++ /dev/null @@ -1,26 +0,0 @@ -from typing import overload - -from openai import AsyncOpenAI - -from haiku.rag.embeddings.base import EmbedderBase - - -class Embedder(EmbedderBase): - @overload - async def embed(self, text: str) -> list[float]: ... - - @overload - async def embed(self, text: list[str]) -> list[list[float]]: ... - - async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]: - client = AsyncOpenAI() - if not text: - return [] - response = await client.embeddings.create( - model=self._model, - input=text, - ) - if isinstance(text, str): - return response.data[0].embedding - else: - return [item.embedding for item in response.data] diff --git a/haiku_rag_slim/haiku/rag/embeddings/vllm.py b/haiku_rag_slim/haiku/rag/embeddings/vllm.py deleted file mode 100644 index 7d521fa2..00000000 --- a/haiku_rag_slim/haiku/rag/embeddings/vllm.py +++ /dev/null @@ -1,29 +0,0 @@ -from typing import overload - -from openai import AsyncOpenAI - -from haiku.rag.embeddings.base import EmbedderBase - - -class Embedder(EmbedderBase): # pragma: no cover - @overload - async def embed(self, text: str) -> list[float]: ... - - @overload - async def embed(self, text: list[str]) -> list[list[float]]: ... - - async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]: - client = AsyncOpenAI( - base_url=f"{self._config.providers.vllm.embeddings_base_url}/v1", - api_key="dummy", - ) - if not text: - return [] - response = await client.embeddings.create( - model=self._model, - input=text, - ) - if isinstance(text, str): - return response.data[0].embedding - else: - return [item.embedding for item in response.data]