From 5fbbf47581fb578038c026e4fcb7d44dd8154738 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 14 Jan 2025 20:08:23 -0800 Subject: [PATCH] (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 --- litellm/integrations/prometheus.py | 77 +++++++++++++++- .../management_endpoints/team_endpoints.py | 36 +++++++- litellm/proxy/proxy_config.yaml | 3 + .../test_prometheus_unit_tests.py | 90 ++++++++++++++++++- .../test_key_generate_prisma.py | 44 +++++++++ 5 files changed, 247 insertions(+), 3 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 9ca3c48cc4d..ce57309a0cb 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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], diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index e63a472b613..ea07a90ad7f 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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 diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index f1927b613cd..b2edcffab9e 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -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"] diff --git a/tests/logging_callback_tests/test_prometheus_unit_tests.py b/tests/logging_callback_tests/test_prometheus_unit_tests.py index 94b3164a26c..1026368dc69 100644 --- a/tests/logging_callback_tests/test_prometheus_unit_tests.py +++ b/tests/logging_callback_tests/test_prometheus_unit_tests.py @@ -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] diff --git a/tests/proxy_unit_tests/test_key_generate_prisma.py b/tests/proxy_unit_tests/test_key_generate_prisma.py index a7719fe056e..e4382a4543c 100644 --- a/tests/proxy_unit_tests/test_key_generate_prisma.py +++ b/tests/proxy_unit_tests/test_key_generate_prisma.py @@ -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}")