diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index c868d3d22b2..c8c0ce6c8e1 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1113,6 +1113,9 @@ def _coerce_user_id_to_str(value: Any) -> Optional[str]: def get_end_user_id_from_request_body( request_body: dict, request_headers: Optional[dict] = None ) -> Optional[str]: + if litellm.disable_end_user_cost_tracking: + return None + # Import general_settings here to avoid potential circular import issues at module level # and to ensure it's fetched at runtime. from litellm.proxy.proxy_server import general_settings diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index d215294fd04..ba86c41964f 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -298,9 +298,10 @@ def get_logging_payload( # noqa: PLR0915 or standard_logging_payload["metadata"].get("user_api_key_hash") or "" ) - end_user_id = end_user_id or standard_logging_payload["metadata"].get( - "user_api_key_end_user_id" - ) + if not litellm.disable_end_user_cost_tracking: + end_user_id = end_user_id or standard_logging_payload["metadata"].get( + "user_api_key_end_user_id" + ) # BUG FIX: Don't overwrite api_key when standard_logging_payload is None # The api_key was already extracted from metadata (line 243) and hashed (lines 256-259) request_tags = ( diff --git a/tests/proxy_unit_tests/test_disable_end_user_cost_tracking_proxy.py b/tests/proxy_unit_tests/test_disable_end_user_cost_tracking_proxy.py new file mode 100644 index 00000000000..e7350eddbb1 --- /dev/null +++ b/tests/proxy_unit_tests/test_disable_end_user_cost_tracking_proxy.py @@ -0,0 +1,85 @@ +import pytest +import asyncio +import os, sys +from unittest.mock import MagicMock, patch + +# Add the project root to sys.path +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../"))) + +import litellm +from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload +from datetime import datetime + +@pytest.mark.asyncio +async def test_get_logging_payload_honors_disable_flag(): + """ + Test that get_logging_payload correctly suppresses end_user_id + when litellm.disable_end_user_cost_tracking is True. + """ + # 1. Setup + litellm.disable_end_user_cost_tracking = True + + kwargs = { + "litellm_params": { + "metadata": { + "user_api_key_end_user_id": "test-user-123" + } + }, + "call_type": "completion", + "standard_logging_object": { + "metadata": { + "user_api_key_end_user_id": "test-user-123" + }, + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + "model_map_information": {} + } + } + + response_obj = {"id": "chatcmpl-123", "usage": {"total_tokens": 15}} + start_time = datetime.now() + end_time = datetime.now() + + # 2. Execute + payload = get_logging_payload(kwargs, response_obj, start_time, end_time) + + # 3. Assert + assert payload["end_user"] == "" + + # 4. Cleanup + litellm.disable_end_user_cost_tracking = False + +@pytest.mark.asyncio +async def test_get_logging_payload_tracks_when_not_disabled(): + """ + Test that get_logging_payload correctly includes end_user_id + when litellm.disable_end_user_cost_tracking is False. + """ + # 1. Setup + litellm.disable_end_user_cost_tracking = False + + kwargs = { + "litellm_params": { + "metadata": { + "user_api_key_end_user_id": "test-user-456" + } + }, + "call_type": "completion", + "standard_logging_object": { + "metadata": { + "user_api_key_end_user_id": "test-user-456" + }, + "model_map_information": {} + } + } + + response_obj = {"id": "chatcmpl-456"} + start_time = datetime.now() + end_time = datetime.now() + + # 2. Execute + payload = get_logging_payload(kwargs, response_obj, start_time, end_time) + + # 3. Assert + assert payload["end_user"] == "test-user-456"