mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): reuse embeddings processor on failures
This commit is contained in:
parent
3fb91d2c37
commit
caa07aa3f5
3 changed files with 84 additions and 47 deletions
|
|
@ -8,8 +8,8 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
get_litellm_metadata_from_kwargs,
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost
|
||||
|
|
@ -207,8 +207,6 @@ class _ProxyDBLogger(CustomLogger):
|
|||
) -> dict:
|
||||
merged_metadata: dict = {}
|
||||
existing_litellm_params = request_data.get("litellm_params", {}) or {}
|
||||
trusted_metadata_key = _ProxyDBLogger._get_failure_metadata_variable_name(request_data=request_data)
|
||||
|
||||
def merge_metadata(
|
||||
metadata: Any,
|
||||
*,
|
||||
|
|
@ -229,7 +227,11 @@ class _ProxyDBLogger(CustomLogger):
|
|||
merged_metadata[key] = value
|
||||
|
||||
merge_metadata(
|
||||
request_data.get(trusted_metadata_key, {}),
|
||||
request_data.get("litellm_metadata", {}),
|
||||
skip_user_api_key_fields=True,
|
||||
)
|
||||
merge_metadata(
|
||||
request_data.get("metadata", {}),
|
||||
overwrite=True,
|
||||
skip_user_api_key_fields=True,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -96,6 +96,36 @@ def embedding_body_read_raises(monkeypatch):
|
|||
yield
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def embedding_token_decode_raises(monkeypatch):
|
||||
router = MagicMock()
|
||||
router.model_names = ["text-embedding-ada-002"]
|
||||
router.get_deployment_by_model_group_name = MagicMock(
|
||||
return_value={"litellm_params": {"model": "custom/provider-model"}}
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()))
|
||||
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
def _raise_decode(*args, **kwargs):
|
||||
raise ValueError("token-decode")
|
||||
|
||||
async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj, version=None):
|
||||
assert self._failure_call_type == "aembedding"
|
||||
assert self.data["model"] == "text-embedding-ada-002"
|
||||
assert self.data["input"] == [[1, 2, 3]]
|
||||
return ProxyException(message="token-decode", type="bad_request_error", param="input", code=400)
|
||||
|
||||
monkeypatch.setattr(proxy_server.litellm, "decode", _raise_decode)
|
||||
monkeypatch.setattr(
|
||||
common_request_processing.ProxyBaseLLMRequestProcessing,
|
||||
"_handle_llm_api_exception",
|
||||
_handler,
|
||||
)
|
||||
yield
|
||||
|
||||
|
||||
_EMBED_PATHS = [
|
||||
"/v1/embeddings",
|
||||
"/embeddings",
|
||||
|
|
@ -146,3 +176,14 @@ def test_embeddings_body_read_error_preserves_call_type(client, auth_as, embeddi
|
|||
response = client.post(path, json=payload)
|
||||
assert response.status_code == 400
|
||||
assert response.content
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", _EMBED_PATHS)
|
||||
def test_embeddings_preprocessing_error_preserves_parsed_data(
|
||||
client, auth_as, embedding_token_decode_raises, path
|
||||
):
|
||||
payload = {"model": "text-embedding-ada-002", "input": [[1, 2, 3]]}
|
||||
with auth_as():
|
||||
response = client.post(path, json=payload)
|
||||
assert response.status_code == 400
|
||||
assert response.content
|
||||
|
|
|
|||
|
|
@ -56,55 +56,49 @@ async def test_embedding_metadata_population():
|
|||
Test that the embedding endpoint correctly populates metadata
|
||||
from UserAPIKeyAuth.
|
||||
"""
|
||||
captured_data = {}
|
||||
|
||||
async def mock_base_process(self, *args, **kwargs):
|
||||
captured_data.update(self.data)
|
||||
return {"data": []}
|
||||
|
||||
# Setup
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request"
|
||||
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request",
|
||||
new=mock_base_process,
|
||||
):
|
||||
# Create a mock UserAPIKeyAuth object
|
||||
mock_user_auth = MagicMock(spec=UserAPIKeyAuth)
|
||||
mock_user_auth.user_id = "test_user_id_emb"
|
||||
mock_user_auth.team_id = "test_team_id_emb"
|
||||
mock_user_auth.org_id = "test_org_id_emb"
|
||||
|
||||
# Create a mock Request object
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.json = AsyncMock(
|
||||
return_value={"model": "gpt-3.5-turbo", "input": "hello"}
|
||||
)
|
||||
# Mock _read_request_body to return our data
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.__init__",
|
||||
return_value=None,
|
||||
) as mock_base_process_init:
|
||||
# Create a mock UserAPIKeyAuth object
|
||||
mock_user_auth = MagicMock(spec=UserAPIKeyAuth)
|
||||
mock_user_auth.user_id = "test_user_id_emb"
|
||||
mock_user_auth.team_id = "test_team_id_emb"
|
||||
mock_user_auth.org_id = "test_org_id_emb"
|
||||
|
||||
# Create a mock Request object
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.json = AsyncMock(
|
||||
return_value={"model": "gpt-3.5-turbo", "input": "hello"}
|
||||
"litellm.proxy.proxy_server._read_request_body",
|
||||
new=AsyncMock(return_value={"model": "gpt-3.5-turbo", "input": "hello"}),
|
||||
):
|
||||
# Call the endpoint function directly
|
||||
await embeddings(
|
||||
request=mock_request,
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
# Mock _read_request_body to return our data
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server._read_request_body",
|
||||
new=AsyncMock(
|
||||
return_value={"model": "gpt-3.5-turbo", "input": "hello"}
|
||||
),
|
||||
):
|
||||
# Call the endpoint function directly
|
||||
await embeddings(
|
||||
request=mock_request,
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
# Check if ProxyBaseLLMRequestProcessing was initialized with the correct metadata
|
||||
mock_base_process_init.assert_called_once()
|
||||
call_args = mock_base_process_init.call_args
|
||||
# handle both positional and keyword args for data
|
||||
if "data" in call_args.kwargs:
|
||||
data_arg = call_args.kwargs["data"]
|
||||
else:
|
||||
data_arg = call_args.args[0]
|
||||
|
||||
assert (
|
||||
data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_emb"
|
||||
)
|
||||
assert (
|
||||
data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_emb"
|
||||
)
|
||||
assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id_emb"
|
||||
assert (
|
||||
captured_data["metadata"]["user_api_key_user_id"]
|
||||
== "test_user_id_emb"
|
||||
)
|
||||
assert (
|
||||
captured_data["metadata"]["user_api_key_team_id"]
|
||||
== "test_team_id_emb"
|
||||
)
|
||||
assert captured_data["metadata"]["user_api_key_org_id"] == "test_org_id_emb"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue