From 51723308ed504107ea9d120e02b533530000ad62 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Wed, 18 Jun 2025 08:26:11 +0200 Subject: [PATCH] Introduce embedder factory --- src/haiku/rag/config.py | 2 ++ src/haiku/rag/embeddings/__init__.py | 13 +++++++++++++ src/haiku/rag/embeddings/base.py | 4 ++++ src/haiku/rag/store/engine.py | 4 ++-- src/haiku/rag/store/repositories/chunk.py | 6 +++--- tests/test_embedder.py | 6 +++--- 6 files changed, 27 insertions(+), 8 deletions(-) diff --git a/src/haiku/rag/config.py b/src/haiku/rag/config.py index 53f0939a..de22f209 100644 --- a/src/haiku/rag/config.py +++ b/src/haiku/rag/config.py @@ -11,7 +11,9 @@ class AppConfig(BaseModel): OLLAMA_BASE_URL: str = "http://localhost:11434" + EMBEDDING_PROVIDER: str = "ollama" EMBEDDING_MODEL: str = "mxbai-embed-large" + EMBEDDING_VECTOR_DIM: int = 1024 CHUNK_SIZE: int = 256 CHUNK_OVERLAP: int = 32 diff --git a/src/haiku/rag/embeddings/__init__.py b/src/haiku/rag/embeddings/__init__.py index e69de29b..bcaee4d6 100644 --- a/src/haiku/rag/embeddings/__init__.py +++ b/src/haiku/rag/embeddings/__init__.py @@ -0,0 +1,13 @@ +from haiku.rag.config import Config +from haiku.rag.embeddings.base import EmbedderBase +from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder + + +def get_embedder() -> EmbedderBase: + """ + Factory function to get the appropriate embedder based on the configuration. + """ + + if Config.EMBEDDING_PROVIDER == "ollama": + return OllamaEmbedder(Config.EMBEDDING_MODEL, Config.EMBEDDING_VECTOR_DIM) + raise ValueError(f"Unsupported embedding provider: {Config.EMBEDDING_PROVIDER}") diff --git a/src/haiku/rag/embeddings/base.py b/src/haiku/rag/embeddings/base.py index a62deffd..16e19d9c 100644 --- a/src/haiku/rag/embeddings/base.py +++ b/src/haiku/rag/embeddings/base.py @@ -2,6 +2,10 @@ class EmbedderBase: _model: str = "" _vector_dim: int = 0 + def __init__(self, model: str, vector_dim: int): + self._model = model + self._vector_dim = vector_dim + async def embed(self, text: str) -> list[float]: raise NotImplementedError( "Embedder is an abstract class. Please implement the embed method in a subclass." diff --git a/src/haiku/rag/store/engine.py b/src/haiku/rag/store/engine.py index 449b702b..cdf1d2ca 100644 --- a/src/haiku/rag/store/engine.py +++ b/src/haiku/rag/store/engine.py @@ -5,7 +5,7 @@ from typing import Literal import sqlite_vec -from haiku.rag.embeddings.ollama import Embedder +from haiku.rag.embeddings import get_embedder class Store: @@ -43,7 +43,7 @@ class Store: """) # Create vector table for chunk embeddings - embedder = Embedder() + embedder = get_embedder() db.execute(f""" CREATE VIRTUAL TABLE IF NOT EXISTS chunk_embeddings USING vec0( chunk_id INTEGER PRIMARY KEY, diff --git a/src/haiku/rag/store/repositories/chunk.py b/src/haiku/rag/store/repositories/chunk.py index 70026238..4366e065 100644 --- a/src/haiku/rag/store/repositories/chunk.py +++ b/src/haiku/rag/store/repositories/chunk.py @@ -2,7 +2,7 @@ import json import re from haiku.rag.chunker import chunker -from haiku.rag.embeddings.ollama import Embedder +from haiku.rag.embeddings import get_embedder from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.repositories.base import BaseRepository @@ -10,9 +10,9 @@ from haiku.rag.store.repositories.base import BaseRepository class ChunkRepository(BaseRepository[Chunk]): """Repository for Chunk database operations.""" - def __init__(self, store, embedder: Embedder | None = None): + def __init__(self, store): super().__init__(store) - self.embedder = embedder or Embedder() + self.embedder = get_embedder() async def create(self, entity: Chunk, commit: bool = True) -> Chunk: """Create a chunk in the database.""" diff --git a/tests/test_embedder.py b/tests/test_embedder.py index 509b0c7e..5b9b3d35 100644 --- a/tests/test_embedder.py +++ b/tests/test_embedder.py @@ -1,19 +1,19 @@ import numpy as np import pytest -from haiku.rag.embeddings.ollama import Embedder +from haiku.rag.embeddings import get_embedder @pytest.mark.asyncio async def test_embedder(): - embedder = Embedder() + embedder = get_embedder() embedding = await embedder.embed("hello world") assert len(embedding) == embedder._vector_dim @pytest.mark.asyncio async def test_similarity(): - embedder = Embedder() + embedder = get_embedder() phrases = [ "I enjoy eating great food.", "Python is my favorite programming language.",