mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +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
dffe2bab02
commit
5ac0bf521b
4 changed files with 311 additions and 277 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
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