mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
(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:
parent
df7d500d42
commit
5fbbf47581
5 changed files with 247 additions and 3 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue