fix(proxy): reuse embeddings processor on failures

This commit is contained in:
silencedoctor 2026-07-29 17:51:50 +08:00
parent 3fb91d2c37
commit caa07aa3f5
3 changed files with 84 additions and 47 deletions

View file

@ -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,
)

View file

@ -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

View file

@ -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