diff --git a/src/haiku/rag/embeddings/base.py b/src/haiku/rag/embeddings/base.py index 0edf568f..8a8ea4e2 100644 --- a/src/haiku/rag/embeddings/base.py +++ b/src/haiku/rag/embeddings/base.py @@ -1,3 +1,5 @@ +from typing import overload + from haiku.rag.config import Config @@ -9,6 +11,12 @@ class EmbedderBase: self._model = model 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]]: raise NotImplementedError( "Embedder is an abstract class. Please implement the embed method in a subclass." diff --git a/src/haiku/rag/embeddings/ollama.py b/src/haiku/rag/embeddings/ollama.py index a7303ea7..ea565541 100644 --- a/src/haiku/rag/embeddings/ollama.py +++ b/src/haiku/rag/embeddings/ollama.py @@ -1,3 +1,5 @@ +from typing import overload + from openai import AsyncOpenAI from haiku.rag.config import Config @@ -5,6 +7,12 @@ 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"{Config.OLLAMA_BASE_URL}/v1", api_key="dummy") if not text: diff --git a/src/haiku/rag/embeddings/openai.py b/src/haiku/rag/embeddings/openai.py index 5b0ea2ff..7660a2e9 100644 --- a/src/haiku/rag/embeddings/openai.py +++ b/src/haiku/rag/embeddings/openai.py @@ -1,9 +1,17 @@ +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: diff --git a/src/haiku/rag/embeddings/vllm.py b/src/haiku/rag/embeddings/vllm.py index 2d2f77bd..d7e2c608 100644 --- a/src/haiku/rag/embeddings/vllm.py +++ b/src/haiku/rag/embeddings/vllm.py @@ -1,3 +1,5 @@ +from typing import overload + from openai import AsyncOpenAI from haiku.rag.config import Config @@ -5,6 +7,12 @@ 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"{Config.VLLM_EMBEDDINGS_BASE_URL}/v1", api_key="dummy" diff --git a/src/haiku/rag/embeddings/voyageai.py b/src/haiku/rag/embeddings/voyageai.py index 60c8c6b3..4d6af089 100644 --- a/src/haiku/rag/embeddings/voyageai.py +++ b/src/haiku/rag/embeddings/voyageai.py @@ -1,9 +1,17 @@ try: + from typing import overload + from voyageai.client import Client # type: ignore 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 = Client() if not text: