mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(proxy): strictly honor disable_end_user_cost_tracking in SpendLogs
This commit is contained in:
parent
2655d1dd5e
commit
0d08939269
3 changed files with 92 additions and 3 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue