Merge pull request #35091 from fzowl/feat/voyage-context-4

fix(voyage): accept flat list[str] input for contextual embeddings
This commit is contained in:
Mateo Wang 2026-09-10 14:29:23 -07:00 committed by GitHub
commit 56b51db451
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 297 additions and 9 deletions

View file

@ -3,6 +3,7 @@ This module is used to transform the request and response for the Voyage context
This would be used for all the contextualized embeddings models in Voyage.
"""
from collections.abc import Mapping
from typing import Final
import httpx
@ -24,7 +25,10 @@ class VoyageError(BaseLLMException):
):
self.status_code = status_code
self.message = message
self.request = httpx.Request(method="POST", url="https://api.voyageai.com/v1/contextualizedembeddings")
self.request = httpx.Request(
method="POST",
url="https://api.voyageai.com/v1/contextualizedembeddings",
)
self.response = httpx.Response(status_code=status_code, request=self.request)
super().__init__(
status_code=status_code,
@ -56,16 +60,16 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig):
return api_base
return "https://api.voyageai.com/v1/contextualizedembeddings"
def get_supported_openai_params(self, model: str) -> list:
def get_supported_openai_params(self, model: str) -> list: # mutable-ok: base class signature
return ["encoding_format", "dimensions"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
non_default_params: dict, # mutable-ok: base class signature
optional_params: dict, # mutable-ok: base class signature
model: str,
drop_params: bool,
) -> dict:
) -> dict: # mutable-ok: base class signature
"""
Map OpenAI params to Voyage params
@ -79,7 +83,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig):
def validate_environment(
self,
headers: dict,
headers: dict, # mutable-ok: base class signature
model: str,
messages: list[AllMessageValues],
optional_params: dict,
@ -97,6 +101,8 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig):
"Authorization": f"Bearer {api_key}",
}
AUTO_CHUNK_SIZE: Final = 32000
def transform_embedding_request(
self,
model: str,
@ -105,11 +111,27 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig):
headers: dict,
) -> dict:
return {
"inputs": input,
"inputs": [input] if isinstance(input, str) else input,
"model": model,
**self._auto_chunk_params(input, optional_params),
**optional_params,
}
@classmethod
def _auto_chunk_params(
cls,
input: AllEmbeddingInputValues | list[list[str]],
optional_params: Mapping[str, object],
) -> Mapping[str, object]:
is_flat: Final = isinstance(input, str) or all(isinstance(item, str) for item in input)
if not is_flat or optional_params.get("input_type") == "query":
return {}
return {
"enable_auto_chunking": True,
"chunk_size": cls.AUTO_CHUNK_SIZE,
"input_type": "document",
}
def transform_embedding_response(
self,
model: str,
@ -124,9 +146,11 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig):
try:
raw_response_json: Final = raw_response.json()
except Exception:
raise VoyageError(message=raw_response.text, status_code=raw_response.status_code)
raise VoyageError(
message=raw_response.text,
status_code=raw_response.status_code,
)
# model_response.usage
model_response.model = raw_response_json.get("model")
model_response.data = raw_response_json.get("data")
model_response.object = raw_response_json.get("object")

View file

@ -139,6 +139,7 @@ class TestVoyageContextualEmbeddings:
# Test contextual model detection
assert config.is_contextualized_embeddings("voyage-context-3") is True
assert config.is_contextualized_embeddings("voyage-context-4") is True
assert config.is_contextualized_embeddings("voyage-context-2") is True
assert config.is_contextualized_embeddings("context-model") is True

View file

