Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions hindsight-api-slim/hindsight_api/engine/embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -1312,6 +1312,8 @@ async def initialize(self) -> None:
embed_kwargs["dimensions"] = self.output_dimensions
if self.model.startswith("openai/"):
embed_kwargs["allowed_openai_params"] = ["dimensions"]
if self.model.startswith("voyage/"):
embed_kwargs["input_type"] = "document"

# Use async embedding method (standard in litellm)
response = await self._litellm.aembedding(**embed_kwargs)
Expand All @@ -1328,6 +1330,22 @@ async def initialize(self) -> None:
logger.info(f"Embeddings: LiteLLM SDK provider initialized (model: {self.model}, dim: {self._dimension})")

def encode(self, texts: list[str]) -> list[list[float]]:
"""Generate embeddings with provider-default semantics."""
return self._encode_with_input_type(texts)

def encode_query(self, texts: list[str]) -> list[list[float]]:
"""Generate query-side embeddings for asymmetric Voyage retrieval."""
input_type = "query" if self.model.startswith("voyage/") else None
return self._encode_with_input_type(texts, input_type)

def encode_documents(self, texts: list[str]) -> list[list[float]]:
"""Generate document-side embeddings for asymmetric Voyage retrieval."""
input_type = "document" if self.model.startswith("voyage/") else None
return self._encode_with_input_type(texts, input_type)

def _encode_with_input_type(
self, texts: list[str], input_type: Literal["query", "document"] | None = None
) -> list[list[float]]:
"""
Generate embeddings using the LiteLLM SDK.

Expand Down Expand Up @@ -1365,6 +1383,8 @@ def encode(self, texts: list[str]) -> list[list[float]]:
embed_kwargs["dimensions"] = self.output_dimensions
if self.model.startswith("openai/"):
embed_kwargs["allowed_openai_params"] = ["dimensions"]
if input_type is not None:
embed_kwargs["input_type"] = input_type

# Use sync embedding (litellm doesn't have async in thread-safe way)
response = self._litellm.embedding(**embed_kwargs)
Expand Down
44 changes: 44 additions & 0 deletions hindsight-api-slim/tests/test_litellm_sdk_embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,50 @@ async def test_initialization_without_api_key(self, mock_litellm):
call_kwargs = mock_litellm.aembedding.call_args.kwargs
assert "api_key" not in call_kwargs

async def test_voyage_initialization_uses_document_input_type(self, mock_litellm):
with patch(
"builtins.__import__",
side_effect=lambda name, *args: mock_litellm if name == "litellm" else __import__(name, *args),
):
emb = LiteLLMSDKEmbeddings(
api_key="test_key",
model="voyage/voyage-4-large",
output_dimensions=2048,
encoding_format=None,
)
await emb.initialize()

call_kwargs = mock_litellm.aembedding.call_args.kwargs
assert call_kwargs["input_type"] == "document"
assert call_kwargs["dimensions"] == 2048
assert "encoding_format" not in call_kwargs

async def test_voyage_query_and_document_input_types(self, mock_litellm):
emb = LiteLLMSDKEmbeddings(
api_key="test_key",
model="voyage/voyage-4-large",
output_dimensions=2048,
encoding_format=None,
)
emb._litellm = mock_litellm
emb._dimension = 2048
mock_litellm.embedding.return_value.data = [{"embedding": [0.5] * 2048, "index": 0}]

emb.encode_documents(["document"])
assert mock_litellm.embedding.call_args.kwargs["input_type"] == "document"

emb.encode_query(["query"])
assert mock_litellm.embedding.call_args.kwargs["input_type"] == "query"

async def test_non_voyage_query_keeps_provider_default(self, mock_litellm):
emb = LiteLLMSDKEmbeddings(api_key="test_key", model="cohere/embed-english-v3.0")
emb._litellm = mock_litellm
emb._dimension = 768
mock_litellm.embedding.return_value.data = [{"embedding": [0.5] * 768, "index": 0}]

emb.encode_query(["query"])
assert "input_type" not in mock_litellm.embedding.call_args.kwargs

async def test_encode_without_api_key(self, mock_litellm):
"""Test encode omits api_key when not set (IAM/ambient credentials)."""
emb = LiteLLMSDKEmbeddings(
Expand Down