diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index ed51d696462..f8bdd32c762 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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, ) diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py b/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py index b319afcff1f..c56b048946f 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py @@ -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 diff --git a/tests/test_litellm/proxy/test_chat_completion_metadata.py b/tests/test_litellm/proxy/test_chat_completion_metadata.py index 38dcdc13c50..5016992deef 100644 --- a/tests/test_litellm/proxy/test_chat_completion_metadata.py +++ b/tests/test_litellm/proxy/test_chat_completion_metadata.py @@ -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