From 3d6b7f0d3d9f5264a61f34ba05f38f83934d871c Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 5 Dec 2025 14:27:37 +0530 Subject: [PATCH] Add background health checks to db --- .../health_endpoints/_health_endpoints.py | 205 ++++++++++ litellm/proxy/proxy_server.py | 30 +- litellm/proxy/utils.py | 12 +- .../proxy/test_health_check_functions.py | 387 +++++++++++++++++- 4 files changed, 627 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 2226e190901..5e4784d709e 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e1d5a90dc79..fe8e94d7747 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index aca9bd96eb9..7f23345d26e 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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: diff --git a/tests/test_litellm/proxy/test_health_check_functions.py b/tests/test_litellm/proxy/test_health_check_functions.py index 4f014ce1bee..ccae9fb5425 100644 --- a/tests/test_litellm/proxy/test_health_check_functions.py +++ b/tests/test_litellm/proxy/test_health_check_functions.py @@ -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__]) \ No newline at end of file