fix(health): filter background-cache result by targeted model before 503 check

When use_background_health_checks is enabled, /health?model=foo returned
the full cached aggregate across every model — so an unhealthy foo
combined with any other healthy deployment kept healthy_count > 0 and
the targeted-503 path never fired.

Resolve the targeted model/model_id to a deployment-id set first
(mirroring perform_health_check's match-on-model_name-or-litellm_model
semantics) and narrow the cache to those IDs before _post_process
evaluates healthy_count, so the 503 contract holds for both the live
and cache code paths.
This commit is contained in:
Ryan Crabbe 2026-05-01 13:47:16 -07:00
parent 7635955c91
commit 3340533cfb
No known key found for this signature in database
2 changed files with 141 additions and 7 deletions

View file

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

View file

@ -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():
"""