mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[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
This commit is contained in:
parent
1c8761111f
commit
67833590d6
6 changed files with 1171 additions and 116 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
120
ui/litellm-dashboard/package-lock.json
generated
120
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1045,30 +1045,14 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
|
|||
<div className="w-full mx-4 h-[75vh]">
|
||||
<Grid numItems={1} className="gap-2 p-8 w-full mt-2">
|
||||
<Col numColSpan={1} className="flex flex-col gap-2">
|
||||
{/* Price Data Reload Section */}
|
||||
{/* Model Management Header */}
|
||||
<div className="flex justify-between items-center mb-4">
|
||||
<div>
|
||||
<h2 className="text-lg font-semibold">Model Management</h2>
|
||||
<p className="text-sm text-gray-600">
|
||||
Manage your models and pricing data
|
||||
Manage your models and configurations
|
||||
</p>
|
||||
</div>
|
||||
{all_admin_roles.includes(userRole) && (
|
||||
<PriceDataReload
|
||||
accessToken={accessToken}
|
||||
onReloadSuccess={() => {
|
||||
// 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"
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
{selectedModelId ? (
|
||||
<ModelInfoView
|
||||
|
|
@ -1134,6 +1118,9 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
|
|||
{all_admin_roles.includes(userRole) && (
|
||||
<Tab>Model Group Alias</Tab>
|
||||
)}
|
||||
{all_admin_roles.includes(userRole) && (
|
||||
<Tab>Price Data Reload</Tab>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="flex items-center space-x-2">
|
||||
|
|
@ -1905,6 +1892,31 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
|
|||
onAliasUpdate={setModelGroupAlias}
|
||||
/>
|
||||
</TabPanel>
|
||||
<TabPanel>
|
||||
<div className="p-6">
|
||||
<div className="mb-6">
|
||||
<Title>Price Data Management</Title>
|
||||
<Text className="text-tremor-content">
|
||||
Manage model pricing data and configure automatic reload schedules
|
||||
</Text>
|
||||
</div>
|
||||
<PriceDataReload
|
||||
accessToken={accessToken}
|
||||
onReloadSuccess={() => {
|
||||
// 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"
|
||||
/>
|
||||
</div>
|
||||
</TabPanel>
|
||||
</TabPanels>
|
||||
</TabGroup>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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<PriceDataReloadProps> = ({
|
|||
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<number>(6);
|
||||
const [reloadStatus, setReloadStatus] = useState<ReloadStatus | null>(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<PriceDataReloadProps> = ({
|
|||
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 (
|
||||
<Popconfirm
|
||||
title="Reload Price Data"
|
||||
description="This will fetch the latest pricing information from the remote source. Continue?"
|
||||
onConfirm={handleReload}
|
||||
okText="Yes"
|
||||
cancelText="No"
|
||||
okButtonProps={{
|
||||
style: {
|
||||
backgroundColor: "#6366f1",
|
||||
borderColor: "#6366f1",
|
||||
color: "white",
|
||||
fontWeight: "500",
|
||||
borderRadius: "0.375rem",
|
||||
padding: "0.375rem 0.75rem",
|
||||
height: "auto",
|
||||
fontSize: "0.875rem",
|
||||
lineHeight: "1.25rem",
|
||||
transition: "all 0.2s ease-in-out",
|
||||
},
|
||||
onMouseEnter: (e) => {
|
||||
e.currentTarget.style.backgroundColor = "#4f46e5";
|
||||
e.currentTarget.style.borderColor = "#4f46e5";
|
||||
},
|
||||
onMouseLeave: (e) => {
|
||||
e.currentTarget.style.backgroundColor = "#6366f1";
|
||||
e.currentTarget.style.borderColor = "#6366f1";
|
||||
},
|
||||
}}
|
||||
>
|
||||
<Tooltip title="Reload latest pricing data from remote source">
|
||||
<Button
|
||||
type={type}
|
||||
size={size}
|
||||
loading={isLoading}
|
||||
className={className}
|
||||
style={{
|
||||
backgroundColor: "#6366f1", // Tremor primary color
|
||||
borderColor: "#6366f1",
|
||||
color: "white",
|
||||
fontWeight: "500",
|
||||
borderRadius: "0.5rem",
|
||||
padding: "0.5rem 1rem",
|
||||
height: "auto",
|
||||
display: "inline-flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
gap: "0.5rem",
|
||||
fontSize: "0.875rem",
|
||||
lineHeight: "1.25rem",
|
||||
transition: "all 0.2s ease-in-out",
|
||||
boxShadow: "0 1px 3px 0 rgba(0, 0, 0, 0.1), 0 1px 2px 0 rgba(0, 0, 0, 0.06)",
|
||||
}}
|
||||
onMouseEnter={(e) => {
|
||||
e.currentTarget.style.backgroundColor = "#4f46e5";
|
||||
e.currentTarget.style.borderColor = "#4f46e5";
|
||||
}}
|
||||
onMouseLeave={(e) => {
|
||||
e.currentTarget.style.backgroundColor = "#6366f1";
|
||||
e.currentTarget.style.borderColor = "#6366f1";
|
||||
<div className={className}>
|
||||
{/* Action Buttons */}
|
||||
<Space direction="horizontal" size="middle" style={{ marginBottom: 16 }}>
|
||||
{/* Hard Refresh Button - Always visible */}
|
||||
<Popconfirm
|
||||
title="Hard Refresh Price Data"
|
||||
description="This will immediately fetch the latest pricing information from the remote source. Continue?"
|
||||
onConfirm={handleHardRefresh}
|
||||
okText="Yes"
|
||||
cancelText="No"
|
||||
okButtonProps={{
|
||||
style: {
|
||||
backgroundColor: "#6366f1",
|
||||
borderColor: "#6366f1",
|
||||
color: "white",
|
||||
fontWeight: "500",
|
||||
borderRadius: "0.375rem",
|
||||
padding: "0.375rem 0.75rem",
|
||||
height: "auto",
|
||||
fontSize: "0.875rem",
|
||||
lineHeight: "1.25rem",
|
||||
transition: "all 0.2s ease-in-out",
|
||||
},
|
||||
onMouseEnter: (e) => {
|
||||
e.currentTarget.style.backgroundColor = "#4f46e5";
|
||||
},
|
||||
onMouseLeave: (e) => {
|
||||
e.currentTarget.style.backgroundColor = "#6366f1";
|
||||
},
|
||||
}}
|
||||
>
|
||||
{showIcon ? "↻ " : ""}{buttonText}
|
||||
</Button>
|
||||
</Tooltip>
|
||||
</Popconfirm>
|
||||
<Button
|
||||
type={type}
|
||||
size={size}
|
||||
loading={isLoading}
|
||||
icon={showIcon ? <ReloadOutlined /> : undefined}
|
||||
style={{
|
||||
backgroundColor: "#6366f1",
|
||||
borderColor: "#6366f1",
|
||||
color: "white",
|
||||
fontWeight: "500",
|
||||
borderRadius: "0.375rem",
|
||||
padding: "0.375rem 0.75rem",
|
||||
height: "auto",
|
||||
fontSize: "0.875rem",
|
||||
lineHeight: "1.25rem",
|
||||
transition: "all 0.2s ease-in-out",
|
||||
}}
|
||||
onMouseEnter={(e) => {
|
||||
e.currentTarget.style.backgroundColor = "#4f46e5";
|
||||
}}
|
||||
onMouseLeave={(e) => {
|
||||
e.currentTarget.style.backgroundColor = "#6366f1";
|
||||
}}
|
||||
>
|
||||
{buttonText}
|
||||
</Button>
|
||||
</Popconfirm>
|
||||
|
||||
{/* Periodic Reload Controls */}
|
||||
{!reloadStatus?.scheduled ? (
|
||||
<Button
|
||||
type="default"
|
||||
size={size}
|
||||
icon={<ClockCircleOutlined />}
|
||||
onClick={() => setShowScheduleModal(true)}
|
||||
style={{
|
||||
borderColor: "#d9d9d9",
|
||||
color: "#6366f1",
|
||||
fontWeight: "500",
|
||||
borderRadius: "0.375rem",
|
||||
padding: "0.375rem 0.75rem",
|
||||
height: "auto",
|
||||
fontSize: "0.875rem",
|
||||
lineHeight: "1.25rem",
|
||||
}}
|
||||
>
|
||||
Set Up Periodic Reload
|
||||
</Button>
|
||||
) : (
|
||||
<Button
|
||||
type="default"
|
||||
size={size}
|
||||
danger
|
||||
icon={<StopOutlined />}
|
||||
loading={isCancelling}
|
||||
onClick={handleCancelReload}
|
||||
style={{
|
||||
borderColor: "#ff4d4f",
|
||||
color: "#ff4d4f",
|
||||
fontWeight: "500",
|
||||
borderRadius: "0.375rem",
|
||||
padding: "0.375rem 0.75rem",
|
||||
height: "auto",
|
||||
fontSize: "0.875rem",
|
||||
lineHeight: "1.25rem",
|
||||
}}
|
||||
>
|
||||
Cancel Periodic Reload
|
||||
</Button>
|
||||
)}
|
||||
</Space>
|
||||
|
||||
{/* Status Card */}
|
||||
{reloadStatus && (
|
||||
<Card
|
||||
size="small"
|
||||
style={{
|
||||
backgroundColor: '#f8f9fa',
|
||||
border: '1px solid #e9ecef',
|
||||
borderRadius: 8
|
||||
}}
|
||||
>
|
||||
<Space direction="vertical" size="small" style={{ width: '100%' }}>
|
||||
{reloadStatus.scheduled ? (
|
||||
<div>
|
||||
<Tag color="green" icon={<ClockCircleOutlined />}>
|
||||
Scheduled every {reloadStatus.interval_hours} hours
|
||||
</Tag>
|
||||
</div>
|
||||
) : (
|
||||
<Text type="secondary">No periodic reload scheduled</Text>
|
||||
)}
|
||||
|
||||
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center' }}>
|
||||
<Text type="secondary" style={{ fontSize: '12px' }}>Last run:</Text>
|
||||
<Text style={{ fontSize: '12px' }}>{formatDateTime(reloadStatus.last_run)}</Text>
|
||||
</div>
|
||||
|
||||
{reloadStatus.scheduled && (
|
||||
<>
|
||||
{reloadStatus.next_run && (
|
||||
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center' }}>
|
||||
<Text type="secondary" style={{ fontSize: '12px' }}>Next run:</Text>
|
||||
<Text style={{ fontSize: '12px' }}>{formatDateTime(reloadStatus.next_run)}</Text>
|
||||
</div>
|
||||
)}
|
||||
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center' }}>
|
||||
<Text type="secondary" style={{ fontSize: '12px' }}>Status:</Text>
|
||||
<Tag color={getStatusColor()}>{getStatusText()}</Tag>
|
||||
</div>
|
||||
</>
|
||||
)}
|
||||
</Space>
|
||||
</Card>
|
||||
)}
|
||||
|
||||
{/* Schedule Modal */}
|
||||
<Modal
|
||||
title="Set Up Periodic Reload"
|
||||
open={showScheduleModal}
|
||||
onOk={handleScheduleReload}
|
||||
onCancel={() => setShowScheduleModal(false)}
|
||||
confirmLoading={isScheduling}
|
||||
okText="Schedule"
|
||||
cancelText="Cancel"
|
||||
okButtonProps={{
|
||||
style: {
|
||||
backgroundColor: "#6366f1",
|
||||
borderColor: "#6366f1",
|
||||
color: "white",
|
||||
},
|
||||
}}
|
||||
>
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<Text>Set up automatic reload of price data every:</Text>
|
||||
</div>
|
||||
<div style={{ marginBottom: 16 }}>
|
||||
<InputNumber
|
||||
min={1}
|
||||
max={168} // 1 week max
|
||||
value={hours}
|
||||
onChange={(value) => setHours(value || 6)}
|
||||
addonAfter="hours"
|
||||
style={{ width: '100%' }}
|
||||
/>
|
||||
</div>
|
||||
<div>
|
||||
<Text type="secondary">
|
||||
This will automatically fetch the latest pricing data from the remote source every {hours} hours.
|
||||
</Text>
|
||||
</div>
|
||||
</Modal>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue