feat(embedding): default OpenAI-path encoding_format to float

Made-with: Cursor
This commit is contained in:
Sameer Kankute 2026-05-01 16:26:17 +05:30
parent eab0075353
commit 8473b70dd8
No known key found for this signature in database
2 changed files with 56 additions and 5 deletions

View file

@ -4920,11 +4920,12 @@ def embedding( # noqa: PLR0915
if headers is not None and headers != {}:
optional_params["extra_headers"] = headers
if encoding_format is not None:
optional_params["encoding_format"] = encoding_format
else:
# Omiting causes openai sdk to add default value of "float"
optional_params["encoding_format"] = None
optional_params["encoding_format"] = (
encoding_format
or optional_params.get("encoding_format")
or get_secret_str("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT")
or "float"
)
api_version = None

View file

@ -0,0 +1,50 @@
from unittest.mock import MagicMock, patch
import pytest
from litellm import embedding
@pytest.mark.parametrize(
"set_env, env_value, expected",
[
(False, None, "float"),
(True, "base64", "base64"),
],
)
def test_openai_embedding_encoding_format_default(
monkeypatch, set_env, env_value, expected
):
monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False)
if set_env:
monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", env_value)
mock_response = MagicMock()
mock_response.parse.return_value = MagicMock(
model_dump=lambda: {
"data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}],
"model": "text-embedding-ada-002",
"object": "list",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
}
)
mock_response.headers = {}
with patch(
"litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client"
) as mock_get_client:
mock_client_instance = MagicMock()
mock_get_client.return_value = mock_client_instance
mock_client_instance.embeddings.with_raw_response.create.return_value = (
mock_response
)
embedding(
model="text-embedding-ada-002",
input="Hello world",
)
call_kwargs = (
mock_client_instance.embeddings.with_raw_response.create.call_args[1]
)
assert call_kwargs["encoding_format"] == expected