mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(embedding): default OpenAI-path encoding_format to float
Made-with: Cursor
This commit is contained in:
parent
eab0075353
commit
8473b70dd8
2 changed files with 56 additions and 5 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue