diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 3e05ee3c484..058f2f4ed9d 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 271fb72e183..42740c24f45 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 """ diff --git a/litellm/router.py b/litellm/router.py index 5cd4f837782..8d0e3334cb2 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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]: diff --git a/tests/test_litellm/router_utils/test_health_check_routing.py b/tests/test_litellm/router_utils/test_health_check_routing.py new file mode 100644 index 00000000000..f40144b44c9 --- /dev/null +++ b/tests/test_litellm/router_utils/test_health_check_routing.py @@ -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 == {}