mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
test(ollama): cover embedding prefix normalization
This commit is contained in:
parent
039b747318
commit
6e2861c3f9
1 changed files with 76 additions and 1 deletions
|
|
@ -2,7 +2,11 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.ollama.completion.handler import ollama_aembeddings, ollama_embeddings
|
||||
from litellm.llms.ollama.completion.handler import (
|
||||
_prepare_ollama_embedding_payload,
|
||||
ollama_aembeddings,
|
||||
ollama_embeddings,
|
||||
)
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
|
||||
|
|
@ -51,6 +55,52 @@ def test_ollama_embeddings(mock_response_data, mock_embedding_response, mock_enc
|
|||
assert response.usage.total_tokens == 5
|
||||
|
||||
|
||||
def test_prepare_ollama_embedding_payload_strips_ollama_prefix():
|
||||
payload = _prepare_ollama_embedding_payload(
|
||||
model="ollama/nomic-embed-text",
|
||||
prompts=["hello"],
|
||||
optional_params={},
|
||||
)
|
||||
|
||||
assert payload["model"] == "nomic-embed-text"
|
||||
assert payload["input"] == ["hello"]
|
||||
|
||||
|
||||
def test_prepare_ollama_embedding_payload_keeps_unprefixed_model():
|
||||
payload = _prepare_ollama_embedding_payload(
|
||||
model="nomic-embed-text",
|
||||
prompts=["hello"],
|
||||
optional_params={},
|
||||
)
|
||||
|
||||
assert payload["model"] == "nomic-embed-text"
|
||||
assert payload["input"] == ["hello"]
|
||||
|
||||
|
||||
def test_ollama_embeddings_strips_prefix_in_embed_payload(
|
||||
mock_response_data, mock_embedding_response, mock_encoding
|
||||
):
|
||||
with patch("litellm.module_level_client.post") as mock_post, patch(
|
||||
"litellm.OllamaConfig.get_config", return_value={"truncate": 512}
|
||||
):
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = mock_response_data
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
ollama_embeddings(
|
||||
api_base="http://localhost:11434",
|
||||
model="ollama/nomic-embed-text",
|
||||
prompts=["hello", "world"],
|
||||
optional_params={},
|
||||
model_response=mock_embedding_response,
|
||||
logging_obj=None,
|
||||
encoding=mock_encoding,
|
||||
)
|
||||
|
||||
assert mock_post.call_args.kwargs["json"]["model"] == "nomic-embed-text"
|
||||
assert mock_post.call_args.kwargs["json"]["input"] == ["hello", "world"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ollama_aembeddings(
|
||||
mock_response_data, mock_embedding_response, mock_encoding
|
||||
|
|
@ -80,6 +130,31 @@ async def test_ollama_aembeddings(
|
|||
assert response.usage.total_tokens == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ollama_aembeddings_strips_prefix_in_embed_payload(
|
||||
mock_response_data, mock_embedding_response, mock_encoding
|
||||
):
|
||||
mock_response = AsyncMock()
|
||||
mock_response.json = MagicMock(return_value=mock_response_data)
|
||||
with patch(
|
||||
"litellm.module_level_aclient.post", return_value=mock_response
|
||||
) as mock_post, patch(
|
||||
"litellm.OllamaConfig.get_config", return_value={"truncate": 512}
|
||||
):
|
||||
await ollama_aembeddings(
|
||||
api_base="http://localhost:11434",
|
||||
model="ollama/nomic-embed-text",
|
||||
prompts=["hello", "world"],
|
||||
optional_params={},
|
||||
model_response=mock_embedding_response,
|
||||
logging_obj=None,
|
||||
encoding=mock_encoding,
|
||||
)
|
||||
|
||||
assert mock_post.call_args.kwargs["json"]["model"] == "nomic-embed-text"
|
||||
assert mock_post.call_args.kwargs["json"]["input"] == ["hello", "world"]
|
||||
|
||||
|
||||
def test_prompt_eval_fallback_when_missing(mock_embedding_response, mock_encoding):
|
||||
response_data = {
|
||||
"embeddings": [[0.1, 0.2, 0.3]],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue