mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #18785 from BerriAI/litellm_user_promethus_metrics
[Feature] User Metrics for Promethus
This commit is contained in:
commit
6a7edd8f2b
9 changed files with 615 additions and 5 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue