mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
feat(router): add health-check-driven routing behind opt-in flag
Background health checks now feed deployment health state into the router candidate-filtering pipeline. Unhealthy deployments are excluded proactively instead of waiting for request failures to trigger cooldown. Gated by `enable_health_check_routing: true` in general_settings. Off by default — zero behavior change for existing users. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
397f9ab0fc
commit
94816c7d83
4 changed files with 340 additions and 13 deletions
|
|
@ -207,17 +207,23 @@ async def _perform_health_check(
|
|||
|
||||
for is_healthy, model in zip(results, model_list):
|
||||
litellm_params = model["litellm_params"]
|
||||
_model_id = (model.get("model_info") or {}).get("id")
|
||||
|
||||
if isinstance(is_healthy, dict) and "error" not in is_healthy:
|
||||
healthy_endpoints.append(
|
||||
_clean_endpoint_data({**litellm_params, **is_healthy}, details)
|
||||
)
|
||||
endpoint_data = {**litellm_params, **is_healthy}
|
||||
if _model_id:
|
||||
endpoint_data["model_id"] = _model_id
|
||||
healthy_endpoints.append(_clean_endpoint_data(endpoint_data, details))
|
||||
elif isinstance(is_healthy, dict):
|
||||
unhealthy_endpoints.append(
|
||||
_clean_endpoint_data({**litellm_params, **is_healthy}, details)
|
||||
)
|
||||
endpoint_data = {**litellm_params, **is_healthy}
|
||||
if _model_id:
|
||||
endpoint_data["model_id"] = _model_id
|
||||
unhealthy_endpoints.append(_clean_endpoint_data(endpoint_data, details))
|
||||
else:
|
||||
unhealthy_endpoints.append(_clean_endpoint_data(litellm_params, details))
|
||||
endpoint_data = {**litellm_params}
|
||||
if _model_id:
|
||||
endpoint_data["model_id"] = _model_id
|
||||
unhealthy_endpoints.append(_clean_endpoint_data(endpoint_data, details))
|
||||
|
||||
return healthy_endpoints, unhealthy_endpoints
|
||||
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ import websockets
|
|||
import websockets.exceptions
|
||||
from pydantic import BaseModel, Json
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm._litellm_uuid import uuid
|
||||
from litellm.constants import (
|
||||
AIOHTTP_CONNECTOR_LIMIT,
|
||||
AIOHTTP_CONNECTOR_LIMIT_PER_HOST,
|
||||
|
|
@ -2112,6 +2112,37 @@ def _schedule_background_health_check_db_save(
|
|||
)
|
||||
|
||||
|
||||
def _write_health_state_to_router_cache(
|
||||
healthy_endpoints: list,
|
||||
unhealthy_endpoints: list,
|
||||
) -> None:
|
||||
"""
|
||||
Write deployment health states to the router's health state cache
|
||||
for health-check-driven routing. No-op if the feature is disabled.
|
||||
"""
|
||||
from litellm.proxy.health_check import build_deployment_health_states
|
||||
|
||||
try:
|
||||
if llm_router is None or not llm_router.enable_health_check_routing:
|
||||
return
|
||||
|
||||
states = build_deployment_health_states(
|
||||
healthy_endpoints=healthy_endpoints,
|
||||
unhealthy_endpoints=unhealthy_endpoints,
|
||||
)
|
||||
if states:
|
||||
llm_router.health_state_cache.set_deployment_health_states(states)
|
||||
verbose_proxy_logger.debug(
|
||||
"health_check_routing_state_updated healthy=%d unhealthy=%d",
|
||||
sum(1 for s in states.values() if s.get("is_healthy")),
|
||||
sum(1 for s in states.values() if not s.get("is_healthy")),
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to write health state to router cache: %s", str(e)
|
||||
)
|
||||
|
||||
|
||||
async def _run_background_health_check():
|
||||
"""
|
||||
Periodically run health checks in the background on the endpoints.
|
||||
|
|
@ -6282,9 +6313,8 @@ class ProxyStartupEvent:
|
|||
KeyRotationManager,
|
||||
)
|
||||
|
||||
# Get prisma_client and proxy_logging_obj from global scope
|
||||
# Get prisma_client from global scope
|
||||
global prisma_client
|
||||
global proxy_logging_obj
|
||||
if prisma_client is not None:
|
||||
key_rotation_manager = KeyRotationManager(prisma_client)
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
|
|||
|
|
@ -46,15 +46,19 @@ import litellm
|
|||
import litellm.litellm_core_utils
|
||||
import litellm.litellm_core_utils.exception_mapping_utils
|
||||
from litellm import get_secret_str
|
||||
from litellm._litellm_uuid import uuid
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.caching import (
|
||||
DualCache,
|
||||
InMemoryCache,
|
||||
RedisCache,
|
||||
RedisClusterCache,
|
||||
)
|
||||
from litellm.constants import DEFAULT_MAX_LRU_CACHE_SIZE
|
||||
from litellm.constants import (
|
||||
DEFAULT_HEALTH_CHECK_INTERVAL,
|
||||
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER,
|
||||
DEFAULT_MAX_LRU_CACHE_SIZE,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.asyncify import run_async_function
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
|
|
@ -113,6 +117,7 @@ from litellm.router_utils.handle_error import (
|
|||
async_raise_no_deployment_exception,
|
||||
send_llm_exception_alert,
|
||||
)
|
||||
from litellm.router_utils.health_state_cache import DeploymentHealthCache
|
||||
from litellm.router_utils.pre_call_checks.deployment_affinity_check import (
|
||||
DeploymentAffinityCheck,
|
||||
)
|
||||
|
|
@ -303,6 +308,8 @@ class Router:
|
|||
deployment_affinity_ttl_seconds: int = 3600,
|
||||
model_group_affinity_config: Optional[Dict[str, List[str]]] = None,
|
||||
ignore_invalid_deployments: bool = False,
|
||||
enable_health_check_routing: bool = False,
|
||||
health_check_staleness_threshold: Optional[int] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
|
||||
|
|
@ -493,6 +500,13 @@ class Router:
|
|||
cache=self.cache, default_cooldown_time=self.cooldown_time
|
||||
)
|
||||
self.disable_cooldowns = disable_cooldowns
|
||||
self.enable_health_check_routing = enable_health_check_routing
|
||||
_staleness = health_check_staleness_threshold or (
|
||||
DEFAULT_HEALTH_CHECK_INTERVAL * DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER
|
||||
)
|
||||
self.health_state_cache = DeploymentHealthCache(
|
||||
cache=self.cache, staleness_threshold=float(_staleness)
|
||||
)
|
||||
self.failed_calls = (
|
||||
InMemoryCache()
|
||||
) # cache to track failed call per deployment, if num failed calls within 1 minute > allowed fails, then add it to cooldown
|
||||
|
|
@ -9150,6 +9164,14 @@ class Router:
|
|||
if isinstance(healthy_deployments, dict):
|
||||
return healthy_deployments
|
||||
|
||||
# Health-check-based filtering (before cooldown)
|
||||
healthy_deployments = (
|
||||
await self._async_filter_health_check_unhealthy_deployments(
|
||||
healthy_deployments=healthy_deployments,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
)
|
||||
|
||||
cooldown_deployments = await _async_get_cooldown_deployments(
|
||||
litellm_router_instance=self, parent_otel_span=parent_otel_span
|
||||
)
|
||||
|
|
@ -9581,6 +9603,13 @@ class Router:
|
|||
parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(
|
||||
request_kwargs
|
||||
)
|
||||
|
||||
# Health-check-based filtering (before cooldown)
|
||||
healthy_deployments = self._filter_health_check_unhealthy_deployments(
|
||||
healthy_deployments=healthy_deployments,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
cooldown_deployments = _get_cooldown_deployments(
|
||||
litellm_router_instance=self, parent_otel_span=parent_otel_span
|
||||
)
|
||||
|
|
@ -9746,10 +9775,14 @@ class Router:
|
|||
llm_provider="",
|
||||
)
|
||||
|
||||
# 4. Apply cooldown filtering
|
||||
# 4. Apply health-check and cooldown filtering
|
||||
parent_otel_span: Optional[Span] = _get_parent_otel_span_from_kwargs(
|
||||
request_kwargs
|
||||
)
|
||||
pass_through_deployments = self._filter_health_check_unhealthy_deployments(
|
||||
healthy_deployments=pass_through_deployments,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
cooldown_deployments = _get_cooldown_deployments(
|
||||
litellm_router_instance=self, parent_otel_span=parent_otel_span
|
||||
)
|
||||
|
|
@ -9871,6 +9904,67 @@ class Router:
|
|||
if deployment["model_info"]["id"] not in cooldown_set
|
||||
]
|
||||
|
||||
async def _async_filter_health_check_unhealthy_deployments(
|
||||
self,
|
||||
healthy_deployments: List[Dict],
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
) -> List[Dict]:
|
||||
"""
|
||||
Filter out deployments marked unhealthy by background health checks.
|
||||
No-op when enable_health_check_routing is False.
|
||||
Returns all deployments if health state is unavailable, stale, or would
|
||||
exclude every candidate (safety net).
|
||||
"""
|
||||
if not self.enable_health_check_routing:
|
||||
return healthy_deployments
|
||||
|
||||
unhealthy_ids = (
|
||||
await self.health_state_cache.async_get_unhealthy_deployment_ids(
|
||||
parent_otel_span=parent_otel_span
|
||||
)
|
||||
)
|
||||
if not unhealthy_ids:
|
||||
return healthy_deployments
|
||||
|
||||
filtered = [
|
||||
d for d in healthy_deployments if d["model_info"]["id"] not in unhealthy_ids
|
||||
]
|
||||
|
||||
if not filtered:
|
||||
verbose_router_logger.warning(
|
||||
"All deployments marked unhealthy by health checks, bypassing health filter"
|
||||
)
|
||||
return healthy_deployments
|
||||
|
||||
return filtered
|
||||
|
||||
def _filter_health_check_unhealthy_deployments(
|
||||
self,
|
||||
healthy_deployments: List[Dict],
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
) -> List[Dict]:
|
||||
"""Sync version of _async_filter_health_check_unhealthy_deployments."""
|
||||
if not self.enable_health_check_routing:
|
||||
return healthy_deployments
|
||||
|
||||
unhealthy_ids = self.health_state_cache.get_unhealthy_deployment_ids(
|
||||
parent_otel_span=parent_otel_span
|
||||
)
|
||||
if not unhealthy_ids:
|
||||
return healthy_deployments
|
||||
|
||||
filtered = [
|
||||
d for d in healthy_deployments if d["model_info"]["id"] not in unhealthy_ids
|
||||
]
|
||||
|
||||
if not filtered:
|
||||
verbose_router_logger.warning(
|
||||
"All deployments marked unhealthy by health checks, bypassing health filter"
|
||||
)
|
||||
return healthy_deployments
|
||||
|
||||
return filtered
|
||||
|
||||
def _filter_pass_through_deployments(
|
||||
self, healthy_deployments: List[Dict]
|
||||
) -> List[Dict]:
|
||||
|
|
|
|||
197
tests/test_litellm/router_utils/test_health_check_routing.py
Normal file
197
tests/test_litellm/router_utils/test_health_check_routing.py
Normal file
|
|
@ -0,0 +1,197 @@
|
|||
"""
|
||||
Tests for health-check-driven routing filter in the Router.
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.router_utils.health_state_cache import DeploymentHealthCache
|
||||
|
||||
|
||||
def _make_deployment(model_id: str, model_name: str = "gpt-4") -> dict:
|
||||
"""Helper to create a deployment dict for testing."""
|
||||
return {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {"model": model_name, "api_key": "fake"},
|
||||
"model_info": {"id": model_id},
|
||||
}
|
||||
|
||||
|
||||
def _make_health_cache(
|
||||
unhealthy_ids: set = None, staleness_threshold: float = 60.0
|
||||
) -> DeploymentHealthCache:
|
||||
"""Create a health cache pre-populated with unhealthy deployment IDs."""
|
||||
cache = DualCache()
|
||||
health_cache = DeploymentHealthCache(
|
||||
cache=cache, staleness_threshold=staleness_threshold
|
||||
)
|
||||
if unhealthy_ids:
|
||||
now = time.time()
|
||||
states = {}
|
||||
for uid in unhealthy_ids:
|
||||
states[uid] = {
|
||||
"is_healthy": False,
|
||||
"timestamp": now,
|
||||
"reason": "test_unhealthy",
|
||||
}
|
||||
health_cache.set_deployment_health_states(states)
|
||||
return health_cache
|
||||
|
||||
|
||||
class TestFilterHealthCheckUnhealthyDeployments:
|
||||
"""Test the sync filter method."""
|
||||
|
||||
def _make_router_like(self, enable: bool, health_cache: DeploymentHealthCache):
|
||||
"""Create a minimal object that behaves like Router for filter testing."""
|
||||
|
||||
class FakeRouter:
|
||||
def __init__(self):
|
||||
self.enable_health_check_routing = enable
|
||||
self.health_state_cache = health_cache
|
||||
|
||||
# Import the actual method and bind it
|
||||
from litellm.router import Router
|
||||
|
||||
fake = FakeRouter()
|
||||
# Use the unbound method
|
||||
fake._filter_health_check_unhealthy_deployments = (
|
||||
Router._filter_health_check_unhealthy_deployments.__get__(fake, FakeRouter)
|
||||
)
|
||||
return fake
|
||||
|
||||
def test_filter_removes_unhealthy_deployments(self):
|
||||
"""Unhealthy deployments should be removed from candidates."""
|
||||
health_cache = _make_health_cache(unhealthy_ids={"deploy-2"})
|
||||
router = self._make_router_like(enable=True, health_cache=health_cache)
|
||||
|
||||
deployments = [
|
||||
_make_deployment("deploy-1"),
|
||||
_make_deployment("deploy-2"),
|
||||
_make_deployment("deploy-3"),
|
||||
]
|
||||
result = router._filter_health_check_unhealthy_deployments(deployments)
|
||||
assert len(result) == 2
|
||||
assert all(d["model_info"]["id"] != "deploy-2" for d in result)
|
||||
|
||||
def test_filter_noop_when_disabled(self):
|
||||
"""When enable_health_check_routing=False, filter should be a no-op."""
|
||||
health_cache = _make_health_cache(unhealthy_ids={"deploy-1"})
|
||||
router = self._make_router_like(enable=False, health_cache=health_cache)
|
||||
|
||||
deployments = [
|
||||
_make_deployment("deploy-1"),
|
||||
_make_deployment("deploy-2"),
|
||||
]
|
||||
result = router._filter_health_check_unhealthy_deployments(deployments)
|
||||
assert len(result) == 2 # no filtering
|
||||
|
||||
def test_filter_returns_all_when_all_unhealthy(self):
|
||||
"""Safety net: if ALL deployments are unhealthy, return all (don't cause outage)."""
|
||||
health_cache = _make_health_cache(
|
||||
unhealthy_ids={"deploy-1", "deploy-2", "deploy-3"}
|
||||
)
|
||||
router = self._make_router_like(enable=True, health_cache=health_cache)
|
||||
|
||||
deployments = [
|
||||
_make_deployment("deploy-1"),
|
||||
_make_deployment("deploy-2"),
|
||||
_make_deployment("deploy-3"),
|
||||
]
|
||||
result = router._filter_health_check_unhealthy_deployments(deployments)
|
||||
assert len(result) == 3 # all returned, safety net
|
||||
|
||||
def test_filter_returns_all_when_cache_empty(self):
|
||||
"""When cache is empty, all deployments should pass through."""
|
||||
health_cache = _make_health_cache() # empty
|
||||
router = self._make_router_like(enable=True, health_cache=health_cache)
|
||||
|
||||
deployments = [
|
||||
_make_deployment("deploy-1"),
|
||||
_make_deployment("deploy-2"),
|
||||
]
|
||||
result = router._filter_health_check_unhealthy_deployments(deployments)
|
||||
assert len(result) == 2
|
||||
|
||||
|
||||
class TestAsyncFilterHealthCheckUnhealthyDeployments:
|
||||
"""Test the async filter method."""
|
||||
|
||||
def _make_router_like(self, enable: bool, health_cache: DeploymentHealthCache):
|
||||
from litellm.router import Router
|
||||
|
||||
class FakeRouter:
|
||||
def __init__(self):
|
||||
self.enable_health_check_routing = enable
|
||||
self.health_state_cache = health_cache
|
||||
|
||||
fake = FakeRouter()
|
||||
fake._async_filter_health_check_unhealthy_deployments = (
|
||||
Router._async_filter_health_check_unhealthy_deployments.__get__(
|
||||
fake, FakeRouter
|
||||
)
|
||||
)
|
||||
return fake
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_filter_removes_unhealthy(self):
|
||||
"""Async version: unhealthy deployments removed."""
|
||||
health_cache = _make_health_cache(unhealthy_ids={"deploy-2"})
|
||||
router = self._make_router_like(enable=True, health_cache=health_cache)
|
||||
|
||||
deployments = [
|
||||
_make_deployment("deploy-1"),
|
||||
_make_deployment("deploy-2"),
|
||||
_make_deployment("deploy-3"),
|
||||
]
|
||||
result = await router._async_filter_health_check_unhealthy_deployments(
|
||||
healthy_deployments=deployments
|
||||
)
|
||||
assert len(result) == 2
|
||||
assert all(d["model_info"]["id"] != "deploy-2" for d in result)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_filter_safety_net(self):
|
||||
"""Async version: safety net when all unhealthy."""
|
||||
health_cache = _make_health_cache(unhealthy_ids={"deploy-1", "deploy-2"})
|
||||
router = self._make_router_like(enable=True, health_cache=health_cache)
|
||||
|
||||
deployments = [
|
||||
_make_deployment("deploy-1"),
|
||||
_make_deployment("deploy-2"),
|
||||
]
|
||||
result = await router._async_filter_health_check_unhealthy_deployments(
|
||||
healthy_deployments=deployments
|
||||
)
|
||||
assert len(result) == 2 # safety net
|
||||
|
||||
|
||||
class TestBuildDeploymentHealthStates:
|
||||
"""Test the build_deployment_health_states function."""
|
||||
|
||||
def test_builds_states_from_endpoints(self):
|
||||
from litellm.proxy.health_check import build_deployment_health_states
|
||||
|
||||
healthy = [{"model": "gpt-4", "model_id": "deploy-1"}]
|
||||
unhealthy = [{"model": "gpt-4", "model_id": "deploy-2", "error": "timeout"}]
|
||||
|
||||
states = build_deployment_health_states(healthy, unhealthy)
|
||||
assert states["deploy-1"]["is_healthy"] is True
|
||||
assert states["deploy-2"]["is_healthy"] is False
|
||||
|
||||
def test_no_model_id_skipped(self):
|
||||
from litellm.proxy.health_check import build_deployment_health_states
|
||||
|
||||
healthy = [{"model": "gpt-4"}] # no model_id
|
||||
unhealthy = [{"model": "gpt-4", "model_id": "deploy-2"}]
|
||||
|
||||
states = build_deployment_health_states(healthy, unhealthy)
|
||||
assert "deploy-1" not in states
|
||||
assert states["deploy-2"]["is_healthy"] is False
|
||||
|
||||
def test_empty_endpoints(self):
|
||||
from litellm.proxy.health_check import build_deployment_health_states
|
||||
|
||||
states = build_deployment_health_states([], [])
|
||||
assert states == {}
|
||||
Loading…
Add table
Reference in a new issue