@ -0,0 +1,263 @@
import json
from unittest.mock import MagicMock
import pytest
class TestVoyageContextualEmbeddings:
def test_contextual_model_detection(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
assert VoyageContextualEmbeddingConfig.is_contextualized_embeddings("voyage-context-3")
assert VoyageContextualEmbeddingConfig.is_contextualized_embeddings("voyage-context-4")
assert not VoyageContextualEmbeddingConfig.is_contextualized_embeddings("voyage-3-lite")
def test_url_generation(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
assert (
config.get_complete_url(None, None, "voyage-context-4", {}, {})
== "https://api.voyageai.com/v1/contextualizedembeddings"
)
assert (
config.get_complete_url("https://custom.api.com", None, "voyage-context-4", {}, {})
== "https://custom.api.com/contextualizedembeddings"
)
assert (
config.get_complete_url(
"https://custom.api.com/contextualizedembeddings",
None,
"voyage-context-4",
{},
{},
)
== "https://custom.api.com/contextualizedembeddings"
)
def test_get_supported_openai_params(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
assert config.get_supported_openai_params("voyage-context-4") == [
"encoding_format",
"dimensions",
]
def test_map_openai_params(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
result = config.map_openai_params(
{"encoding_format": "float", "dimensions": 512}, {}, "voyage-context-4", False
)
assert result["encoding_format"] == "float"
assert result["output_dimension"] == 512
def test_validate_environment_with_api_key(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
headers = config.validate_environment(
{}, "voyage-context-4", [], {}, {}, api_key="test-key"
)
assert headers == {"Authorization": "Bearer test-key"}
def test_validate_environment_secret_fallback(self, monkeypatch):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
monkeypatch.setenv("VOYAGE_API_KEY", "secret-key")
config = VoyageContextualEmbeddingConfig()
headers = config.validate_environment(
{}, "voyage-context-4", [], {}, {}, api_key=None
)
assert headers == {"Authorization": "Bearer secret-key"}
def test_nested_list_passthrough(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
nested = [["Hello", "world"], ["Test"]]
transformed = config.transform_embedding_request(
"voyage-context-4", nested, {}, {}
)
assert transformed["inputs"] == nested
assert transformed["model"] == "voyage-context-4"
assert "enable_auto_chunking" not in transformed
def test_flat_list_str_auto_chunked(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
transformed = config.transform_embedding_request(
"voyage-context-4", ["Hello", "world"], {}, {}
)
assert transformed["inputs"] == ["Hello", "world"]
assert transformed["enable_auto_chunking"] is True
assert transformed["chunk_size"] == 32000
assert transformed["input_type"] == "document"
def test_flat_list_str_query_no_auto_chunk(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
transformed = config.transform_embedding_request(
"voyage-context-4", ["Hello", "world"], {"input_type": "query"}, {}
)
assert transformed["inputs"] == ["Hello", "world"]
assert transformed["input_type"] == "query"
assert "enable_auto_chunking" not in transformed
def test_flat_list_str_document_preserves_input_type(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
transformed = config.transform_embedding_request(
"voyage-context-4", ["Hello"], {"input_type": "document"}, {}
)
assert transformed["input_type"] == "document"
assert transformed["enable_auto_chunking"] is True
def test_flat_list_str_caller_chunk_params_win(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
transformed = config.transform_embedding_request(
"voyage-context-4",
["Hello", "world"],
{"input_type": "document", "chunk_size": 512, "chunk_overlap": 32},
{},
)
assert transformed["enable_auto_chunking"] is True
assert transformed["chunk_size"] == 512
assert transformed["chunk_overlap"] == 32
assert transformed["input_type"] == "document"
def test_flat_list_str_caller_can_disable_auto_chunking(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
transformed = config.transform_embedding_request(
"voyage-context-4", ["Hello"], {"enable_auto_chunking": False}, {}
)
assert transformed["enable_auto_chunking"] is False
assert transformed["input_type"] == "document"
def test_nested_list_keeps_caller_params(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
transformed = config.transform_embedding_request(
"voyage-context-4", [["Hello", "world"]], {"input_type": "document", "output_dimension": 512}, {}
)
assert transformed == {
"inputs": [["Hello", "world"]],
"model": "voyage-context-4",
"input_type": "document",
"output_dimension": 512,
}
def test_single_string_auto_chunked(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
transformed = config.transform_embedding_request(
"voyage-context-4", "Hello", {}, {}
)
assert transformed["inputs"] == ["Hello"]
assert transformed["enable_auto_chunking"] is True
assert transformed["input_type"] == "document"
def test_single_string_query_no_auto_chunk(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
config = VoyageContextualEmbeddingConfig()
transformed = config.transform_embedding_request(
"voyage-context-4", "Hello", {"input_type": "query"}, {}
)
assert transformed["inputs"] == ["Hello"]
assert transformed["input_type"] == "query"
assert "enable_auto_chunking" not in transformed
def test_response_transformation(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
)
from litellm.types.utils import EmbeddingResponse
config = VoyageContextualEmbeddingConfig()
response_payload = {
"object": "list",
"data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}],
"model": "voyage-context-4",
"usage": {"total_tokens": 24},
}
raw_response = MagicMock()
raw_response.json.return_value = response_payload
raw_response.status_code = 200
raw_response.text = json.dumps(response_payload)
model_response = EmbeddingResponse()
transformed = config.transform_embedding_response(
"voyage-context-4", raw_response, model_response, MagicMock()
)
assert transformed.model == "voyage-context-4"
assert transformed.object == "list"
assert transformed.data == response_payload["data"]
assert transformed.usage.prompt_tokens == 24
assert transformed.usage.total_tokens == 24
def test_error_response_and_error_class(self):
from litellm.llms.voyage.embedding.transformation_contextual import (
VoyageContextualEmbeddingConfig,
VoyageError,
)
from litellm.types.utils import EmbeddingResponse
config = VoyageContextualEmbeddingConfig()
raw_response = MagicMock()
raw_response.json.side_effect = ValueError("not json")
raw_response.status_code = 400
raw_response.text = "bad request"
with pytest.raises(VoyageError) as exc_info:
config.transform_embedding_response(
"voyage-context-4", raw_response, EmbeddingResponse(), MagicMock()
)
assert exc_info.value.status_code == 400
assert exc_info.value.message == "bad request"
error = config.get_error_class("rate limited", 429, {"x-test": "1"})
assert isinstance(error, VoyageError)
assert error.status_code == 429
assert error.message == "rate limited"