diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index a6b2ec47325..b2ca2c30131 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -776,6 +776,35 @@ def _strip_admin_only_fields_from_health_result(result: dict) -> dict: return out +def _resolve_targeted_model_ids( + model_list: list, model: Optional[str], model_id: Optional[str] +) -> Optional[set]: + """ + Resolve a ``/health`` ``model`` / ``model_id`` query param to the set of + deployment IDs the response should be scoped to. + + Mirrors the live-path semantics in ``perform_health_check()``: ``model`` + matches either the deployment's ``model_name`` alias or its + ``litellm_params.model`` provider string. ``model_id`` is taken as-is. + + Returns ``None`` when no targeting is requested — callers should treat + that as "no filter." + """ + if not model and not model_id: + return None + if model_id: + return {model_id} + target_ids: set = set() + for m in model_list: + deployment_id = (m.get("model_info") or {}).get("id") + if not deployment_id: + continue + litellm_model = (m.get("litellm_params") or {}).get("model") + if m.get("model_name") == model or litellm_model == model: + target_ids.add(deployment_id) + return target_ids + + def _filter_health_check_results_by_model_ids( results: dict, allowed_model_ids: set ) -> dict: @@ -982,16 +1011,29 @@ async def health_endpoint( m for m in _llm_model_list if m.get("model_name") in allowed_models ] if use_background_health_checks: + # The cached background result covers every model. When the + # caller targets a specific model/model_id we have to narrow the + # cache to that deployment before _post_process evaluates + # healthy_count, otherwise an unhealthy "foo" combined with any + # other healthy model would still report healthy_count > 0 and + # the targeted-503 path would never fire. + targeted_ids = _resolve_targeted_model_ids(_llm_model_list, model, model_id) if len(user_api_key_dict.models) > 0: allowed_model_ids = { (m.get("model_info") or {}).get("id") for m in _llm_model_list if (m.get("model_info") or {}).get("id") } - filtered = _filter_health_check_results_by_model_ids( - health_check_results, allowed_model_ids + # _llm_model_list is already scoped to the caller's allowed + # model_names above, so targeted_ids is implicitly the + # intersection of "targeted" and "allowed." + filter_ids = ( + targeted_ids if targeted_ids is not None else allowed_model_ids ) - if not allowed_model_ids: + filtered = _filter_health_check_results_by_model_ids( + health_check_results, filter_ids + ) + if targeted_ids is None and not allowed_model_ids: # Caller has accessible model_names but none of the # matching deployments expose a model_info.id, so the # cache filter (which keys on model_id) drops every @@ -1012,6 +1054,15 @@ async def health_endpoint( "to populate model_info.id for these models." ] return _post_process(filtered) + if targeted_ids is not None: + # Admin caller targeting a specific model: filter the cache + # so the response (and the targeted-503 check) reflects only + # that deployment, not the global aggregate. + return _post_process( + _filter_health_check_results_by_model_ids( + health_check_results, targeted_ids + ) + ) return _post_process(health_check_results) else: router_result = await _perform_health_check_and_save( diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index c7d31908dab..d5102daa438 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -927,8 +927,14 @@ async def test_health_endpoint_filters_background_cache_by_user_access(): ): from fastapi import Response + # Pass model=None, model_id=None explicitly: direct calls to the + # handler skip FastAPI's Query() resolution, so unspecified params + # would otherwise carry the Query() sentinel (which is truthy). result = await health_endpoint( - response=Response(), user_api_key_dict=user_api_key_dict + response=Response(), + user_api_key_dict=user_api_key_dict, + model=None, + model_id=None, ) # Sanity: the source cache had two entries before scoping; the scoping @@ -1016,10 +1022,16 @@ async def test_health_endpoint_admin_sees_routing_fields_non_admin_does_not(): admin_response = Response() non_admin_response = Response() admin_result = await health_endpoint( - response=admin_response, user_api_key_dict=admin_key + response=admin_response, + user_api_key_dict=admin_key, + model=None, + model_id=None, ) non_admin_result = await health_endpoint( - response=non_admin_response, user_api_key_dict=non_admin_key + response=non_admin_response, + user_api_key_dict=non_admin_key, + model=None, + model_id=None, ) finally: for p in common_patches: @@ -1113,7 +1125,10 @@ async def test_health_endpoint_warns_when_scoped_models_lack_model_id(): patch("litellm.proxy.proxy_server.health_check_concurrency", 1), ): result = await health_endpoint( - response=Response(), user_api_key_dict=user_api_key_dict + response=Response(), + user_api_key_dict=user_api_key_dict, + model=None, + model_id=None, ) assert result["healthy_count"] == 0 @@ -1125,6 +1140,74 @@ async def test_health_endpoint_warns_when_scoped_models_lack_model_id(): assert any("model_info.id" in w for w in result["warnings"]) +@pytest.mark.asyncio +async def test_health_endpoint_503_for_targeted_unhealthy_model_under_background_cache_admin(): + """ + With background_health_checks enabled, an admin calling /health?model=foo + must get 503 when foo specifically has zero healthy endpoints — even if + other unrelated models in the cache are healthy. Without the cache-path + filter, the global healthy_count would mask the targeted failure. + """ + from fastapi import Response + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + full_model_list = [ + { + "model_name": "model-a", # the unhealthy target + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": "id-a"}, + }, + { + "model_name": "model-b", # an unrelated healthy model + "litellm_params": {"model": "openai/gpt-4o"}, + "model_info": {"id": "id-b"}, + }, + ] + + cached_results = { + "healthy_endpoints": [ + {"model": "openai/gpt-4o", "model_id": "id-b"}, + ], + "unhealthy_endpoints": [ + {"model": "openai/gpt-4o", "model_id": "id-a", "error": "boom"}, + ], + "healthy_count": 1, + "unhealthy_count": 1, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-admin", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + response = Response() + with ( + patch("litellm.proxy.proxy_server.llm_model_list", full_model_list), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch("litellm.proxy.proxy_server.use_background_health_checks", True), + patch("litellm.proxy.proxy_server.user_model", None), + patch("litellm.proxy.proxy_server.health_check_results", cached_results), + patch("litellm.proxy.proxy_server.health_check_details", True), + patch("litellm.proxy.proxy_server.health_check_concurrency", 1), + ): + result = await health_endpoint( + response=response, + user_api_key_dict=user_api_key_dict, + model="model-a", + model_id=None, + ) + + assert response.status_code == 503 + # Body must be scoped to the targeted model — not the global cache. + assert result["healthy_count"] == 0 + assert result["unhealthy_count"] == 1 + returned_ids = {ep["model_id"] for ep in result.get("unhealthy_endpoints", [])} + assert returned_ids == {"id-a"} + + @pytest.mark.asyncio async def test_health_endpoint_returns_503_when_requested_model_has_no_healthy_endpoints(): """