mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
7635955c91
commit
3340533cfb
2 changed files with 141 additions and 7 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue