From 86e31b8d2952a55748c4afd5a869f6cdb4b1636e Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Fri, 26 Dec 2025 11:50:47 +0200 Subject: [PATCH] VoyageAI embeddings, should end up as a PR for pydantic-ai --- .../haiku/rag/embeddings/__init__.py | 5 +- .../haiku/rag/embeddings/voyageai.py | 199 +++++++++++++++--- 2 files changed, 177 insertions(+), 27 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/embeddings/__init__.py b/haiku_rag_slim/haiku/rag/embeddings/__init__.py index 8c031aba..7894ea69 100644 --- a/haiku_rag_slim/haiku/rag/embeddings/__init__.py +++ b/haiku_rag_slim/haiku/rag/embeddings/__init__.py @@ -128,14 +128,15 @@ def get_embedder(config: AppConfig = Config) -> EmbedderWrapper: if provider == "voyageai": try: - from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder + from haiku.rag.embeddings.voyageai import VoyageAIEmbeddingModel except ImportError: raise ImportError( "VoyageAI embedder requires the 'voyageai' package. " "Please install haiku.rag with the 'voyageai' extra: " "uv pip install haiku.rag[voyageai]" ) - return VoyageAIEmbedder(model_name, vector_dim, config) # type: ignore[return-value] + model = VoyageAIEmbeddingModel(model_name) + return EmbedderWrapper(Embedder(model), vector_dim) if provider == "cohere": return EmbedderWrapper(Embedder(f"cohere:{model_name}"), vector_dim) diff --git a/haiku_rag_slim/haiku/rag/embeddings/voyageai.py b/haiku_rag_slim/haiku/rag/embeddings/voyageai.py index ff6f16b0..c38b2a23 100644 --- a/haiku_rag_slim/haiku/rag/embeddings/voyageai.py +++ b/haiku_rag_slim/haiku/rag/embeddings/voyageai.py @@ -1,33 +1,182 @@ +from collections.abc import Sequence +from dataclasses import dataclass, field +from typing import Literal, cast + +from pydantic_ai.embeddings.base import EmbeddingModel +from pydantic_ai.embeddings.result import EmbeddingResult, EmbedInputType +from pydantic_ai.embeddings.settings import EmbeddingSettings +from pydantic_ai.exceptions import ModelAPIError +from pydantic_ai.usage import RequestUsage + try: - from voyageai.client import Client # type: ignore + from voyageai.client_async import AsyncClient + from voyageai.error import VoyageError +except ImportError as _import_error: + raise ImportError( + "Please install `voyageai` to use the VoyageAI embeddings model, " + "you can use — `pip install voyageai`" + ) from _import_error - from haiku.rag.config import AppConfig +LatestVoyageAIEmbeddingModelNames = Literal[ + "voyage-3-large", + "voyage-3.5", + "voyage-3.5-lite", + "voyage-code-3", + "voyage-finance-2", + "voyage-law-2", + "voyage-code-2", +] +"""Latest VoyageAI embedding models. - class Embedder: - """VoyageAI embedder with explicit query/document methods.""" +See [VoyageAI Embeddings](https://docs.voyageai.com/docs/embeddings) +for available models and their capabilities. +""" - def __init__(self, model: str, vector_dim: int, config: AppConfig): - self._model = model - self._vector_dim = vector_dim - self._config = config +VoyageAIEmbeddingModelName = str | LatestVoyageAIEmbeddingModelNames +"""Possible VoyageAI embedding model names.""" - async def embed_query(self, text: str) -> list[float]: - """Embed a search query.""" - client = Client() - res = client.embed( - [text], model=self._model, input_type="query", output_dtype="float" + +class VoyageAIEmbeddingSettings(EmbeddingSettings, total=False): + """Settings used for a VoyageAI embedding model request. + + All fields from [`EmbeddingSettings`][pydantic_ai.embeddings.EmbeddingSettings] are supported, + plus VoyageAI-specific settings prefixed with `voyageai_`. + """ + + # ALL FIELDS MUST BE `voyageai_` PREFIXED SO YOU CAN MERGE THEM WITH OTHER MODELS. + + voyageai_truncation: bool + """Whether to truncate inputs that exceed the model's context length. + + Defaults to True. If False, an error is raised for inputs that are too long. + """ + + voyageai_output_dtype: Literal["float", "int8", "uint8", "binary", "ubinary"] + """The output data type for embeddings. + + - `'float'` (default): 32-bit floats + - `'int8'`: Signed 8-bit integers (quantized) + - `'uint8'`: Unsigned 8-bit integers (quantized) + - `'binary'`: Binary embeddings + - `'ubinary'`: Unsigned binary embeddings + """ + + +_MAX_INPUT_TOKENS: dict[VoyageAIEmbeddingModelName, int] = { + "voyage-3-large": 32000, + "voyage-3.5": 32000, + "voyage-3.5-lite": 32000, + "voyage-code-3": 32000, + "voyage-finance-2": 32000, + "voyage-law-2": 16000, + "voyage-code-2": 16000, +} + + +@dataclass(init=False) +class VoyageAIEmbeddingModel(EmbeddingModel): + """VoyageAI embedding model implementation. + + VoyageAI provides state-of-the-art embedding models optimized for + retrieval, with specialized models for code, finance, and legal domains. + + Example: + ```python + from pydantic_ai.embeddings.voyageai import VoyageAIEmbeddingModel + + model = VoyageAIEmbeddingModel('voyage-3.5') + ``` + """ + + _model_name: VoyageAIEmbeddingModelName = field(repr=False) + _client: AsyncClient = field(repr=False) + + def __init__( + self, + model_name: VoyageAIEmbeddingModelName, + *, + api_key: str | None = None, + max_retries: int = 0, + timeout: int | None = None, + settings: EmbeddingSettings | None = None, + ): + """Initialize a VoyageAI embedding model. + + Args: + model_name: The name of the VoyageAI model to use. + See [VoyageAI models](https://docs.voyageai.com/docs/embeddings) + for available options. + api_key: The VoyageAI API key. If not provided, uses the + `VOYAGE_API_KEY` environment variable. + max_retries: Maximum number of retries for failed requests. + timeout: Request timeout in seconds. + settings: Model-specific [`EmbeddingSettings`][pydantic_ai.embeddings.EmbeddingSettings] + to use as defaults for this model. + """ + self._model_name = model_name + self._client = AsyncClient( + api_key=api_key, + max_retries=max_retries, + timeout=timeout, + ) + + super().__init__(settings=settings) + + @property + def model_name(self) -> VoyageAIEmbeddingModelName: + """The embedding model name.""" + return self._model_name + + @property + def system(self) -> str: + """The embedding model provider.""" + return "voyageai" + + async def embed( + self, + inputs: str | Sequence[str], + *, + input_type: EmbedInputType, + settings: EmbeddingSettings | None = None, + ) -> EmbeddingResult: + inputs, settings = self.prepare_embed(inputs, settings) + settings = cast(VoyageAIEmbeddingSettings, settings) + + voyageai_input_type = "document" if input_type == "document" else "query" + + try: + response = await self._client.embed( + texts=list(inputs), + model=self.model_name, + input_type=voyageai_input_type, + truncation=settings.get("voyageai_truncation", True), + output_dtype=settings.get("voyageai_output_dtype", "float"), + output_dimension=settings.get("dimensions"), ) - return res.embeddings[0] # type: ignore[return-value] + except VoyageError as e: + raise ModelAPIError(model_name=self.model_name, message=str(e)) from e - async def embed_documents(self, texts: list[str]) -> list[list[float]]: - """Embed documents/chunks for indexing.""" - if not texts: - return [] - client = Client() - res = client.embed( - texts, model=self._model, input_type="document", output_dtype="float" - ) - return res.embeddings # type: ignore[return-value] + return EmbeddingResult( + embeddings=response.embeddings, + inputs=inputs, + input_type=input_type, + usage=_map_usage(response.total_tokens, self.model_name), + model_name=self.model_name, + provider_name=self.system, + ) -except ImportError: - pass + async def max_input_tokens(self) -> int | None: + return _MAX_INPUT_TOKENS.get(self.model_name) + + +def _map_usage(total_tokens: int, model: str) -> RequestUsage: + usage_data = {"total_tokens": total_tokens} + response_data = {"model": model, "usage": usage_data} + + return RequestUsage.extract( + response_data, + provider="voyageai", + provider_url="https://api.voyageai.com", + provider_fallback="voyageai", + api_flavor="embeddings", + )