From 67833590d695a630c2f70ca84ce59c6515367dc3 Mon Sep 17 00:00:00 2001 From: "Jugal D. Bhatt" <55304795+jugaldb@users.noreply.github.com> Date: Sat, 9 Aug 2025 16:12:13 -0700 Subject: [PATCH] [Proxy changes] Litellm add model price reload schedule for multi-pod (#13470) * added mcp guardrails doc in mcp.md * add button to reload models * Added button changes * added button for scheduling reload * add multi pod support to reloading the model price json * fix ruff --- litellm/proxy/proxy_server.py | 350 +++++++++++++++- tests/test_litellm/proxy/test_proxy_server.py | 312 +++++++++++++- ui/litellm-dashboard/package-lock.json | 120 ++++++ .../src/components/model_dashboard.tsx | 48 ++- .../src/components/networking.tsx | 72 ++++ .../src/components/price_data_reload.tsx | 385 ++++++++++++++---- 6 files changed, 1171 insertions(+), 116 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6b33e92c48a..df6f734b987 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11,6 +11,7 @@ import time import traceback import uuid import warnings +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from datetime import datetime, timedelta from typing import ( TYPE_CHECKING, @@ -981,6 +982,11 @@ proxy_logging_obj = ProxyLogging( async_result = None celery_app_conn = None celery_fn = None # Redis Queue for handling requests + +# Global variables for model cost map reload scheduling +scheduler = None +last_model_cost_map_reload = None + ### DB WRITER ### db_writer_client: Optional[AsyncHTTPHandler] = None ### logger ### @@ -2939,6 +2945,91 @@ class ProxyConfig: await self._init_mcp_servers_in_db() await self._init_pass_through_endpoints_in_db() await self._init_prompts_in_db(prisma_client=prisma_client) + await self._check_and_reload_model_cost_map(prisma_client=prisma_client) + + async def _check_and_reload_model_cost_map(self, prisma_client: PrismaClient): + """ + Check if model cost map needs to be reloaded based on database configuration. + This function runs every 10 seconds as part of _init_non_llm_objects_in_db. + """ + try: + # Get model cost map reload configuration from database + config_record = await prisma_client.db.litellm_config.find_unique( + where={"param_name": "model_cost_map_reload_config"} + ) + + if config_record is None or config_record.param_value is None: + return # No configuration found, skip reload + + config = config_record.param_value + interval_hours = config.get("interval_hours") + force_reload = config.get("force_reload", False) + + if interval_hours is None and force_reload is False: + return # No interval configured, skip reload + + current_time = datetime.utcnow() + + # Check if we need to reload based on interval or force reload + should_reload = False + + if force_reload: + should_reload = True + verbose_proxy_logger.info("Model cost map reload triggered by force reload flag") + elif interval_hours is not None: + # Use pod's in-memory last reload time + global last_model_cost_map_reload + if last_model_cost_map_reload is not None: + try: + last_reload_time = datetime.fromisoformat(last_model_cost_map_reload) + time_since_last_reload = current_time - last_reload_time + hours_since_last_reload = time_since_last_reload.total_seconds() / 3600 + + if hours_since_last_reload >= interval_hours: + should_reload = True + verbose_proxy_logger.info(f"Model cost map reload triggered by interval. Hours since last reload: {hours_since_last_reload:.2f}, Interval: {interval_hours}") + except Exception as e: + verbose_proxy_logger.warning(f"Error parsing last reload time: {e}") + # If we can't parse the last reload time, reload anyway + should_reload = True + else: + # No last reload time recorded, reload now + should_reload = True + verbose_proxy_logger.info("Model cost map reload triggered - no previous reload time recorded") + + if should_reload: + # Perform the reload + from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map + model_cost_map_url = litellm.model_cost_map_url + new_model_cost_map = get_model_cost_map(url=model_cost_map_url) + litellm.model_cost = new_model_cost_map + + # Update pod's in-memory last reload time + last_model_cost_map_reload = current_time.isoformat() + + # Clear force reload flag in database + await prisma_client.db.litellm_config.upsert( + where={"param_name": "model_cost_map_reload_config"}, + data={ + "create": { + "param_name": "model_cost_map_reload_config", + "param_value": safe_dumps({ + "interval_hours": interval_hours, + "force_reload": False + }) + }, + "update": { + "param_value": safe_dumps({ + "force_reload": False + }) + } + } + ) + + verbose_proxy_logger.info(f"Model cost map reloaded successfully. Models count: {len(new_model_cost_map) if new_model_cost_map else 0}") + + except Exception as e: + verbose_proxy_logger.exception(f"Error in _check_and_reload_model_cost_map: {str(e)}") async def _init_prompts_in_db(self, prisma_client: PrismaClient): from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY @@ -3537,6 +3628,7 @@ class ProxyStartupEvent: args=[prisma_client, db_writer_client, proxy_logging_obj], ) + ### ADD NEW MODELS ### store_model_in_db = ( get_secret_bool("STORE_MODEL_IN_DB", store_model_in_db) or store_model_in_db @@ -8919,24 +9011,51 @@ async def reload_model_cost_map( ) try: + global prisma_client + if prisma_client is None: + raise HTTPException( + status_code=500, + detail="Database connection not available" + ) + + # Immediately reload the model cost map in the current pod from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map - - # Get the current URL from litellm configuration model_cost_map_url = litellm.model_cost_map_url - - # Reload the model cost map new_model_cost_map = get_model_cost_map(url=model_cost_map_url) - - # Update the global model_cost variable litellm.model_cost = new_model_cost_map - verbose_proxy_logger.info("Model cost map reloaded successfully") + # Update pod's in-memory last reload time + global last_model_cost_map_reload + current_time = datetime.utcnow() + last_model_cost_map_reload = current_time.isoformat() + + # Set force reload flag in database for other pods + await prisma_client.db.litellm_config.upsert( + where={"param_name": "model_cost_map_reload_config"}, + data={ + "create": { + "param_name": "model_cost_map_reload_config", + "param_value": safe_dumps({ + "interval_hours": None, + "force_reload": True + }) + }, + "update": { + "param_value": safe_dumps({ + "force_reload": True + }) + } + } + ) + + models_count = len(new_model_cost_map) if new_model_cost_map else 0 + verbose_proxy_logger.info(f"Model cost map reloaded successfully in current pod. Models count: {models_count}") return { - "message": "Model cost map reloaded successfully", + "message": f"Price data reloaded successfully! {models_count} models updated.", "status": "success", - "timestamp": datetime.utcnow().isoformat(), - "models_count": len(new_model_cost_map) if new_model_cost_map else 0 + "models_count": models_count, + "timestamp": current_time.isoformat() } except Exception as e: verbose_proxy_logger.exception(f"Failed to reload model cost map: {str(e)}") @@ -8946,6 +9065,217 @@ async def reload_model_cost_map( ) +@router.post( + "/schedule/model_cost_map_reload", + tags=["model management"], + dependencies=[Depends(user_api_key_auth)], + include_in_schema=False, +) +async def schedule_model_cost_map_reload( + hours: int, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + ADMIN ONLY / MASTER KEY Only Endpoint + + Schedule periodic reload of the model cost map. + This will create a background job that reloads the model cost map every specified hours. + """ + # Check if user is admin + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", + ) + + if hours <= 0: + raise HTTPException( + status_code=400, + detail="Hours must be greater than 0" + ) + + try: + global prisma_client + if prisma_client is None: + raise HTTPException( + status_code=500, + detail="Database connection not available" + ) + + # Update database with new reload configuration + await prisma_client.db.litellm_config.upsert( + where={"param_name": "model_cost_map_reload_config"}, + data={ + "create": { + "param_name": "model_cost_map_reload_config", + "param_value": safe_dumps({ + "interval_hours": hours, + "force_reload": False + }) + }, + "update": { + "param_value": safe_dumps({ + "interval_hours": hours, + "force_reload": False + }) + } + } + ) + + verbose_proxy_logger.info(f"Model cost map reload scheduled for every {hours} hours") + + return { + "message": f"Model cost map reload scheduled for every {hours} hours", + "status": "success", + "interval_hours": hours, + "timestamp": datetime.utcnow().isoformat() + } + except Exception as e: + verbose_proxy_logger.exception(f"Failed to schedule model cost map reload: {str(e)}") + raise HTTPException( + status_code=500, + detail=f"Failed to schedule model cost map reload: {str(e)}" + ) + + +@router.delete( + "/schedule/model_cost_map_reload", + tags=["model management"], + dependencies=[Depends(user_api_key_auth)], + include_in_schema=False, +) +async def cancel_model_cost_map_reload( + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + ADMIN ONLY / MASTER KEY Only Endpoint + + Cancel the scheduled periodic reload of the model cost map. + """ + # Check if user is admin + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", + ) + + try: + global prisma_client + if prisma_client is None: + raise HTTPException( + status_code=500, + detail="Database connection not available" + ) + + # Remove reload configuration from database + await prisma_client.db.litellm_config.delete( + where={"param_name": "model_cost_map_reload_config"} + ) + + verbose_proxy_logger.info("Model cost map reload schedule cancelled") + + return { + "message": "Model cost map reload schedule cancelled", + "status": "success", + "timestamp": datetime.utcnow().isoformat() + } + except Exception as e: + verbose_proxy_logger.exception(f"Failed to cancel model cost map reload: {str(e)}") + raise HTTPException( + status_code=500, + detail=f"Failed to cancel model cost map reload: {str(e)}" + ) + + +@router.get( + "/schedule/model_cost_map_reload/status", + tags=["model management"], + dependencies=[Depends(user_api_key_auth)], + include_in_schema=False, +) +async def get_model_cost_map_reload_status( + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + ADMIN ONLY / MASTER KEY Only Endpoint + + Get the status of the scheduled model cost map reload job. + """ + # Check if user is admin + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", + ) + + try: + global prisma_client, last_model_cost_map_reload + + verbose_proxy_logger.info(f"Checking model cost map reload status. Last reload: {last_model_cost_map_reload}") + + if prisma_client is None: + verbose_proxy_logger.info("No database connection, returning not scheduled") + return { + "scheduled": False, + "interval_hours": None, + "last_run": None, + "next_run": None + } + + # Get reload configuration from database + config_record = await prisma_client.db.litellm_config.find_unique( + where={"param_name": "model_cost_map_reload_config"} + ) + + if config_record is None or config_record.param_value is None: + verbose_proxy_logger.info("No model cost map reload configuration found") + return { + "scheduled": False, + "interval_hours": None, + "last_run": None, + "next_run": None + } + + config = config_record.param_value + interval_hours = config.get("interval_hours") + + if interval_hours is None: + verbose_proxy_logger.info("No interval configured, returning not scheduled") + return { + "scheduled": False, + "interval_hours": None, + "last_run": None, + "next_run": None + } + + current_time = datetime.utcnow() + next_run = None + + # Use pod's in-memory last reload time + if last_model_cost_map_reload is not None: + try: + last_reload_time = datetime.fromisoformat(last_model_cost_map_reload) + time_since_last_reload = current_time - last_reload_time + hours_since_last_reload = time_since_last_reload.total_seconds() / 3600 + + if hours_since_last_reload < interval_hours: + next_run = (last_reload_time + timedelta(hours=interval_hours)).isoformat() + except Exception as e: + verbose_proxy_logger.warning(f"Error parsing last reload time: {e}") + + return { + "scheduled": True, + "interval_hours": interval_hours, + "last_run": last_model_cost_map_reload, + "next_run": next_run + } + except Exception as e: + verbose_proxy_logger.exception(f"Failed to get model cost map reload status: {str(e)}") + raise HTTPException( + status_code=500, + detail=f"Failed to get model cost map reload status: {str(e)}" + ) + @router.get("/", dependencies=[Depends(user_api_key_auth)]) async def home(request: Request): return "LiteLLM: RUNNING" diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index feb3f154f3a..c774602e5ca 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5,6 +5,7 @@ import os import socket import subprocess import sys +from datetime import datetime from unittest import mock from unittest.mock import AsyncMock, MagicMock, mock_open, patch @@ -1222,15 +1223,21 @@ class TestPriceDataReloadAPI: """Test that admin users can access the reload endpoint""" with patch('litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map') as mock_get_map: mock_get_map.return_value = {"gpt-3.5-turbo": {"input_cost_per_token": 0.001}} - - response = client_with_auth.post("/reload/model_cost_map") - - assert response.status_code == 200 - data = response.json() - assert data["status"] == "success" - assert "message" in data - assert "timestamp" in data - assert "models_count" in data + # Mock the database connection + with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) + + response = client_with_auth.post("/reload/model_cost_map") + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "success" + assert "message" in data + assert "timestamp" in data + assert "models_count" in data + # The new implementation immediately reloads and returns the count + assert "Price data reloaded successfully! 1 models updated." in data["message"] + assert data["models_count"] == 1 def test_reload_model_cost_map_non_admin_access(self, client_with_auth): """Test that non-admin users cannot access the reload endpoint""" @@ -1274,11 +1281,155 @@ class TestPriceDataReloadAPI: with patch('litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map') as mock_get_map: mock_get_map.side_effect = Exception("Network error") - response = client_with_auth.post("/reload/model_cost_map") - - assert response.status_code == 500 + # Mock the database connection + with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) + + response = client_with_auth.post("/reload/model_cost_map") + + assert response.status_code == 500 # The new implementation immediately reloads and fails on error + data = response.json() + assert "Failed to reload model cost map" in data["detail"] + + def test_schedule_model_cost_map_reload_admin_access(self, client_with_auth): + """Test that admin users can schedule periodic reload""" + with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + # Mock database upsert + mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) + + response = client_with_auth.post("/schedule/model_cost_map_reload?hours=6") + + assert response.status_code == 200 data = response.json() - assert "Failed to reload model cost map" in data["detail"] + assert data["status"] == "success" + assert data["interval_hours"] == 6 + assert "message" in data + assert "timestamp" in data + + def test_schedule_model_cost_map_reload_non_admin_access(self, client_with_auth): + """Test that non-admin users cannot schedule periodic reload""" + # Mock non-admin user + mock_auth = MagicMock() + mock_auth.user_role = "user" # Non-admin role + app.dependency_overrides[user_api_key_auth] = lambda: mock_auth + + response = client_with_auth.post("/schedule/model_cost_map_reload?hours=6") + + assert response.status_code == 403 + data = response.json() + assert "Access denied" in data["detail"] + assert "Admin role required" in data["detail"] + + def test_schedule_model_cost_map_reload_invalid_hours(self, client_with_auth): + """Test that invalid hours parameter is rejected""" + response = client_with_auth.post("/schedule/model_cost_map_reload?hours=0") + + assert response.status_code == 400 + data = response.json() + assert "Hours must be greater than 0" in data["detail"] + + def test_cancel_model_cost_map_reload_admin_access(self, client_with_auth): + """Test that admin users can cancel periodic reload""" + with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + # Mock database delete + mock_prisma.db.litellm_config.delete = AsyncMock(return_value=None) + + response = client_with_auth.delete("/schedule/model_cost_map_reload") + + assert response.status_code == 200 + data = response.json() + assert data["status"] == "success" + assert "message" in data + assert "timestamp" in data + + def test_cancel_model_cost_map_reload_non_admin_access(self, client_with_auth): + """Test that non-admin users cannot cancel periodic reload""" + # Mock non-admin user + mock_auth = MagicMock() + mock_auth.user_role = "user" # Non-admin role + app.dependency_overrides[user_api_key_auth] = lambda: mock_auth + + response = client_with_auth.delete("/schedule/model_cost_map_reload") + + assert response.status_code == 403 + data = response.json() + assert "Access denied" in data["detail"] + assert "Admin role required" in data["detail"] + + def test_get_model_cost_map_reload_status_admin_access(self, client_with_auth): + """Test that admin users can get reload status""" + with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + # Mock database config record + mock_config = MagicMock() + mock_config.param_value = { + "interval_hours": 6, + "force_reload": False + } + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_config) + + # Mock the last reload time and current time + with patch('litellm.proxy.proxy_server.last_model_cost_map_reload', "2024-01-01T06:00:00"): + with patch('litellm.proxy.proxy_server.datetime') as mock_datetime: + # Mock current time to be 1 hour after last reload + mock_datetime.utcnow.return_value = datetime(2024, 1, 1, 7, 0, 0) + mock_datetime.fromisoformat = datetime.fromisoformat + + response = client_with_auth.get("/schedule/model_cost_map_reload/status") + + assert response.status_code == 200 + data = response.json() + assert data["scheduled"] == True + assert data["interval_hours"] == 6 + assert data["last_run"] == "2024-01-01T06:00:00" + assert data["next_run"] == "2024-01-01T12:00:00" + + def test_get_model_cost_map_reload_status_non_admin_access(self, client_with_auth): + """Test that non-admin users cannot get reload status""" + # Mock non-admin user + mock_auth = MagicMock() + mock_auth.user_role = "user" # Non-admin role + app.dependency_overrides[user_api_key_auth] = lambda: mock_auth + + response = client_with_auth.get("/schedule/model_cost_map_reload/status") + + assert response.status_code == 403 + data = response.json() + assert "Access denied" in data["detail"] + assert "Admin role required" in data["detail"] + + def test_get_model_cost_map_reload_status_no_config(self, client_with_auth): + """Test that status returns not scheduled when no config exists""" + with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + + response = client_with_auth.get("/schedule/model_cost_map_reload/status") + + assert response.status_code == 200 + data = response.json() + assert data["scheduled"] == False + assert data["interval_hours"] == None + assert data["last_run"] == None + assert data["next_run"] == None + + def test_get_model_cost_map_reload_status_no_interval(self, client_with_auth): + """Test that status returns not scheduled when no interval is configured""" + with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + # Mock config with no interval + mock_config = MagicMock() + mock_config.param_value = { + "interval_hours": None, + "force_reload": False + } + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_config) + + response = client_with_auth.get("/schedule/model_cost_map_reload/status") + + assert response.status_code == 200 + data = response.json() + assert data["scheduled"] == False + assert data["interval_hours"] == None + assert data["last_run"] == None + assert data["next_run"] == None class TestPriceDataReloadIntegration: @@ -1319,13 +1470,70 @@ class TestPriceDataReloadIntegration: with patch('litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map') as mock_get_map: mock_get_map.return_value = mock_cost_map - # Test reload endpoint - response = client_with_auth.post("/reload/model_cost_map") - assert response.status_code == 200 + # Mock the database connection + with patch('litellm.proxy.proxy_server.prisma_client') as mock_prisma: + mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) + + # Test reload endpoint + response = client_with_auth.post("/reload/model_cost_map") + assert response.status_code == 200 + + # Test get endpoint + response = client_with_auth.get("/get/litellm_model_cost_map") + assert response.status_code == 200 + + def test_distributed_reload_check_function(self): + """Test the _check_and_reload_model_cost_map function""" + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + + # Mock prisma client + mock_prisma = MagicMock() + + # Test case 1: No config in database + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) + + # Should return early without reloading + asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) + + # Test case 2: Config with interval but not time to reload + mock_config = MagicMock() + mock_config.param_value = { + "interval_hours": 6, + "force_reload": False + } + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_config) + + # Mock current time and last reload time + with patch('litellm.proxy.proxy_server.last_model_cost_map_reload', "2024-01-01T06:00:00"): + with patch('litellm.proxy.proxy_server.datetime') as mock_datetime: + mock_datetime.utcnow.return_value = datetime(2024, 1, 1, 7, 0, 0) # 1 hour later + + # Should not reload (only 1 hour passed, need 6) + asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) + + # Test case 3: Config with force reload + mock_config.param_value = { + "interval_hours": 6, + "force_reload": True + } + mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_config) + mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) + + with patch('litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map') as mock_get_map: + mock_get_map.return_value = {"gpt-3.5-turbo": {"input_cost_per_token": 0.001}} - # Test get endpoint - response = client_with_auth.get("/get/litellm_model_cost_map") - assert response.status_code == 200 + # Should reload due to force flag + asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma)) + + # Verify force_reload was reset to False + mock_prisma.db.litellm_config.upsert.assert_called() + call_args = mock_prisma.db.litellm_config.upsert.call_args + # The param_value is now a JSON string, so we need to parse it + param_value_json = call_args[1]['data']['update']['param_value'] + param_value_dict = json.loads(param_value_json) + assert param_value_dict['force_reload'] == False def test_config_file_parsing(self): """Test parsing of config file with reload settings""" @@ -1354,3 +1562,69 @@ model_list: # Verify models are present assert "model_list" in config assert len(config["model_list"]) == 2 + def test_database_config_storage(self): + """Test that configuration is properly stored in database""" + # Mock prisma client + mock_prisma = MagicMock() + + # Test the database upsert call that would be made by the schedule endpoint + mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) + + # Simulate the database call that the schedule endpoint would make + asyncio.run(mock_prisma.db.litellm_config.upsert( + where={"param_name": "model_cost_map_reload_config"}, + data={ + "create": { + "param_name": "model_cost_map_reload_config", + "param_value": { + "interval_hours": 6, + "force_reload": False + } + }, + "update": { + "param_value": { + "interval_hours": 6, + "force_reload": False + } + } + } + )) + + # Verify database upsert was called with correct data + mock_prisma.db.litellm_config.upsert.assert_called_once() + call_args = mock_prisma.db.litellm_config.upsert.call_args + assert call_args[1]['where']['param_name'] == "model_cost_map_reload_config" + assert call_args[1]['data']['create']['param_value']['interval_hours'] == 6 + assert call_args[1]['data']['create']['param_value']['force_reload'] == False + + def test_manual_reload_force_flag(self): + """Test that manual reload sets force flag correctly""" + # Mock prisma client + mock_prisma = MagicMock() + + # Test the database upsert call that would be made by the manual reload endpoint + mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) + + # Simulate the database call that the manual reload endpoint would make + asyncio.run(mock_prisma.db.litellm_config.upsert( + where={"param_name": "model_cost_map_reload_config"}, + data={ + "create": { + "param_name": "model_cost_map_reload_config", + "param_value": { + "interval_hours": None, + "force_reload": True + } + }, + "update": { + "param_value": { + "force_reload": True + } + } + } + )) + + # Verify force_reload flag was set + mock_prisma.db.litellm_config.upsert.assert_called_once() + call_args = mock_prisma.db.litellm_config.upsert.call_args + assert call_args[1]['data']['update']['param_value']['force_reload'] == True diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index db4a360babb..a80d65485e9 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -20484,6 +20484,126 @@ "type": "github", "url": "https://github.com/sponsors/wooorm" } + }, + "node_modules/@next/swc-darwin-x64": { + "version": "14.2.30", + "resolved": "https://registry.npmjs.org/@next/swc-darwin-x64/-/swc-darwin-x64-14.2.30.tgz", + "integrity": "sha512-TyO7Wz1IKE2kGv8dwQ0bmPL3s44EKVencOqwIY69myoS3rdpO1NPg5xPM5ymKu7nfX4oYJrpMxv8G9iqLsnL4A==", + "cpu": [ + "x64" + ], + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@next/swc-linux-arm64-gnu": { + "version": "14.2.30", + "resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-gnu/-/swc-linux-arm64-gnu-14.2.30.tgz", + "integrity": "sha512-I5lg1fgPJ7I5dk6mr3qCH1hJYKJu1FsfKSiTKoYwcuUf53HWTrEkwmMI0t5ojFKeA6Vu+SfT2zVy5NS0QLXV4Q==", + "cpu": [ + "arm64" + ], + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@next/swc-linux-arm64-musl": { + "version": "14.2.30", + "resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-musl/-/swc-linux-arm64-musl-14.2.30.tgz", + "integrity": "sha512-8GkNA+sLclQyxgzCDs2/2GSwBc92QLMrmYAmoP2xehe5MUKBLB2cgo34Yu242L1siSkwQkiV4YLdCnjwc/Micw==", + "cpu": [ + "arm64" + ], + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@next/swc-linux-x64-gnu": { + "version": "14.2.30", + "resolved": "https://registry.npmjs.org/@next/swc-linux-x64-gnu/-/swc-linux-x64-gnu-14.2.30.tgz", + "integrity": "sha512-8Ly7okjssLuBoe8qaRCcjGtcMsv79hwzn/63wNeIkzJVFVX06h5S737XNr7DZwlsbTBDOyI6qbL2BJB5n6TV/w==", + "cpu": [ + "x64" + ], + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@next/swc-linux-x64-musl": { + "version": "14.2.30", + "resolved": "https://registry.npmjs.org/@next/swc-linux-x64-musl/-/swc-linux-x64-musl-14.2.30.tgz", + "integrity": "sha512-dBmV1lLNeX4mR7uI7KNVHsGQU+OgTG5RGFPi3tBJpsKPvOPtg9poyav/BYWrB3GPQL4dW5YGGgalwZ79WukbKQ==", + "cpu": [ + "x64" + ], + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@next/swc-win32-arm64-msvc": { + "version": "14.2.30", + "resolved": "https://registry.npmjs.org/@next/swc-win32-arm64-msvc/-/swc-win32-arm64-msvc-14.2.30.tgz", + "integrity": "sha512-6MMHi2Qc1Gkq+4YLXAgbYslE1f9zMGBikKMdmQRHXjkGPot1JY3n5/Qrbg40Uvbi8//wYnydPnyvNhI1DMUW1g==", + "cpu": [ + "arm64" + ], + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@next/swc-win32-ia32-msvc": { + "version": "14.2.30", + "resolved": "https://registry.npmjs.org/@next/swc-win32-ia32-msvc/-/swc-win32-ia32-msvc-14.2.30.tgz", + "integrity": "sha512-pVZMnFok5qEX4RT59mK2hEVtJX+XFfak+/rjHpyFh7juiT52r177bfFKhnlafm0UOSldhXjj32b+LZIOdswGTg==", + "cpu": [ + "ia32" + ], + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@next/swc-win32-x64-msvc": { + "version": "14.2.30", + "resolved": "https://registry.npmjs.org/@next/swc-win32-x64-msvc/-/swc-win32-x64-msvc-14.2.30.tgz", + "integrity": "sha512-4KCo8hMZXMjpTzs3HOqOGYYwAXymXIy7PEPAXNEcEOyKqkjiDlECumrWziy+JEF0Oi4ILHGxzgQ3YiMGG2t/Lg==", + "cpu": [ + "x64" + ], + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">= 10" + } } } } diff --git a/ui/litellm-dashboard/src/components/model_dashboard.tsx b/ui/litellm-dashboard/src/components/model_dashboard.tsx index a4c11560fa7..cd4fd66e583 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard.tsx @@ -1045,30 +1045,14 @@ const ModelDashboard: React.FC = ({
- {/* Price Data Reload Section */} + {/* Model Management Header */}

Model Management

- Manage your models and pricing data + Manage your models and configurations

- {all_admin_roles.includes(userRole) && ( - { - // Refresh the model map after successful reload - const fetchModelMap = async () => { - const data = await modelCostMap(accessToken); - setModelMap(data); - }; - fetchModelMap(); - }} - buttonText="Reload Price Data" - size="small" - type="primary" - /> - )}
{selectedModelId ? ( = ({ {all_admin_roles.includes(userRole) && ( Model Group Alias )} + {all_admin_roles.includes(userRole) && ( + Price Data Reload + )}
@@ -1905,6 +1892,31 @@ const ModelDashboard: React.FC = ({ onAliasUpdate={setModelGroupAlias} /> + +
+
+ Price Data Management + + Manage model pricing data and configure automatic reload schedules + +
+ { + // Refresh the model map after successful reload + const fetchModelMap = async () => { + const data = await modelCostMap(accessToken); + setModelMap(data); + }; + fetchModelMap(); + }} + buttonText="Reload Price Data" + size="middle" + type="primary" + className="w-full" + /> +
+
)} diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 8097f9c4976..12a3faa8ad0 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -278,6 +278,78 @@ export const reloadModelCostMap = async (accessToken: string) => { throw error; } }; + +export const scheduleModelCostMapReload = async (accessToken: string, hours: number) => { + try { + const url = proxyBaseUrl + ? `${proxyBaseUrl}/schedule/model_cost_map_reload?hours=${hours}` + : `/schedule/model_cost_map_reload?hours=${hours}`; + const response = await fetch(url, { + method: "POST", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + const jsonData = await response.json(); + console.log(`Schedule model cost map reload response: ${jsonData}`); + return jsonData; + } catch (error) { + console.error("Failed to schedule model cost map reload:", error); + throw error; + } +}; + +export const cancelModelCostMapReload = async (accessToken: string) => { + try { + const url = proxyBaseUrl + ? `${proxyBaseUrl}/schedule/model_cost_map_reload` + : `/schedule/model_cost_map_reload`; + const response = await fetch(url, { + method: "DELETE", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + const jsonData = await response.json(); + console.log(`Cancel model cost map reload response: ${jsonData}`); + return jsonData; + } catch (error) { + console.error("Failed to cancel model cost map reload:", error); + throw error; + } +}; + +export const getModelCostMapReloadStatus = async (accessToken: string) => { + try { + const url = proxyBaseUrl + ? `${proxyBaseUrl}/schedule/model_cost_map_reload/status` + : `/schedule/model_cost_map_reload/status`; + console.log("Fetching status from URL:", url); + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + console.error(`Status request failed with status: ${response.status}`); + const errorText = await response.text(); + console.error("Error response:", errorText); + throw new Error(`HTTP ${response.status}: ${errorText}`); + } + + const jsonData = await response.json(); + console.log(`Model cost map reload status:`, jsonData); + return jsonData; + } catch (error) { + console.error("Failed to get model cost map reload status:", error); + throw error; + } +}; export const modelCreateCall = async ( accessToken: string, formValues: Model diff --git a/ui/litellm-dashboard/src/components/price_data_reload.tsx b/ui/litellm-dashboard/src/components/price_data_reload.tsx index e31fa7d37a1..05d3a2e102c 100644 --- a/ui/litellm-dashboard/src/components/price_data_reload.tsx +++ b/ui/litellm-dashboard/src/components/price_data_reload.tsx @@ -1,6 +1,17 @@ -import React, { useState } from "react"; -import { Button, message, Popconfirm, Tooltip } from "antd"; -import { reloadModelCostMap } from "./networking"; +import React, { useState, useEffect } from "react"; +import { Button, Popconfirm, message, Modal, InputNumber, Space, Typography, Tag, Card } from "antd"; +import { ReloadOutlined, ClockCircleOutlined, StopOutlined } from "@ant-design/icons"; +import { reloadModelCostMap, scheduleModelCostMapReload, cancelModelCostMapReload, getModelCostMapReloadStatus } from "./networking"; + +const { Text } = Typography; + +interface ReloadStatus { + scheduled: boolean; + interval_hours: number | null; + last_run: string | null; + next_run: string | null; +} + interface PriceDataReloadProps { accessToken: string; @@ -22,22 +33,65 @@ const PriceDataReload: React.FC = ({ className = "", }) => { const [isLoading, setIsLoading] = useState(false); + const [isScheduling, setIsScheduling] = useState(false); + const [isCancelling, setIsCancelling] = useState(false); + const [showScheduleModal, setShowScheduleModal] = useState(false); + const [hours, setHours] = useState(6); + const [reloadStatus, setReloadStatus] = useState(null); + const [loadingStatus, setLoadingStatus] = useState(false); - const handleReload = async () => { + // Fetch status on component mount and periodically + useEffect(() => { + fetchReloadStatus(); + + // Refresh status every 30 seconds to keep it up to date + const interval = setInterval(() => { + fetchReloadStatus(); + }, 30000); + + return () => clearInterval(interval); + }, [accessToken]); + + const fetchReloadStatus = async () => { + if (!accessToken) return; + + setLoadingStatus(true); + try { + console.log("Fetching reload status..."); + const status = await getModelCostMapReloadStatus(accessToken); + console.log("Received status:", status); + setReloadStatus(status); + } catch (error) { + console.error("Failed to fetch reload status:", error); + // Set a default status to prevent UI issues + setReloadStatus({ + scheduled: false, + interval_hours: null, + last_run: null, + next_run: null + }); + } finally { + setLoadingStatus(false); + } + }; + + const handleHardRefresh = async () => { if (!accessToken) { message.error("No access token available"); return; } - + setIsLoading(true); try { const response = await reloadModelCostMap(accessToken); - + if (response.status === "success") { message.success( `Price data reloaded successfully! ${response.models_count || 0} models updated.` ); onReloadSuccess?.(); + // Refresh status after successful reload + await fetchReloadStatus(); } else { message.error("Failed to reload price data"); } @@ -48,73 +102,266 @@ const PriceDataReload: React.FC = ({ setIsLoading(false); } }; + const handleScheduleReload = async () => { + if (!accessToken) { + message.error("No access token available"); + return; + } + + if (hours <= 0) { + message.error("Hours must be greater than 0"); + return; + } + + setIsScheduling(true); + try { + const response = await scheduleModelCostMapReload(accessToken, hours); + + if (response.status === "success") { + message.success(`Periodic reload scheduled for every ${hours} hours`); + setShowScheduleModal(false); + await fetchReloadStatus(); + } else { + message.error("Failed to schedule periodic reload"); + } + } catch (error) { + console.error("Error scheduling reload:", error); + message.error("Failed to schedule periodic reload. Please try again."); + } finally { + setIsScheduling(false); + } + }; + + const handleCancelReload = async () => { + if (!accessToken) { + message.error("No access token available"); + return; + } + + setIsCancelling(true); + try { + const response = await cancelModelCostMapReload(accessToken); + + if (response.status === "success") { + message.success("Periodic reload cancelled successfully"); + await fetchReloadStatus(); + } else { + message.error("Failed to cancel periodic reload"); + } + } catch (error) { + console.error("Error cancelling reload:", error); + message.error("Failed to cancel periodic reload. Please try again."); + } finally { + setIsCancelling(false); + } + }; + + const formatDateTime = (dateTimeString: string | null) => { + if (!dateTimeString) return "Never"; + try { + return new Date(dateTimeString).toLocaleString(); + } catch { + return dateTimeString; + } + }; + + const getStatusText = () => { + if (!reloadStatus?.scheduled) return 'Not scheduled'; + if (!reloadStatus.last_run) return 'Ready'; + return 'Active'; + }; + + const getStatusColor = () => { + if (!reloadStatus?.scheduled) return 'default'; + if (!reloadStatus.last_run) return 'processing'; + return 'success'; + }; return ( - { - e.currentTarget.style.backgroundColor = "#4f46e5"; - e.currentTarget.style.borderColor = "#4f46e5"; - }, - onMouseLeave: (e) => { - e.currentTarget.style.backgroundColor = "#6366f1"; - e.currentTarget.style.borderColor = "#6366f1"; - }, - }} - > - - - - + + + + {/* Periodic Reload Controls */} + {!reloadStatus?.scheduled ? ( + + ) : ( + + )} + + + {/* Status Card */} + {reloadStatus && ( + + + {reloadStatus.scheduled ? ( +
+ }> + Scheduled every {reloadStatus.interval_hours} hours + +
+ ) : ( + No periodic reload scheduled + )} + +
+ Last run: + {formatDateTime(reloadStatus.last_run)} +
+ + {reloadStatus.scheduled && ( + <> + {reloadStatus.next_run && ( +
+ Next run: + {formatDateTime(reloadStatus.next_run)} +
+ )} +
+ Status: + {getStatusText()} +
+ + )} +
+
+ )} + + {/* Schedule Modal */} + setShowScheduleModal(false)} + confirmLoading={isScheduling} + okText="Schedule" + cancelText="Cancel" + okButtonProps={{ + style: { + backgroundColor: "#6366f1", + borderColor: "#6366f1", + color: "white", + }, + }} + > +
+ Set up automatic reload of price data every: +
+
+ setHours(value || 6)} + addonAfter="hours" + style={{ width: '100%' }} + /> +
+
+ + This will automatically fetch the latest pricing data from the remote source every {hours} hours. + +
+
+
); };