mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
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:
commit
56b51db451
3 changed files with 297 additions and 9 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue