(Feat) prometheus - emit remaining team budget metric on proxy startup (#7777)

* fix get_paginated_teams

* use _initialize_remaining_budget_metrics

* fix prom metric

* run ci/cd again

* fix run async func

* fix _initialize_prometheus_startup_metrics

* fix _initialize_prometheus_startup_metrics

* prom unit tests

* test_get_paginated_teams
This commit is contained in:
Ishaan Jaff 2025-01-14 20:08:23 -08:00 • committed by GitHub
parent df7d500d42
commit 5fbbf47581
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 247 additions and 3 deletions

View file

@ -1,9 +1,10 @@
# used for /metrics endpoint on LiteLLM Proxy
#### What this does ####
# On success, log events to Prometheus
import asyncio
import sys
from datetime import datetime, timedelta
from typing import List, Optional, cast
from typing import TYPE_CHECKING, Any, List, Optional, cast
import litellm
from litellm._logging import print_verbose, verbose_logger
@ -13,6 +14,11 @@ from litellm.types.integrations.prometheus import *
from litellm.types.utils import StandardLoggingPayload
from litellm.utils import get_end_user_id_for_cost_tracking
if TYPE_CHECKING:
from litellm.proxy._types import LiteLLM_TeamTable
else:
LiteLLM_TeamTable = Any
class PrometheusLogger(CustomLogger):
# Class variables or attributes
@ -306,6 +312,8 @@ class PrometheusLogger(CustomLogger):
),
)
self._initialize_prometheus_startup_metrics()
except Exception as e:
print_verbose(f"Got exception on init prometheus client {str(e)}")
raise e
@ -1243,6 +1251,73 @@ class PrometheusLogger(CustomLogger):
return max_budget - spend
def _initialize_prometheus_startup_metrics(self):
"""
Initialize prometheus startup metrics
Helper to create tasks for initializing metrics that are required on startup - eg. remaining budget metrics
"""
try:
if asyncio.get_running_loop():
asyncio.create_task(self._initialize_remaining_budget_metrics())
except RuntimeError as e: # no running event loop
verbose_logger.exception(
f"No running event loop - skipping budget metrics initialization: {str(e)}"
)
async def _initialize_remaining_budget_metrics(self):
"""
Initialize remaining budget metrics for all teams to avoid metric discrepancies.
Runs when prometheus logger starts up.
"""
from litellm.proxy.management_endpoints.team_endpoints import (
get_paginated_teams,
)
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return
try:
page = 1
page_size = 50
teams, total_count = await get_paginated_teams(
prisma_client=prisma_client, page_size=page_size, page=page
)
# Calculate total pages needed
total_pages = (total_count + page_size - 1) // page_size
# Set metrics for first page of teams
await self._set_team_budget_metrics(teams)
# Get and set metrics for remaining pages
for page in range(2, total_pages + 1):
teams, _ = await get_paginated_teams(
prisma_client=prisma_client, page_size=page_size, page=page
)
await self._set_team_budget_metrics(teams)
except Exception as e:
verbose_logger.exception(
f"Error initializing team budget metrics: {str(e)}"
)
async def _set_team_budget_metrics(self, teams: List[LiteLLM_TeamTable]):
"""Helper function to set budget metrics for a list of teams"""
for team in teams:
if team.max_budget is not None:
self.litellm_remaining_team_budget_metric.labels(
team.team_id,
team.team_alias or "",
).set(
self._safe_get_remaining_budget(
max_budget=team.max_budget,
spend=team.spend,
)
)
def prometheus_label_factory(
supported_enum_labels: List[str],

View file

@ -14,7 +14,7 @@ import json
import traceback
import uuid
from datetime import datetime, timedelta, timezone
from typing import List, Optional, Union, cast
from typing import List, Optional, Tuple, Union, cast
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
@ -1465,3 +1465,37 @@ async def list_team(
# Sort the responses by team_alias
returned_responses.sort(key=lambda x: (getattr(x, "team_alias", "") or ""))
return returned_responses
async def get_paginated_teams(
prisma_client: PrismaClient,
page_size: int = 10,
page: int = 1,
) -> Tuple[List[LiteLLM_TeamTable], int]:
"""
Get paginated list of teams from team table
Parameters:
prisma_client: PrismaClient - The database client
page_size: int - Number of teams per page
page: int - Page number (1-based)
Returns:
Tuple[List[LiteLLM_TeamTable], int] - (list of teams, total count)
"""
try:
# Calculate skip for pagination
skip = (page - 1) * page_size
# Get total count
total_count = await prisma_client.db.litellm_teamtable.count()
# Get paginated teams
teams = await prisma_client.db.litellm_teamtable.find_many(
skip=skip, take=page_size, order={"team_alias": "asc"} # Sort by team_alias
)
return teams, total_count
except Exception as e:
verbose_proxy_logger.exception(
f"[Non-Blocking] Error getting paginated teams: {e}"
)
return [], 0

View file

@ -11,3 +11,6 @@ model_list:
api_key: os.environ/ANTHROPIC_API_KEY
model_info:
health_check_model: anthropic/claude-3-5-sonnet-20240620
litellm_settings:
callbacks: ["prometheus"]

View file

@ -27,7 +27,7 @@ from litellm.types.utils import (
StandardLoggingModelInformation,
)
import pytest
from unittest.mock import MagicMock, patch
from unittest.mock import MagicMock, patch, call
from datetime import datetime, timedelta
from litellm.integrations.prometheus import PrometheusLogger
from litellm.proxy._types import UserAPIKeyAuth
@ -902,3 +902,91 @@ def test_get_custom_labels_from_metadata(monkeypatch):
"metadata_foo": "bar",
"metadata_bar": "baz",
}
@pytest.mark.asyncio(scope="session")
async def test_initialize_remaining_budget_metrics(prometheus_logger):
"""
Test that _initialize_remaining_budget_metrics correctly sets budget metrics for all teams
"""
# Mock the prisma client and get_paginated_teams function
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams:
# Create mock team data
mock_teams = [
MagicMock(team_id="team1", team_alias="alias1", max_budget=100, spend=30),
MagicMock(team_id="team2", team_alias="alias2", max_budget=200, spend=50),
MagicMock(team_id="team3", team_alias=None, max_budget=300, spend=100),
]
# Mock get_paginated_teams to return our test data
mock_get_teams.return_value = (mock_teams, len(mock_teams))
# Mock the Prometheus metric
prometheus_logger.litellm_remaining_team_budget_metric = MagicMock()
# Call the function
await prometheus_logger._initialize_remaining_budget_metrics()
# Verify the metric was set correctly for each team
expected_calls = [
call.labels("team1", "alias1").set(70), # 100 - 30
call.labels("team2", "alias2").set(150), # 200 - 50
call.labels("team3", "").set(200), # 300 - 100
]
prometheus_logger.litellm_remaining_team_budget_metric.assert_has_calls(
expected_calls, any_order=True
)
@pytest.mark.asyncio
async def test_initialize_remaining_budget_metrics_exception_handling(
prometheus_logger,
):
"""
Test that _initialize_remaining_budget_metrics properly handles exceptions
"""
# Mock the prisma client and get_paginated_teams function to raise an exception
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch(
"litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams"
) as mock_get_teams:
# Make get_paginated_teams raise an exception
mock_get_teams.side_effect = Exception("Database error")
# Mock the Prometheus metric
prometheus_logger.litellm_remaining_team_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 the error was logged
mock_logger.assert_called_once()
assert (
"Error initializing team budget metrics" in mock_logger.call_args[0][0]
)
# Verify the metric was never called
prometheus_logger.litellm_remaining_team_budget_metric.assert_not_called()
def test_initialize_prometheus_startup_metrics_no_loop(prometheus_logger):
"""
Test that _initialize_prometheus_startup_metrics handles case when no event loop exists
"""
# Mock asyncio.get_running_loop to raise RuntimeError
with patch(
"asyncio.get_running_loop", side_effect=RuntimeError("No running event loop")
), patch("litellm._logging.verbose_logger.exception") as mock_logger:
# Call the function
prometheus_logger._initialize_prometheus_startup_metrics()
# Verify the error was logged
mock_logger.assert_called_once()
assert "No running event loop" in mock_logger.call_args[0][0]

View file

@ -3759,3 +3759,47 @@ def test_should_track_cost_callback():
team_id=None,
end_user_id="1234",
)
@pytest.mark.asyncio
async def test_get_paginated_teams(prisma_client):
"""
Test the get_paginated_teams function:
1. Test pagination returns valid results
2. Test total count matches across pages
3. Test page size is respected
"""
from litellm.proxy.management_endpoints.team_endpoints import get_paginated_teams
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
await litellm.proxy.proxy_server.prisma_client.connect()
try:
# Get first page with page_size=2
teams_page_1, total_count_1 = await get_paginated_teams(
prisma_client=prisma_client, page_size=2, page=1
)
print("teams_page_1=", teams_page_1)
print("total_count_1=", total_count_1)
# Get second page
teams_page_2, total_count_2 = await get_paginated_teams(
prisma_client=prisma_client, page_size=2, page=2
)
print("teams_page_2=", teams_page_2)
print("total_count_2=", total_count_2)
# Verify results
assert isinstance(teams_page_1, list) # Should return a list
assert isinstance(total_count_1, int) # Should return an integer count
assert (
total_count_1 == total_count_2
) # Total count should be consistent across pages
assert len(teams_page_1) <= 2 # Should respect page_size limit
except Exception as e:
print(f"Error occurred: {e}")
pytest.fail(f"Test failed with exception: {e}")