mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Changes Made
1. litellm/llms/openai/openai.py — P1 fix + anti-pattern removal - _normalize_embedding_response: Replaced the silent list-wrapping fallback (which returned a fake dict with zeroed usage) with an OpenAIError raise. Now when raw_response.parse() returns a list (indicating a misconfigured endpoint), the method raises a clear error with guidance to check the api_base includes the API prefix (e.g. /v1). This is consistent with the openai_like handler behavior. - Removed the try/except Exception as e: raise e no-op anti-pattern from both make_openai_embedding_request and make_sync_openai_embedding_request. 2. tests/test_litellm/llms/openai_like/embedding/test_openai_like_embedding.py — P2 fix - Added test_async_embedding_response_list_raises_type_error: verifies await handler.aembedding(...) raises OpenAILikeError for list responses - Added test_async_embedding_response_dict_uses_convert_to_model_response_object: verifies await handler.aembedding(...) correctly processes dict responses 3. tests/test_litellm/llms/openai/test_openai_common_utils.py — Coverage fix - Added test_normalize_embedding_response_raises_for_list: verifies OpenAIChatCompletion._normalize_embedding_response() raises OpenAIError when given a list Test results - 10/10 openai_like embedding tests pass (including 2 new async tests) - 17/17 openai common utils tests pass (including 1 new error path test) - 4/4 openai empty response tests pass - No regressions The changes address all three Greptile findings (P1, P2) and should bring patch coverage above the 71.33% threshold by eliminating the untested list-wrapping fallback and adding tests for both the error path and async handler path.
This commit is contained in:
parent
1c82e5da61
commit
dc6e1cc680
3 changed files with 114 additions and 26 deletions
|
|
@ -1204,15 +1204,12 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
|
||||
Returns (headers, response_dict) where response_dict is always a dict.
|
||||
"""
|
||||
try:
|
||||
raw_response = await openai_aclient.embeddings.with_raw_response.create(
|
||||
**data, timeout=timeout
|
||||
) # type: ignore
|
||||
headers = dict(raw_response.headers)
|
||||
parsed = raw_response.parse()
|
||||
return headers, self._normalize_embedding_response(parsed, data)
|
||||
except Exception as e:
|
||||
raise e
|
||||
raw_response = await openai_aclient.embeddings.with_raw_response.create(
|
||||
**data, timeout=timeout
|
||||
) # type: ignore
|
||||
headers = dict(raw_response.headers)
|
||||
parsed = raw_response.parse()
|
||||
return headers, self._normalize_embedding_response(parsed, data)
|
||||
|
||||
def _normalize_embedding_response(
|
||||
self,
|
||||
|
|
@ -1223,12 +1220,16 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
return response.model_dump()
|
||||
if isinstance(response, dict):
|
||||
return response
|
||||
return {
|
||||
"object": "list",
|
||||
"data": list(response) if isinstance(response, list) else [],
|
||||
"model": request_data.get("model", ""),
|
||||
"usage": {"prompt_tokens": 0, "total_tokens": 0},
|
||||
}
|
||||
raise OpenAIError(
|
||||
status_code=500,
|
||||
message=(
|
||||
"Embedding response is not a mapping or Pydantic model. "
|
||||
"If you are using an OpenAI-compatible server, "
|
||||
"make sure the api_base includes the API prefix (e.g. /v1). "
|
||||
f"Received type: {type(response).__name__}. "
|
||||
f"Response: {str(response)[:500]}"
|
||||
),
|
||||
)
|
||||
|
||||
@track_llm_api_timing()
|
||||
def make_sync_openai_embedding_request(
|
||||
|
|
@ -1245,16 +1246,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
|
||||
Returns (headers, response_dict) where response_dict is always a dict.
|
||||
"""
|
||||
try:
|
||||
raw_response = openai_client.embeddings.with_raw_response.create(
|
||||
**data, timeout=timeout
|
||||
) # type: ignore
|
||||
raw_response = openai_client.embeddings.with_raw_response.create(
|
||||
**data, timeout=timeout
|
||||
) # type: ignore
|
||||
|
||||
headers = dict(raw_response.headers)
|
||||
parsed = raw_response.parse()
|
||||
return headers, self._normalize_embedding_response(parsed, data)
|
||||
except Exception as e:
|
||||
raise e
|
||||
headers = dict(raw_response.headers)
|
||||
parsed = raw_response.parse()
|
||||
return headers, self._normalize_embedding_response(parsed, data)
|
||||
|
||||
async def aembedding(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, call, patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -175,3 +175,18 @@ def test_get_openai_client_cache_key(client_type):
|
|||
)
|
||||
assert isinstance(key, str)
|
||||
assert "api_key=sk-test" in key
|
||||
|
||||
|
||||
def test_normalize_embedding_response_raises_for_list():
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
from litellm.llms.openai.openai import OpenAIChatCompletion
|
||||
|
||||
handler = OpenAIChatCompletion()
|
||||
|
||||
with pytest.raises(OpenAIError) as exc_info:
|
||||
handler._normalize_embedding_response(
|
||||
response=[{"embedding": [0.1, 0.2]}],
|
||||
request_data={"model": "test-model"},
|
||||
)
|
||||
assert "not a mapping" in str(exc_info.value.message)
|
||||
assert "api_base includes the API prefix" in str(exc_info.value.message)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ Test cases for OpenAI-like embedding handler
|
|||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -416,3 +416,78 @@ class TestOpenAILikeEmbeddingHandler:
|
|||
assert response.object == "list"
|
||||
assert len(response.data) == 1
|
||||
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_embedding_response_list_raises_type_error(self):
|
||||
handler = OpenAILikeEmbeddingHandler()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_response = Mock()
|
||||
mock_response.json.return_value = [
|
||||
{"object": "embedding", "embedding": [0.1, 0.2, 0.3], "index": 0}
|
||||
]
|
||||
mock_response.raise_for_status = Mock()
|
||||
mock_client.post.return_value = mock_response
|
||||
|
||||
mock_logging = MagicMock()
|
||||
model_response = EmbeddingResponse()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.openai_like.embedding.handler.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
with pytest.raises(OpenAILikeError) as exc_info:
|
||||
await handler.aembedding(
|
||||
input=["test input"],
|
||||
data={"model": "test-model", "input": ["test input"]},
|
||||
model_response=model_response,
|
||||
timeout=60.0,
|
||||
logging_obj=mock_logging,
|
||||
api_key="test-key",
|
||||
api_base="http://test.com/v1/embeddings",
|
||||
headers={},
|
||||
client=None,
|
||||
)
|
||||
assert "not a mapping" in str(exc_info.value.message)
|
||||
assert "api_base includes the API prefix" in str(exc_info.value.message)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_embedding_response_dict_uses_convert_to_model_response_object(
|
||||
self,
|
||||
):
|
||||
handler = OpenAILikeEmbeddingHandler()
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_response = Mock()
|
||||
mock_response.json.return_value = {
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "embedding": [0.1, 0.2, 0.3], "index": 0}],
|
||||
"model": "test-model",
|
||||
"usage": {"prompt_tokens": 5, "total_tokens": 5},
|
||||
}
|
||||
mock_response.raise_for_status = Mock()
|
||||
mock_client.post.return_value = mock_response
|
||||
|
||||
mock_logging = MagicMock()
|
||||
model_response = EmbeddingResponse()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.openai_like.embedding.handler.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
response = await handler.aembedding(
|
||||
input=["test input"],
|
||||
data={"model": "test-model", "input": ["test input"]},
|
||||
model_response=model_response,
|
||||
timeout=60.0,
|
||||
logging_obj=mock_logging,
|
||||
api_key="test-key",
|
||||
api_base="http://test.com/v1/embeddings",
|
||||
headers={},
|
||||
client=None,
|
||||
)
|
||||
|
||||
assert response.model == "test-model"
|
||||
assert response.object == "list"
|
||||
assert len(response.data) == 1
|
||||
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue