Merge pull request #18785 from BerriAI/litellm_user_promethus_metrics

[Feature] User Metrics for Promethus
This commit is contained in:
yuneng-jiang 2026-01-15 15:51:02 -08:00 • committed by GitHub
commit 6a7edd8f2b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 615 additions and 5 deletions

View file

@ -21,7 +21,7 @@ from typing import (
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth
from litellm.proxy._types import LiteLLM_TeamTable, LiteLLM_UserTable, UserAPIKeyAuth
from litellm.types.integrations.prometheus import *
from litellm.types.integrations.prometheus import _sanitize_prometheus_label_name
from litellm.types.utils import StandardLoggingPayload
@ -193,6 +193,30 @@ class PrometheusLogger(CustomLogger):
),
)
# Remaining Budget for User
self.litellm_remaining_user_budget_metric = self._gauge_factory(
"litellm_remaining_user_budget_metric",
"Remaining budget for user",
labelnames=self.get_labels_for_metric(
"litellm_remaining_user_budget_metric"
),
)
# Max Budget for User
self.litellm_user_max_budget_metric = self._gauge_factory(
"litellm_user_max_budget_metric",
"Maximum budget set for user",
labelnames=self.get_labels_for_metric("litellm_user_max_budget_metric"),
)
self.litellm_user_budget_remaining_hours_metric = self._gauge_factory(
"litellm_user_budget_remaining_hours_metric",
"Remaining hours for user budget to be reset",
labelnames=self.get_labels_for_metric(
"litellm_user_budget_remaining_hours_metric"
),
)
########################################
# LiteLLM Virtual API KEY metrics
########################################
@ -960,6 +984,7 @@ class PrometheusLogger(CustomLogger):
user_api_key_alias=user_api_key_alias,
litellm_params=litellm_params,
response_cost=response_cost,
user_id=user_id,
)
# set proxy virtual key rpm/tpm metrics
@ -1120,6 +1145,7 @@ class PrometheusLogger(CustomLogger):
user_api_key_alias: Optional[str],
litellm_params: dict,
response_cost: float,
user_id: Optional[str] = None,
):
_team_spend = litellm_params.get("metadata", {}).get(
"user_api_key_team_spend", None
@ -1134,6 +1160,14 @@ class PrometheusLogger(CustomLogger):
_api_key_max_budget = litellm_params.get("metadata", {}).get(
"user_api_key_max_budget", None
)
_user_spend = litellm_params.get("metadata", {}).get(
"user_api_key_user_spend", None
)
_user_max_budget = litellm_params.get("metadata", {}).get(
"user_api_key_user_max_budget", None
)
await self._set_api_key_budget_metrics_after_api_request(
user_api_key=user_api_key,
user_api_key_alias=user_api_key_alias,
@ -1150,6 +1184,13 @@ class PrometheusLogger(CustomLogger):
response_cost=response_cost,
)
await self._set_user_budget_metrics_after_api_request(
user_id=user_id,
user_spend=_user_spend,
user_max_budget=_user_max_budget,
response_cost=response_cost,
)
def _increment_top_level_request_and_spend_metrics(
self,
end_user_id: Optional[str],
@ -2229,6 +2270,37 @@ class PrometheusLogger(CustomLogger):
data_type="keys",
)
async def _initialize_user_budget_metrics(self):
"""
Initialize user budget metrics by reusing the generic pagination logic.
"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
verbose_logger.debug(
"Prometheus: skipping user metrics initialization, DB not initialized"
)
return
async def fetch_users(
page_size: int, page: int
) -> Tuple[List[LiteLLM_UserTable], Optional[int]]:
skip = (page - 1) * page_size
users = await prisma_client.db.litellm_usertable.find_many(
skip=skip,
take=page_size,
order={"created_at": "desc"},
)
total_count = await prisma_client.db.litellm_usertable.count()
return users, total_count
await self._initialize_budget_metrics(
data_fetch_function=fetch_users,
set_metrics_function=self._set_user_list_budget_metrics,
data_type="users",
)
async def initialize_remaining_budget_metrics(self):
"""
Handler for initializing remaining budget metrics for all teams to avoid metric discrepancies.
@ -2261,11 +2333,12 @@ class PrometheusLogger(CustomLogger):
async def _initialize_remaining_budget_metrics(self):
"""
Helper to initialize remaining budget metrics for all teams and API keys.
Helper to initialize remaining budget metrics for all teams, API keys, and users.
"""
verbose_logger.debug("Emitting key, team budget metrics....")
verbose_logger.debug("Emitting key, team, user budget metrics....")
await self._initialize_team_budget_metrics()
await self._initialize_api_key_budget_metrics()
await self._initialize_user_budget_metrics()
async def _set_key_list_budget_metrics(
self, keys: List[Union[str, UserAPIKeyAuth]]
@ -2280,6 +2353,11 @@ class PrometheusLogger(CustomLogger):
for team in teams:
self._set_team_budget_metrics(team)
async def _set_user_list_budget_metrics(self, users: List[LiteLLM_UserTable]):
"""Helper function to set budget metrics for a list of users"""
for user in users:
self._set_user_budget_metrics(user)
async def _set_team_budget_metrics_after_api_request(
self,
user_api_team: Optional[str],
@ -2497,6 +2575,122 @@ class PrometheusLogger(CustomLogger):
return user_api_key_dict
async def _set_user_budget_metrics_after_api_request(
self,
user_id: Optional[str],
user_spend: Optional[float],
user_max_budget: Optional[float],
response_cost: float,
):
"""
Set user budget metrics after an LLM API request
- Assemble a LiteLLM_UserTable object
- looks up user info from db if not available in metadata
- Set user budget metrics
"""
if user_id:
user_object = await self._assemble_user_object(
user_id=user_id,
spend=user_spend,
max_budget=user_max_budget,
response_cost=response_cost,
)
self._set_user_budget_metrics(user_object)
async def _assemble_user_object(
self,
user_id: str,
spend: Optional[float],
max_budget: Optional[float],
response_cost: float,
) -> LiteLLM_UserTable:
"""
Assemble a LiteLLM_UserTable object
for fields not available in metadata, we fetch from db
Fields not available in metadata:
- `budget_reset_at`
"""
from litellm.proxy.auth.auth_checks import get_user_object
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
_total_user_spend = (spend or 0) + response_cost
user_object = LiteLLM_UserTable(
user_id=user_id,
spend=_total_user_spend,
max_budget=max_budget,
)
try:
user_info = await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
check_db_only=True,
)
except Exception as e:
verbose_logger.debug(
f"[Non-Blocking] Prometheus: Error getting user info: {str(e)}"
)
return user_object
if user_info:
user_object.budget_reset_at = user_info.budget_reset_at
return user_object
def _set_user_budget_metrics(
self,
user: LiteLLM_UserTable,
):
"""
Set user budget metrics for a single user
- Remaining Budget
- Max Budget
- Budget Reset At
"""
enum_values = UserAPIKeyLabelValues(
user=user.user_id,
)
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_remaining_user_budget_metric"
),
enum_values=enum_values,
)
self.litellm_remaining_user_budget_metric.labels(**_labels).set(
self._safe_get_remaining_budget(
max_budget=user.max_budget,
spend=user.spend,
)
)
if user.max_budget is not None:
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_user_max_budget_metric"
),
enum_values=enum_values,
)
self.litellm_user_max_budget_metric.labels(**_labels).set(user.max_budget)
if user.budget_reset_at is not None:
_labels = prometheus_label_factory(
supported_enum_labels=self.get_labels_for_metric(
metric_name="litellm_user_budget_remaining_hours_metric"
),
enum_values=enum_values,
)
self.litellm_user_budget_remaining_hours_metric.labels(**_labels).set(
self._get_remaining_hours_for_budget_reset(
budget_reset_at=user.budget_reset_at
)
)
def _get_remaining_hours_for_budget_reset(self, budget_reset_at: datetime) -> float:
"""
Get remaining hours for budget reset

View file

@ -2189,6 +2189,8 @@ class UserAPIKeyAuth(
user_tpm_limit: Optional[int] = None
user_rpm_limit: Optional[int] = None
user_email: Optional[str] = None
user_spend: Optional[float] = None
user_max_budget: Optional[float] = None
request_route: Optional[str] = None
user: Optional[Any] = None # Expanded user object when expand=user is used

View file

@ -1335,6 +1335,8 @@ async def _return_user_api_key_auth_obj(
user_tpm_limit=user_obj.tpm_limit,
user_rpm_limit=user_obj.rpm_limit,
user_email=user_obj.user_email,
user_spend=getattr(user_obj, "spend", None),
user_max_budget=getattr(user_obj, "max_budget", None),
)
if user_obj is not None and _is_user_proxy_admin(user_obj=user_obj):
user_api_key_kwargs.update(

View file

@ -1000,6 +1000,13 @@ async def add_litellm_data_to_request( # noqa: PLR0915
"user_api_key_model_max_budget"
] = user_api_key_dict.model_max_budget
# User spend, budget - used by prometheus.py
# Follow same pattern as team and API key budgets
data[_metadata_variable_name]["user_api_key_user_spend"] = user_api_key_dict.user_spend
data[_metadata_variable_name][
"user_api_key_user_max_budget"
] = user_api_key_dict.user_max_budget
data[_metadata_variable_name]["user_api_key_metadata"] = user_api_key_dict.metadata
_headers = dict(request.headers)
_headers.pop(

View file

@ -175,6 +175,9 @@ DEFINED_PROMETHEUS_METRICS = Literal[
"litellm_remaining_api_key_budget_metric",
"litellm_api_key_max_budget_metric",
"litellm_api_key_budget_remaining_hours_metric",
"litellm_remaining_user_budget_metric",
"litellm_user_max_budget_metric",
"litellm_user_budget_remaining_hours_metric",
"litellm_deployment_state",
"litellm_deployment_failure_responses",
"litellm_deployment_total_requests",
@ -421,6 +424,18 @@ class PrometheusMetricLabels:
litellm_remaining_api_key_budget_metric
)
litellm_remaining_user_budget_metric = [
UserAPIKeyLabelNames.USER.value,
]
litellm_user_max_budget_metric = [
UserAPIKeyLabelNames.USER.value,
]
litellm_user_budget_remaining_hours_metric = [
UserAPIKeyLabelNames.USER.value,
]
# Add deployment metrics
litellm_deployment_failure_responses = [
UserAPIKeyLabelNames.REQUESTED_MODEL.value,

View file

@ -1124,6 +1124,150 @@ def test_get_custom_labels_from_metadata_tags(monkeypatch):
assert get_custom_labels_from_metadata(metadata) == {}
def test_get_custom_labels_from_top_level_metadata(monkeypatch):
"""
Test that get_custom_labels_from_metadata can extract fields from top-level metadata,
such as requester_ip_address, not just from nested dictionaries like requester_metadata.
"""
monkeypatch.setattr(
"litellm.custom_prometheus_metadata_labels",
["requester_ip_address", "user_api_key_alias"],
)
# Simulate metadata structure with top-level fields
metadata = {
"requester_ip_address": "10.48.203.20", # Top-level field
"user_api_key_alias": "TestAlias", # Top-level field
"requester_metadata": {"nested_field": "nested_value"}, # Nested dict (excluded)
"user_api_key_auth_metadata": {"another_nested": "value"}, # Nested dict (excluded)
}
result = get_custom_labels_from_metadata(metadata)
assert result == {
"requester_ip_address": "10.48.203.20",
"user_api_key_alias": "TestAlias",
}
def test_get_custom_labels_from_top_level_and_nested_metadata(monkeypatch):
"""
Test that get_custom_labels_from_metadata can extract fields from both top-level
and nested metadata (requester_metadata, user_api_key_auth_metadata).
"""
monkeypatch.setattr(
"litellm.custom_prometheus_metadata_labels",
[
"requester_ip_address", # Top-level
"metadata.foo", # From requester_metadata
"metadata.bar", # From user_api_key_auth_metadata
],
)
# Simulate combined_metadata structure as it would appear after merging
# This is what gets passed to get_custom_labels_from_metadata
combined_metadata = {
"requester_ip_address": "10.48.203.20", # Top-level field
"foo": "bar_value", # From requester_metadata (spread)
"bar": "baz_value", # From user_api_key_auth_metadata (spread)
}
result = get_custom_labels_from_metadata(combined_metadata)
assert result == {
"requester_ip_address": "10.48.203.20",
"metadata_foo": "bar_value",
"metadata_bar": "baz_value",
}
async def test_async_log_success_event_with_top_level_metadata(prometheus_logger, monkeypatch):
"""
Test that async_log_success_event correctly extracts custom labels from top-level metadata
fields like requester_ip_address, not just from nested dictionaries.
"""
# Configure custom metadata labels to extract requester_ip_address
monkeypatch.setattr(
"litellm.custom_prometheus_metadata_labels", ["requester_ip_address"]
)
# Create standard logging payload with requester_ip_address at top-level metadata
standard_logging_object = create_standard_logging_payload()
standard_logging_object["metadata"]["requester_ip_address"] = "10.48.203.20"
standard_logging_object["metadata"]["requester_metadata"] = {} # Empty nested dict
standard_logging_object["metadata"]["user_api_key_auth_metadata"] = {} # Empty nested dict
kwargs = {
"model": "gpt-3.5-turbo",
"stream": True,
"litellm_params": {
"metadata": {
"user_api_key": "test_key",
"user_api_key_user_id": "test_user",
"user_api_key_team_id": "test_team",
"user_api_key_end_user_id": "test_end_user",
}
},
"start_time": datetime.now(),
"completion_start_time": datetime.now(),
"api_call_start_time": datetime.now(),
"end_time": datetime.now() + timedelta(seconds=1),
"standard_logging_object": standard_logging_object,
}
response_obj = MagicMock()
# Mock the prometheus client methods
# Create mock chain that accepts any labels (including custom labels like requester_ip_address)
def create_mock_metric():
mock_metric = MagicMock()
mock_labels = MagicMock()
mock_metric.labels = MagicMock(return_value=mock_labels)
mock_labels.inc = MagicMock()
mock_labels.observe = MagicMock()
mock_labels.set = MagicMock()
return mock_metric
prometheus_logger.litellm_requests_metric = create_mock_metric()
prometheus_logger.litellm_spend_metric = create_mock_metric()
prometheus_logger.litellm_tokens_metric = create_mock_metric()
prometheus_logger.litellm_input_tokens_metric = create_mock_metric()
prometheus_logger.litellm_output_tokens_metric = create_mock_metric()
prometheus_logger.litellm_remaining_team_budget_metric = create_mock_metric()
prometheus_logger.litellm_remaining_api_key_budget_metric = create_mock_metric()
prometheus_logger.litellm_remaining_user_budget_metric = create_mock_metric()
prometheus_logger.litellm_user_max_budget_metric = create_mock_metric()
prometheus_logger.litellm_user_budget_remaining_hours_metric = create_mock_metric()
prometheus_logger.litellm_remaining_api_key_requests_for_model = create_mock_metric()
prometheus_logger.litellm_remaining_api_key_tokens_for_model = create_mock_metric()
prometheus_logger.litellm_llm_api_time_to_first_token_metric = create_mock_metric()
prometheus_logger.litellm_llm_api_latency_metric = create_mock_metric()
prometheus_logger.litellm_request_total_latency_metric = create_mock_metric()
# Cache metrics
prometheus_logger.litellm_cache_hits_metric = create_mock_metric()
prometheus_logger.litellm_cache_misses_metric = create_mock_metric()
prometheus_logger.litellm_cached_tokens_metric = create_mock_metric()
# Deployment metrics
prometheus_logger.litellm_deployment_state = create_mock_metric()
prometheus_logger.litellm_deployment_success_responses = create_mock_metric()
prometheus_logger.litellm_deployment_total_requests = create_mock_metric()
prometheus_logger.litellm_deployment_latency_per_output_token = create_mock_metric()
prometheus_logger.litellm_remaining_requests_metric = create_mock_metric()
prometheus_logger.litellm_remaining_tokens_metric = create_mock_metric()
prometheus_logger.litellm_overhead_latency_metric = create_mock_metric()
prometheus_logger.litellm_proxy_total_requests_metric = create_mock_metric()
await prometheus_logger.async_log_success_event(
kwargs, response_obj, kwargs["start_time"], kwargs["end_time"]
)
# Verify that the metrics were called with labels
# The custom labels (like requester_ip_address) should be extracted and included in the label factory
# Since we're using mocks that accept any labels, we just verify that labels() was called
# This confirms that the custom label extraction logic ran without errors
assert prometheus_logger.litellm_requests_metric.labels.called
assert prometheus_logger.litellm_spend_metric.labels.called
# Verify that the labels() method was called with some arguments (either positional or keyword)
# This ensures the custom label extraction happened and didn't cause a "Incorrect label names" error
call_args = prometheus_logger.litellm_requests_metric.labels.call_args
assert call_args is not None
# The test passes if labels() was called successfully, which means custom labels were handled correctly
def test_get_custom_labels_from_tags(monkeypatch):
from litellm.integrations.prometheus import get_custom_labels_from_tags
@ -1410,18 +1554,28 @@ async def test_initialize_remaining_budget_metrics_exception_handling(
# Make get_paginated_teams raise an exception
mock_get_teams.side_effect = Exception("Database error")
mock_list_keys.side_effect = Exception("Key listing error")
# Mock prisma_client structure to raise an exception for user budget metrics
# The code accesses prisma_client.db.litellm_usertable.find_many and count
mock_usertable = MagicMock()
mock_usertable.find_many = MagicMock(side_effect=Exception("User database error"))
mock_usertable.count = MagicMock(side_effect=Exception("User count error"))
mock_db = MagicMock()
mock_db.litellm_usertable = mock_usertable
mock_prisma.db = mock_db
# Mock the Prometheus metrics
prometheus_logger.litellm_remaining_team_budget_metric = MagicMock()
prometheus_logger.litellm_remaining_api_key_budget_metric = MagicMock()
prometheus_logger.litellm_remaining_user_budget_metric = MagicMock()
# Mock the logger to capture the error
with patch("litellm._logging.verbose_logger.exception") as mock_logger:
# Call the function
await prometheus_logger._initialize_remaining_budget_metrics()
# Verify both errors were logged
assert mock_logger.call_count == 2
# Verify all three errors were logged (teams, keys, and users)
assert mock_logger.call_count == 3
assert (
"Error initializing teams budget metrics"
in mock_logger.call_args_list[0][0][0]
@ -1430,10 +1584,15 @@ async def test_initialize_remaining_budget_metrics_exception_handling(
"Error initializing keys budget metrics"
in mock_logger.call_args_list[1][0][0]
)
assert (
"Error initializing users budget metrics"
in mock_logger.call_args_list[2][0][0]
)
# Verify the metrics were never called
prometheus_logger.litellm_remaining_team_budget_metric.assert_not_called()
prometheus_logger.litellm_remaining_api_key_budget_metric.assert_not_called()
prometheus_logger.litellm_remaining_user_budget_metric.assert_not_called()
@pytest.mark.asyncio(scope="session")

View file

@ -442,6 +442,24 @@ async def get_key_info(session: aiohttp.ClientSession, key: str) -> Dict[str, An
return await response.json()
async def get_user_info(session: aiohttp.ClientSession, user_id: str) -> Dict[str, Any]:
"""Fetch user info and return the response"""
from urllib.parse import quote
# URL encode user_id to handle special characters
encoded_user_id = quote(user_id, safe="")
url = f"http://0.0.0.0:4000/user/info?user_id={encoded_user_id}"
headers = {
"Authorization": "Bearer sk-1234",
}
async with session.get(url, headers=headers) as response:
assert (
response.status == 200
), f"Failed to get user info. Status: {response.status}"
return await response.json()
def extract_key_budget_metrics(metrics_text: str, key_id: str) -> Dict[str, float]:
"""Extract budget-related metrics for a specific key"""
import re
@ -466,6 +484,33 @@ def extract_key_budget_metrics(metrics_text: str, key_id: str) -> Dict[str, floa
return metrics
def extract_user_budget_metrics(metrics_text: str, user_id: str) -> Dict[str, float]:
"""Extract budget-related metrics for a specific user"""
import re
metrics = {}
# Escape user_id for regex pattern matching
escaped_user_id = re.escape(user_id)
# Get remaining budget
remaining_pattern = f'litellm_remaining_user_budget_metric{{user="{escaped_user_id}"}} ([0-9.]+)'
remaining_match = re.search(remaining_pattern, metrics_text)
metrics["remaining"] = float(remaining_match.group(1)) if remaining_match else None
# Get total budget
total_pattern = f'litellm_user_max_budget_metric{{user="{escaped_user_id}"}} ([0-9.]+)'
total_match = re.search(total_pattern, metrics_text)
metrics["total"] = float(total_match.group(1)) if total_match else None
# Get remaining hours
hours_pattern = f'litellm_user_budget_remaining_hours_metric{{user="{escaped_user_id}"}} ([0-9.]+)'
hours_match = re.search(hours_pattern, metrics_text)
metrics["remaining_hours"] = float(hours_match.group(1)) if hours_match else None
return metrics
@pytest.mark.asyncio
async def test_key_budget_metrics():
"""
@ -476,6 +521,8 @@ async def test_key_budget_metrics():
4. Verify request costs are being tracked correctly
5. Verify prometheus metrics match /key/info spend data
"""
from datetime import datetime, timedelta, timezone
async with aiohttp.ClientSession() as session:
# Setup test key with unique alias
unique_alias = f"budget_test_key_{uuid.uuid4()}"
@ -483,6 +530,7 @@ async def test_key_budget_metrics():
"key_alias": unique_alias,
"max_budget": 10,
"budget_duration": "7d",
"budget_reset_at": (datetime.now(timezone.utc) + timedelta(days=7)).isoformat(),
}
key = await create_test_key_with_budget(session, key_data)
@ -543,6 +591,94 @@ async def test_key_budget_metrics():
), f"Spend mismatch: Prometheus={key_info_remaining_budget}, Key Info={first_budget['remaining']}"
@pytest.mark.asyncio
async def test_user_budget_metrics():
"""
Test user budget tracking metrics:
1. Create a user with max_budget
2. Make chat completion requests using OpenAI SDK with the user's key
3. Verify budget decreases over time
4. Verify request costs are being tracked correctly
5. Verify prometheus metrics match /user/info spend data
"""
from datetime import datetime, timedelta, timezone
async with aiohttp.ClientSession() as session:
# Setup test user with unique user_id
unique_user_id = f"budget_test_user_{uuid.uuid4()}"
user_data = {
"user_id": unique_user_id,
"max_budget": 10,
"budget_duration": "7d",
"budget_reset_at": (datetime.now(timezone.utc) + timedelta(days=7)).isoformat(),
}
user_info = await create_test_user(session, user_data)
print("user_info", user_info)
user_id = user_info["user_id"]
print("user_id", user_id)
# Get the key that was created with the user
key = user_info["key"]
# Initialize OpenAI client with the user's key
client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=key)
# Make initial request and check budget
await client.chat.completions.create(
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}],
)
await asyncio.sleep(11) # Wait for metrics to update
# Get metrics after request
metrics_after_first = await get_prometheus_metrics(session)
print("metrics_after_first request", metrics_after_first)
first_budget = extract_user_budget_metrics(metrics_after_first, user_id)
print(f"Budget after 1 request: {first_budget}")
assert (
first_budget["remaining"] is not None
), "remaining budget metric should be present"
assert (
first_budget["total"] is not None
), "total budget metric should be present"
assert (
first_budget["remaining"] < 10.0
), "remaining budget should be less than 10.0 after first request"
assert first_budget["total"] == 10.0, "Total budget metric is incorrect"
print("first_budget['remaining_hours']", first_budget["remaining_hours"])
# The budget reset time is now standardized - for "7d" it resets on Monday at midnight
# So we'll check if it's within a reasonable range (0-7 days depending on current day of week)
assert (
first_budget["remaining_hours"] is not None
), "remaining hours metric should be present"
assert (
0 <= first_budget["remaining_hours"] <= 168
), "Budget remaining hours should be within a reasonable range (0-7 days depending on day of week)"
# Get user info and verify spend matches prometheus metrics
user_info_response = await get_user_info(session, user_id)
print("user_info_response", user_info_response)
_user_info_data = user_info_response["user_info"]
# Calculate spend from prometheus (total - remaining)
user_info_spend = float(_user_info_data["spend"])
user_info_max_budget = float(_user_info_data["max_budget"])
user_info_remaining_budget = user_info_max_budget - user_info_spend
print("\n\n\n###### Final budget metrics ######\n\n\n")
print("user_info_remaining_budget", user_info_remaining_budget)
print("prometheus_remaining_budget", first_budget["remaining"])
print(
"diff between user_info_remaining_budget and prometheus_remaining_budget",
user_info_remaining_budget - first_budget["remaining"],
)
# Verify spends match within a small delta (floating point comparison)
assert (
abs(user_info_remaining_budget - first_budget["remaining"]) <= 0.001
), f"Spend mismatch: Prometheus={user_info_remaining_budget}, User Info={first_budget['remaining']}"
@pytest.mark.asyncio
async def test_user_email_metrics():
"""

View file

@ -359,3 +359,60 @@ async def test_proxy_admin_expired_key_from_cache():
# Clean up - restore original values if needed
pass
@pytest.mark.asyncio
async def test_return_user_api_key_auth_obj_user_spend_and_budget():
"""
Test that _return_user_api_key_auth_obj correctly sets user_spend and user_max_budget
from user_obj attributes.
"""
from datetime import datetime
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import _return_user_api_key_auth_obj
user_obj = type(
"LiteLLM_UserTable",
(),
{
"tpm_limit": 1000,
"rpm_limit": 100,
"user_email": "test@example.com",
"spend": 250.0,
"max_budget": 1000.0,
"user_role": "internal_user",
},
)
api_key = "sk-test-key"
valid_token_dict = {
"user_id": "test-user",
"org_id": "test-org",
}
route = "/chat/completions"
start_time = datetime.now()
mock_service_logger = MagicMock()
mock_service_logger.async_service_success_hook = AsyncMock()
with patch(
"litellm.proxy.auth.user_api_key_auth.user_api_key_service_logger_obj",
new=mock_service_logger,
):
result = await _return_user_api_key_auth_obj(
user_obj=user_obj,
api_key=api_key,
parent_otel_span=None,
valid_token_dict=valid_token_dict,
route=route,
start_time=start_time,
user_role=None,
)
assert isinstance(result, UserAPIKeyAuth)
assert result.user_spend == 250.0
assert result.user_max_budget == 1000.0
assert result.user_tpm_limit == 1000
assert result.user_rpm_limit == 100
assert result.user_email == "test@example.com"

View file

@ -160,6 +160,44 @@ async def test_add_litellm_data_to_request_parses_string_metadata():
assert updated_data["metadata"]["generation_name"] == "gen123"
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_user_spend_and_budget():
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "Hello"}]}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
user_spend=150.0,
user_max_budget=500.0,
)
updated_data = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
metadata = updated_data.get("metadata", {})
assert metadata["user_api_key_user_spend"] == 150.0
assert metadata["user_api_key_user_max_budget"] == 500.0
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_audio_transcription_multipart():
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request