fix(proxy): strictly honor disable_end_user_cost_tracking in SpendLogs

This commit is contained in:
Dushyant Acharya 2026-06-14 05:26:53 +05:30
parent 2655d1dd5e
commit 0d08939269
No known key found for this signature in database
3 changed files with 92 additions and 3 deletions

View file

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

View file

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

View file

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