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:
Sameer Kankute 2026-03-27 14:43:16 +05:30 • committed by Yuneng Jiang
parent 397f9ab0fc
commit 94816c7d83
No known key found for this signature in database
4 changed files with 340 additions and 13 deletions

View file

@ -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

View file

@ -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(

View file

@ -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]:

View 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 == {}