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
parent dffe2bab02
commit 5ac0bf521b
No known key found for this signature in database
4 changed files with 311 additions and 277 deletions

View file

@ -210,20 +210,20 @@ async def _perform_health_check(
_model_id = (model.get("model_info") or {}).get("id")
if isinstance(is_healthy, dict) and "error" not in is_healthy:
cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details)
endpoint_data = {**litellm_params, **is_healthy}
if _model_id:
cleaned["model_id"] = _model_id
healthy_endpoints.append(cleaned)
endpoint_data["model_id"] = _model_id
healthy_endpoints.append(_clean_endpoint_data(endpoint_data, details))
elif isinstance(is_healthy, dict):
cleaned = _clean_endpoint_data({**litellm_params, **is_healthy}, details)
endpoint_data = {**litellm_params, **is_healthy}
if _model_id:
cleaned["model_id"] = _model_id
unhealthy_endpoints.append(cleaned)
endpoint_data["model_id"] = _model_id
unhealthy_endpoints.append(_clean_endpoint_data(endpoint_data, details))
else:
cleaned = _clean_endpoint_data(litellm_params, details)
endpoint_data = {**litellm_params}
if _model_id:
cleaned["model_id"] = _model_id
unhealthy_endpoints.append(cleaned)
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,
@ -503,9 +503,7 @@ from litellm.proxy.utils import (
get_error_message_str,
get_server_root_path,
handle_exception_on_proxy,
hash_password,
hash_token,
migrate_passwords_to_scrypt_async,
model_dump_with_preserved_fields,
update_spend,
)
@ -872,17 +870,6 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
user_api_key_cache=user_api_key_cache,
)
if prisma_client is not None:
async def _run_pw_migration():
try:
result = await migrate_passwords_to_scrypt_async(prisma_client)
verbose_proxy_logger.info(f"Password migration: {result}")
except Exception as e:
verbose_proxy_logger.warning(f"Password migration skipped: {e}")
asyncio.create_task(_run_pw_migration())
ProxyStartupEvent._initialize_startup_logging(
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
@ -1543,9 +1530,6 @@ shared_aiohttp_session: Optional[
user_api_key_cache = DualCache(
default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value
)
spend_counter_cache = DualCache(
default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value
)
model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(
dual_cache=user_api_key_cache
)
@ -1710,130 +1694,6 @@ def cost_tracking():
)
async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
"""
Read current spend from the cross-pod spend counter.
Reads Redis FIRST (authoritative cross-pod value), not DualCache's
async_get_cache which returns in-memory first. This is critical:
DualCache.async_get_cache returns stale per-pod values because each
pod's in-memory cache is only updated by that pod's own increments.
Fallback chain:
1. Redis counter (cross-pod, authoritative)
2. In-memory counter (single-instance or Redis failure)
3. Cached object's .spend from DB (cold start, no counter yet)
"""
# 1. Try Redis first (cross-pod authoritative)
if spend_counter_cache.redis_cache is not None:
try:
val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key)
if val is not None:
return float(val)
except Exception as e:
verbose_proxy_logger.debug(
"get_current_spend: Redis read failed for %s, falling back to in-memory: %s",
counter_key,
e,
)
# 2. Fall back to in-memory counter (single-instance or Redis failure)
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
if val is not None:
return float(val)
# 3. Final fallback: cached object's spend from DB
return fallback_spend
async def increment_spend_counters(
token: Optional[str],
team_id: Optional[str],
user_id: Optional[str],
response_cost: Optional[float],
):
"""
Atomically increment spend counters for budget enforcement.
Uses spend_counter_cache (DualCache with Redis backend when available)
so counters are shared across all pods. Budget check functions read
from these counters via get_current_spend() (Redis-first).
Awaited (not create_task) in the cost callback, so the counter is
updated before the next request's auth check runs.
"""
if response_cost is None or response_cost == 0:
return
if token is not None:
# token arrives pre-hashed from metadata["user_api_key"] (auth flow
# hashes raw "sk-..." keys before they reach the callback). The
# startswith("sk-") check is a safety net matching update_cache —
# if a raw key somehow arrives, hash it; otherwise use as-is to
# avoid double-hashing (budget checks read valid_token.token which
# is single-hashed).
hashed_token = (
hash_token(token=token)
if isinstance(token, str) and token.startswith("sk-")
else token
)
await _init_and_increment_spend_counter(
counter_key=f"spend:key:{hashed_token}",
source_cache_key=hashed_token,
increment=response_cost,
)
if team_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:team:{team_id}",
source_cache_key=f"team_id:{team_id}",
increment=response_cost,
)
if user_id is not None and team_id is not None:
await _init_and_increment_spend_counter(
counter_key=f"spend:team_member:{user_id}:{team_id}",
source_cache_key=f"team_membership:{user_id}:{team_id}",
increment=response_cost,
)
async def _init_and_increment_spend_counter(
counter_key: str,
source_cache_key: str,
increment: float,
):
"""
Initialize counter from cached object's DB-loaded spend if not yet set,
then atomically increment in both in-memory and Redis.
On first access per pod:
1. Check spend_counter_cache (in-memory -> Redis via DualCache for init check)
2. If not found anywhere, read base spend from user_api_key_cache (DB-loaded object)
3. Seed counter via async_increment_cache (not async_set_cache) to avoid a
check-then-set race: if two pods cold-start simultaneously, both may see
the counter as absent and seed it. Using increment instead of set means
the worst case is over-counting (conservative — blocks slightly early)
rather than under-counting (would allow overspend).
4. Increment atomically (both in-memory + Redis)
"""
current = await spend_counter_cache.async_get_cache(key=counter_key)
if current is None:
source = await user_api_key_cache.async_get_cache(key=source_cache_key)
base_spend = 0.0
if source is not None:
if isinstance(source, dict):
base_spend = source.get("spend", 0.0) or 0.0
else:
base_spend = getattr(source, "spend", 0.0) or 0.0
if base_spend > 0:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=base_spend
)
await spend_counter_cache.async_increment_cache(key=counter_key, value=increment)
async def update_cache( # noqa: PLR0915
token: Optional[str],
user_id: Optional[str],
@ -2278,7 +2138,7 @@ def _write_health_state_to_router_cache(
sum(1 for s in states.values() if not s.get("is_healthy")),
)
except Exception as e:
verbose_proxy_logger.warning(
verbose_proxy_logger.debug(
"Failed to write health state to router cache: %s", str(e)
)
@ -2702,7 +2562,6 @@ class ProxyConfig:
):
## INIT PROXY REDIS USAGE CLIENT ##
redis_usage_cache = litellm.cache.cache
spend_counter_cache.redis_cache = redis_usage_cache
# Note: PKCE verifier storage uses redis_usage_cache directly (not
# user_api_key_cache) to avoid routing all API-key lookups through Redis.
@ -2870,36 +2729,12 @@ class ProxyConfig:
return search_tools_parsed if search_tools_parsed else None
# Environment variable keys that must not be overridden via config because
# they can alter process execution, library loading, or network routing.
_BLOCKED_ENV_KEYS: Set[str] = {
"PATH",
"LD_PRELOAD",
"LD_LIBRARY_PATH",
"DYLD_LIBRARY_PATH",
"DYLD_INSERT_LIBRARIES",
"PYTHONPATH",
"PYTHONSTARTUP",
"PYTHONHOME",
"HOME",
"USER",
"SHELL",
"LOGNAME",
"NO_PROXY",
"no_proxy",
}
def _load_environment_variables(self, config: dict):
## ENVIRONMENT VARIABLES
global premium_user
environment_variables = config.get("environment_variables", None)
if environment_variables:
for key, value in environment_variables.items():
if key in self._BLOCKED_ENV_KEYS:
verbose_proxy_logger.warning(
"Skipping blocked environment variable key: %s", key
)
continue
#########################################################
# handles this scenario:
# ```yaml
@ -5730,11 +5565,6 @@ def _restamp_streaming_chunk_model(
if _is_azure_model_router_request(requested_model_from_client):
return chunk, model_mismatch_logged
# For fastest_response batch completions, preserve the winning model's name
# instead of stamping the comma-separated list the client sent.
if request_data.get("fastest_response", False):
return chunk, model_mismatch_logged
downstream_model = (
chunk.get("model") if isinstance(chunk, dict) else getattr(chunk, "model", None)
)
@ -6501,9 +6331,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(
@ -9230,7 +9059,7 @@ def _add_team_models_to_all_models(
for team_object in team_db_objects_typed:
if (
not team_object.models # None or empty list = all model access
len(team_object.models) == 0 # empty list = all model access
or SpecialModelNames.all_proxy_models.value in team_object.models
):
model_list = llm_router.get_model_list()
@ -9264,75 +9093,6 @@ def _add_team_models_to_all_models(
return team_models
async def _add_access_group_models_to_team_models(
team_db_objects_typed: List[LiteLLM_TeamTable],
llm_router: Router,
prisma_client: PrismaClient,
team_models: Dict[str, Set[str]],
) -> Dict[str, Set[str]]:
"""
Resolve models reachable via team access groups and merge them into team_models.
Batch-fetches all distinct access groups in a single DB query, then resolves
each eligible team's access group models via the pre-fetched map.
This ensures models associated with a team only through access groups
(not directly in team.models) are included in the UI model listing.
"""
# First pass: identify eligible teams and collect all distinct access group IDs
eligible_teams: List[LiteLLM_TeamTable] = []
all_access_group_ids: Set[str] = set()
for team_object in team_db_objects_typed:
if not team_object.access_group_ids:
continue
# Skip teams with empty models list — they already have access to everything
# (handled by _add_team_models_to_all_models)
if (
not team_object.models
or SpecialModelNames.all_proxy_models.value in team_object.models
):
continue
eligible_teams.append(team_object)
all_access_group_ids.update(team_object.access_group_ids)
if not eligible_teams:
return team_models
# Single batch fetch for all access groups
access_group_rows = (
await prisma_client.db.litellm_accessgrouptable.find_many(
where={"access_group_id": {"in": list(all_access_group_ids)}}
)
)
ag_model_map: Dict[str, List[str]] = {
row.access_group_id: row.access_model_names or []
for row in access_group_rows
}
# Second pass: resolve deployments for each eligible team
for team_object in eligible_teams:
model_names: Set[str] = set()
for ag_id in team_object.access_group_ids or [] :
model_names.update(ag_model_map.get(ag_id, []))
for model_name in model_names:
deployments = llm_router.get_model_list(
model_name=model_name, team_id=team_object.team_id
)
if deployments is not None:
for deployment in deployments:
model_id = deployment.get("model_info", {}).get("id", None)
if model_id is not None:
team_models.setdefault(model_id, set()).add(
team_object.team_id
)
return team_models
async def get_all_team_models(
user_teams: Union[List[str], Literal["*"]],
prisma_client: PrismaClient,
@ -9369,14 +9129,6 @@ async def get_all_team_models(
llm_router=llm_router,
)
# Also resolve models reachable via team access groups
team_models = await _add_access_group_models_to_team_models(
team_db_objects_typed=team_db_objects_typed,
llm_router=llm_router,
prisma_client=prisma_client,
team_models=team_models,
)
# convert set to list
returned_team_models: Dict[str, List[str]] = {}
for model_id, team_ids in team_models.items():
@ -9884,7 +9636,7 @@ async def _filter_models_by_team_id(
team_accessible_model_ids: Set[str] = set()
if (
not team_object.models # empty list = all model access
len(team_object.models) == 0 # empty list = all model access
or SpecialModelNames.all_proxy_models.value in team_object.models
):
# Team has access to all models
@ -11770,9 +11522,9 @@ async def claim_onboarding_link(data: InvitationClaim):
},
)
### UPDATE USER OBJECT ###
hashed_pw = hash_password(data.password)
hash_password = hash_token(token=data.password)
user_obj = await prisma_client.db.litellm_usertable.update(
where={"user_id": invite_obj.user_id}, data={"password": hashed_pw}
where={"user_id": invite_obj.user_id}, data={"password": hash_password}
)
if user_obj is None:
@ -11792,8 +11544,6 @@ async def claim_onboarding_link(data: InvitationClaim):
},
)
if user_obj and hasattr(user_obj, "__dict__"):
user_obj.__dict__.pop("password", None)
return user_obj
@ -12232,10 +11982,7 @@ async def invitation_delete(
dependencies=[Depends(user_api_key_auth)],
include_in_schema=False,
)
async def update_config( # noqa: PLR0915
config_info: ConfigYAML,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
async def update_config(config_info: ConfigYAML): # noqa: PLR0915
"""
For Admin UI - allows admin to update config via UI
@ -12243,10 +11990,6 @@ async def update_config( # noqa: PLR0915
"""
global llm_router, llm_model_list, general_settings, proxy_config, proxy_logging_obj, master_key, prisma_client
try:
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=403, detail="Only proxy admins can update config"
)
import base64
"""

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
@ -9154,6 +9168,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
)
@ -9585,6 +9607,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
)
@ -9750,10 +9779,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
)
@ -9875,6 +9908,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 == {}