mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge pull request #17528 from BerriAI/litellm_save_background_checks
Add background health checks to db
This commit is contained in:
commit
a21f1ce21f
4 changed files with 627 additions and 7 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
Loading…
Add table
Reference in a new issue