Merge pull request #17528 from BerriAI/litellm_save_background_checks

Add background health checks to db
This commit is contained in:
Sameer Kankute 2025-12-05 22:25:03 +05:30 • committed by GitHub
commit a21f1ce21f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 627 additions and 7 deletions

View file

@ -397,6 +397,211 @@ async def _save_health_check_to_db(
# Continue execution - don't let database save failure break health checks
def _build_model_param_to_info_mapping(model_list: list) -> dict:
"""
Build a mapping from model parameter to model info (model_name, model_id).
Multiple models might share the same model parameter, so we use a list.
Args:
model_list: List of model configurations
Returns:
Dictionary mapping model parameter to list of model info dicts
"""
model_param_to_info = {}
for model in model_list:
model_info = model.get("model_info", {})
model_name = model.get("model_name")
model_id = model_info.get("id")
litellm_params = model.get("litellm_params", {})
model_param = litellm_params.get("model")
if model_param and model_name:
if model_param not in model_param_to_info:
model_param_to_info[model_param] = []
model_param_to_info[model_param].append({
"model_name": model_name,
"model_id": model_id,
})
return model_param_to_info
def _aggregate_health_check_results(
model_param_to_info: dict,
healthy_endpoints: list,
unhealthy_endpoints: list,
) -> dict:
"""
Aggregate health check results per unique model.
Uses (model_id, model_name) as key, or (None, model_name) if model_id is None.
Args:
model_param_to_info: Mapping from model parameter to model info
healthy_endpoints: List of healthy endpoint results
unhealthy_endpoints: List of unhealthy endpoint results
Returns:
Dictionary mapping (model_id, model_name) to aggregated health check results
"""
model_results = {}
# Process healthy endpoints
for endpoint in healthy_endpoints:
model_param = endpoint.get("model")
if model_param and model_param in model_param_to_info:
for model_info in model_param_to_info[model_param]:
key = (model_info["model_id"], model_info["model_name"])
if key not in model_results:
model_results[key] = {
"model_name": model_info["model_name"],
"model_id": model_info["model_id"],
"healthy_count": 0,
"unhealthy_count": 0,
"error_message": None,
}
model_results[key]["healthy_count"] += 1
# Process unhealthy endpoints
for endpoint in unhealthy_endpoints:
model_param = endpoint.get("model")
error_message = endpoint.get("error")
if model_param and model_param in model_param_to_info:
for model_info in model_param_to_info[model_param]:
key = (model_info["model_id"], model_info["model_name"])
if key not in model_results:
model_results[key] = {
"model_name": model_info["model_name"],
"model_id": model_info["model_id"],
"healthy_count": 0,
"unhealthy_count": 0,
"error_message": None,
}
model_results[key]["unhealthy_count"] += 1
# Use the first error message encountered
if not model_results[key]["error_message"] and error_message:
model_results[key]["error_message"] = str(error_message)[:500]
return model_results
async def _save_health_check_results_if_changed(
prisma_client,
model_results: dict,
latest_checks_map: dict,
start_time: float,
checked_by: Optional[str] = None,
):
"""
Save health check results to database, but only if status changed or >1 hour since last save.
OPTIMIZATION: Only saves to database if the status has changed from the last saved check.
This dramatically reduces database writes when health status remains stable.
- Stable systems: ~1 write/hour per model (instead of 12 writes/hour with 5-min intervals)
- Status changes: Immediate write (no delay)
- Result: ~92% reduction in DB writes for stable systems, while maintaining real-time updates on changes
Args:
prisma_client: Database client
model_results: Dictionary of aggregated health check results per model
latest_checks_map: Dictionary mapping model_id/model_name to latest health check
start_time: Start time of health check for calculating response time
checked_by: Identifier for who/what performed the check
"""
for result in model_results.values():
new_status = "healthy" if result["healthy_count"] > 0 else "unhealthy"
# Check if we should save this result
should_save = True
lookup_key = result["model_id"] if result["model_id"] else result["model_name"]
if lookup_key in latest_checks_map:
last_check = latest_checks_map[lookup_key]
# Only save if status changed or if it's been a while since last check
if last_check.status == new_status:
# Check if last check was recent (within 1 hour)
if last_check.checked_at:
from datetime import datetime, timezone
time_since_last_check = (
datetime.now(timezone.utc) - last_check.checked_at
).total_seconds()
# Only skip if status unchanged AND checked recently (within 1 hour)
# This ensures we still get periodic updates even if status is stable
if time_since_last_check < 3600: # 1 hour threshold
should_save = False
if should_save:
asyncio.create_task(
prisma_client.save_health_check_result(
model_name=result["model_name"],
model_id=result["model_id"],
status=new_status,
healthy_count=result["healthy_count"],
unhealthy_count=result["unhealthy_count"],
error_message=result["error_message"],
response_time_ms=(time.time() - start_time) * 1000,
details=None,
checked_by=checked_by,
)
)
async def _save_background_health_checks_to_db(
prisma_client,
model_list: list,
healthy_endpoints: list,
unhealthy_endpoints: list,
start_time: float,
checked_by: Optional[str] = None,
):
"""
Save background health check results to database for each model.
Maps health check endpoints back to their original models to get model_name and model_id.
Aggregates results per unique model (by model_id if available, otherwise model_name).
OPTIMIZATION: Only saves to database if the status has changed from the last saved check.
This dramatically reduces database writes when health status remains stable.
"""
if prisma_client is None:
return
try:
# Step 1: Build mapping from model parameter to model info
model_param_to_info = _build_model_param_to_info_mapping(model_list)
# Step 2: Aggregate health check results per unique model
model_results = _aggregate_health_check_results(
model_param_to_info,
healthy_endpoints,
unhealthy_endpoints,
)
# Step 3: Get latest health checks for all models in one query to compare status
latest_checks = await prisma_client.get_all_latest_health_checks()
latest_checks_map = {}
for check in latest_checks:
# Use model_id as primary key, fallback to model_name
key = check.model_id if check.model_id else check.model_name
if key not in latest_checks_map:
latest_checks_map[key] = check
# Step 4: Save aggregated results, but only if status changed
await _save_health_check_results_if_changed(
prisma_client,
model_results,
latest_checks_map,
start_time,
checked_by,
)
except Exception as db_error:
verbose_proxy_logger.warning(
f"Failed to save background health checks to database: {db_error}"
)
# Continue execution - don't let database save failure break health checks
async def _perform_health_check_and_save(
model_list,
target_model,

View file

@ -1578,7 +1578,7 @@ async def _run_background_health_check():
Update health_check_results, based on this.
Uses shared health check state when Redis is available to coordinate across pods.
"""
global health_check_results, llm_model_list, health_check_interval, health_check_details, use_shared_health_check, redis_usage_cache
global health_check_results, llm_model_list, health_check_interval, health_check_details, use_shared_health_check, redis_usage_cache, prisma_client
if (
health_check_interval is None
@ -1645,6 +1645,34 @@ async def _run_background_health_check():
health_check_results["healthy_count"] = len(healthy_endpoints)
health_check_results["unhealthy_count"] = len(unhealthy_endpoints)
# Save background health checks to database (non-blocking)
if prisma_client is not None:
import time as time_module
from litellm.proxy.health_endpoints._health_endpoints import (
_save_background_health_checks_to_db,
)
# Use pod_id or a system identifier for checked_by if shared health check is enabled
checked_by = None
if shared_health_manager is not None:
checked_by = shared_health_manager.pod_id
else:
# Use a system identifier for background health checks
checked_by = "background_health_check"
start_time = time_module.time()
asyncio.create_task(
_save_background_health_checks_to_db(
prisma_client,
_llm_model_list,
healthy_endpoints,
unhealthy_endpoints,
start_time,
checked_by=checked_by,
)
)
await asyncio.sleep(health_check_interval)

View file

@ -3094,8 +3094,16 @@ class PrismaClient:
# Group by model_name and get the latest for each
latest_checks = {}
for check in all_checks:
if check.model_name not in latest_checks:
latest_checks[check.model_name] = check
# Create a unique key: prefer model_id if available, otherwise use model_name
# This ensures we get the latest check for each unique model
if check.model_id:
key = (check.model_id, check.model_name)
else:
key = (None, check.model_name)
# Only add if we haven't seen this key yet (since checks are ordered by checked_at desc)
if key not in latest_checks:
latest_checks[key] = check
return list(latest_checks.values())
except Exception as e:

View file

@ -1,13 +1,22 @@
import asyncio
import pytest
from unittest.mock import AsyncMock, MagicMock
import sys
import os
import sys
import time
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.health_endpoints._health_endpoints import (
_aggregate_health_check_results,
_build_model_param_to_info_mapping,
_save_background_health_checks_to_db,
_save_health_check_results_if_changed,
_save_health_check_to_db,
)
from litellm.proxy.utils import PrismaClient
from litellm.proxy.health_endpoints._health_endpoints import _save_health_check_to_db
@pytest.fixture
@ -87,5 +96,375 @@ async def test_save_health_check_to_db_no_client():
assert result is None
# Tests for background health check functions
def test_build_model_param_to_info_mapping():
"""Test building model parameter to info mapping"""
model_list = [
{
"model_name": "gpt-3.5-turbo",
"model_info": {"id": "model-123"},
"litellm_params": {"model": "gpt-3.5-turbo"},
},
{
"model_name": "gpt-4",
"model_info": {"id": "model-456"},
"litellm_params": {"model": "gpt-4"},
},
{
"model_name": "gpt-3.5-turbo-alias",
"model_info": {"id": "model-789"},
"litellm_params": {"model": "gpt-3.5-turbo"}, # Same model param
},
]
result = _build_model_param_to_info_mapping(model_list)
assert "gpt-3.5-turbo" in result
assert "gpt-4" in result
assert len(result["gpt-3.5-turbo"]) == 2 # Two models share same param
assert len(result["gpt-4"]) == 1
assert result["gpt-3.5-turbo"][0]["model_name"] == "gpt-3.5-turbo"
assert result["gpt-3.5-turbo"][0]["model_id"] == "model-123"
assert result["gpt-3.5-turbo"][1]["model_name"] == "gpt-3.5-turbo-alias"
assert result["gpt-3.5-turbo"][1]["model_id"] == "model-789"
def test_build_model_param_to_info_mapping_no_model_name():
"""Test mapping skips models without model_name"""
model_list = [
{
"model_info": {"id": "model-123"},
"litellm_params": {"model": "gpt-3.5-turbo"},
},
]
result = _build_model_param_to_info_mapping(model_list)
assert len(result) == 0
def test_aggregate_health_check_results():
"""Test aggregating health check results per model"""
model_param_to_info = {
"gpt-3.5-turbo": [
{"model_name": "gpt-3.5-turbo", "model_id": "model-123"},
],
"gpt-4": [
{"model_name": "gpt-4", "model_id": "model-456"},
],
}
healthy_endpoints = [
{"model": "gpt-3.5-turbo"},
]
unhealthy_endpoints = [
{"model": "gpt-4", "error": "Rate limit exceeded"},
]
result = _aggregate_health_check_results(
model_param_to_info, healthy_endpoints, unhealthy_endpoints
)
# Check gpt-3.5-turbo is healthy
gpt35_key = ("model-123", "gpt-3.5-turbo")
assert gpt35_key in result
assert result[gpt35_key]["healthy_count"] == 1
assert result[gpt35_key]["unhealthy_count"] == 0
assert result[gpt35_key]["error_message"] is None
# Check gpt-4 is unhealthy
gpt4_key = ("model-456", "gpt-4")
assert gpt4_key in result
assert result[gpt4_key]["healthy_count"] == 0
assert result[gpt4_key]["unhealthy_count"] == 1
assert "Rate limit" in result[gpt4_key]["error_message"]
def test_aggregate_health_check_results_multiple_endpoints():
"""Test aggregation with multiple endpoints for same model"""
model_param_to_info = {
"gpt-3.5-turbo": [
{"model_name": "gpt-3.5-turbo", "model_id": "model-123"},
],
}
healthy_endpoints = [
{"model": "gpt-3.5-turbo"},
{"model": "gpt-3.5-turbo"},
]
unhealthy_endpoints = []
result = _aggregate_health_check_results(
model_param_to_info, healthy_endpoints, unhealthy_endpoints
)
key = ("model-123", "gpt-3.5-turbo")
assert result[key]["healthy_count"] == 2
assert result[key]["unhealthy_count"] == 0
@pytest.mark.asyncio
async def test_save_health_check_results_if_changed_status_changed():
"""Test saving when status changes"""
mock_prisma = MagicMock()
mock_prisma.save_health_check_result = AsyncMock()
model_results = {
("model-123", "gpt-3.5-turbo"): {
"model_name": "gpt-3.5-turbo",
"model_id": "model-123",
"healthy_count": 1,
"unhealthy_count": 0,
"error_message": None,
},
}
# Latest check shows unhealthy, new result is healthy (status changed)
latest_checks_map = {
"model-123": MagicMock(
status="unhealthy",
checked_at=datetime.now(timezone.utc) - timedelta(minutes=5),
),
}
start_time = 1234567890.0
await _save_health_check_results_if_changed(
mock_prisma, model_results, latest_checks_map, start_time, "background_health_check"
)
# Should save because status changed
mock_prisma.save_health_check_result.assert_called_once()
call_kwargs = mock_prisma.save_health_check_result.call_args[1]
assert call_kwargs["status"] == "healthy"
assert call_kwargs["model_name"] == "gpt-3.5-turbo"
assert call_kwargs["checked_by"] == "background_health_check"
@pytest.mark.asyncio
async def test_save_health_check_results_if_changed_status_unchanged_recent():
"""Test skipping save when status unchanged and checked recently"""
mock_prisma = MagicMock()
mock_prisma.save_health_check_result = AsyncMock()
model_results = {
("model-123", "gpt-3.5-turbo"): {
"model_name": "gpt-3.5-turbo",
"model_id": "model-123",
"healthy_count": 1,
"unhealthy_count": 0,
"error_message": None,
},
}
# Latest check shows healthy, new result is healthy (status unchanged)
# And checked recently (within 1 hour)
latest_checks_map = {
"model-123": MagicMock(
status="healthy",
checked_at=datetime.now(timezone.utc) - timedelta(minutes=30),
),
}
start_time = 1234567890.0
await _save_health_check_results_if_changed(
mock_prisma, model_results, latest_checks_map, start_time, "background_health_check"
)
# Should NOT save because status unchanged and checked recently
mock_prisma.save_health_check_result.assert_not_called()
@pytest.mark.asyncio
async def test_save_health_check_results_if_changed_status_unchanged_old():
"""Test saving when status unchanged but last check is old (>1 hour)"""
mock_prisma = MagicMock()
mock_prisma.save_health_check_result = AsyncMock()
model_results = {
("model-123", "gpt-3.5-turbo"): {
"model_name": "gpt-3.5-turbo",
"model_id": "model-123",
"healthy_count": 1,
"unhealthy_count": 0,
"error_message": None,
},
}
# Latest check shows healthy, new result is healthy (status unchanged)
# But checked >1 hour ago
latest_checks_map = {
"model-123": MagicMock(
status="healthy",
checked_at=datetime.now(timezone.utc) - timedelta(hours=2),
),
}
start_time = 1234567890.0
await _save_health_check_results_if_changed(
mock_prisma, model_results, latest_checks_map, start_time, "background_health_check"
)
# Should save because last check is old (>1 hour)
mock_prisma.save_health_check_result.assert_called_once()
@pytest.mark.asyncio
async def test_save_health_check_results_if_changed_no_previous_check():
"""Test saving when there's no previous check"""
mock_prisma = MagicMock()
mock_prisma.save_health_check_result = AsyncMock()
model_results = {
("model-123", "gpt-3.5-turbo"): {
"model_name": "gpt-3.5-turbo",
"model_id": "model-123",
"healthy_count": 1,
"unhealthy_count": 0,
"error_message": None,
},
}
# No previous check
latest_checks_map = {}
start_time = 1234567890.0
await _save_health_check_results_if_changed(
mock_prisma, model_results, latest_checks_map, start_time, "background_health_check"
)
# Should save because no previous check
mock_prisma.save_health_check_result.assert_called_once()
@pytest.mark.asyncio
async def test_save_background_health_checks_to_db():
"""Test the main background health check save function"""
mock_prisma = MagicMock()
mock_prisma.save_health_check_result = AsyncMock()
mock_prisma.get_all_latest_health_checks = AsyncMock(return_value=[])
model_list = [
{
"model_name": "gpt-3.5-turbo",
"model_info": {"id": "model-123"},
"litellm_params": {"model": "gpt-3.5-turbo"},
},
]
healthy_endpoints = [{"model": "gpt-3.5-turbo"}]
unhealthy_endpoints = []
start_time = 1234567890.0
await _save_background_health_checks_to_db(
mock_prisma, model_list, healthy_endpoints, unhealthy_endpoints, start_time, "background_health_check"
)
# Should call get_all_latest_health_checks and save_health_check_result
mock_prisma.get_all_latest_health_checks.assert_called_once()
mock_prisma.save_health_check_result.assert_called_once()
call_kwargs = mock_prisma.save_health_check_result.call_args[1]
assert call_kwargs["model_name"] == "gpt-3.5-turbo"
assert call_kwargs["model_id"] == "model-123"
assert call_kwargs["status"] == "healthy"
assert call_kwargs["checked_by"] == "background_health_check"
@pytest.mark.asyncio
async def test_save_background_health_checks_to_db_no_prisma():
"""Test graceful handling when no prisma client"""
result = await _save_background_health_checks_to_db(
None, [], [], [], 0.0, "background_health_check"
)
assert result is None
@pytest.mark.asyncio
async def test_save_background_health_checks_to_db_exception_handling():
"""Test exception handling in background health check save"""
mock_prisma = MagicMock()
mock_prisma.get_all_latest_health_checks = AsyncMock(side_effect=Exception("DB Error"))
model_list = [
{
"model_name": "gpt-3.5-turbo",
"model_info": {"id": "model-123"},
"litellm_params": {"model": "gpt-3.5-turbo"},
},
]
# Should not raise exception, should handle gracefully
await _save_background_health_checks_to_db(
mock_prisma, model_list, [], [], 0.0, "background_health_check"
)
# Function should complete without raising
@pytest.mark.asyncio
async def test_get_all_latest_health_checks_with_model_id(mock_prisma):
"""Test get_all_latest_health_checks properly groups by model_id"""
# Create mock checks with same model_name but different model_id
mock_check1 = MagicMock()
mock_check1.model_id = "model-123"
mock_check1.model_name = "gpt-3.5-turbo"
mock_check1.checked_at = datetime.now(timezone.utc) - timedelta(minutes=10)
mock_check2 = MagicMock()
mock_check2.model_id = "model-456"
mock_check2.model_name = "gpt-3.5-turbo"
mock_check2.checked_at = datetime.now(timezone.utc) - timedelta(minutes=5)
mock_check3 = MagicMock()
mock_check3.model_id = "model-123"
mock_check3.model_name = "gpt-3.5-turbo"
mock_check3.checked_at = datetime.now(timezone.utc) - timedelta(minutes=1) # Latest for model-123
# Order by checked_at desc
mock_prisma.db.litellm_healthchecktable.find_many = AsyncMock(
return_value=[mock_check3, mock_check2, mock_check1]
)
result = await mock_prisma.get_all_latest_health_checks()
# Should return 2 unique models (by model_id)
assert len(result) == 2
# Should have latest check for each model_id
model_ids = {check.model_id for check in result}
assert "model-123" in model_ids
assert "model-456" in model_ids
# model-123 should have the latest check (1 minute ago)
model123_check = next(c for c in result if c.model_id == "model-123")
assert model123_check.checked_at == mock_check3.checked_at
@pytest.mark.asyncio
async def test_get_all_latest_health_checks_without_model_id(mock_prisma):
"""Test get_all_latest_health_checks groups by model_name when model_id is None"""
mock_check1 = MagicMock()
mock_check1.model_id = None
mock_check1.model_name = "gpt-3.5-turbo"
mock_check1.checked_at = datetime.now(timezone.utc) - timedelta(minutes=10)
mock_check2 = MagicMock()
mock_check2.model_id = None
mock_check2.model_name = "gpt-3.5-turbo"
mock_check2.checked_at = datetime.now(timezone.utc) - timedelta(minutes=1) # Latest
mock_prisma.db.litellm_healthchecktable.find_many = AsyncMock(
return_value=[mock_check2, mock_check1]
)
result = await mock_prisma.get_all_latest_health_checks()
# Should return 1 unique model (by model_name)
assert len(result) == 1
assert result[0].model_name == "gpt-3.5-turbo"
assert result[0].checked_at == mock_check2.checked_at # Latest
if __name__ == "__main__":
pytest.main([__file__])