mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 1641dfe3ee into 80d3b69d9c
This commit is contained in:
commit
610f8edeaf
3 changed files with 90 additions and 3 deletions
1
.github/workflows/test-unit-proxy-db.yml
vendored
1
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -173,6 +173,7 @@ jobs:
|
|||
tests/proxy_unit_tests/test_update_daily_tag_spend.py
|
||||
tests/proxy_unit_tests/test_update_spend.py
|
||||
tests/proxy_unit_tests/test_proxy_encrypt_decrypt.py
|
||||
tests/proxy_unit_tests/test_disable_end_user_cost_tracking_proxy.py
|
||||
workers: 4
|
||||
dist: loadscope
|
||||
timeout: 15
|
||||
|
|
|
|||
|
|
@ -306,9 +306,10 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
|
|||
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