fix(health): attribute health check results by model_id first

_aggregate_health_check_results was iterating over all deployments sharing a
model name via model_param_to_info[model_param], causing cross-attribution
when multiple deployments (different api_base, or model-group aliases) share
the same litellm_params.model.

Added _get_model_infos_for_endpoint() helper that matches by model_id first
(which each endpoint result carries from _perform_health_check), falling back
to the model-param mapping only when no id is available.

Fixes #44154
This commit is contained in:
Suhas Hanamannavar 2026-10-03 05:43:41 +02:00
parent f9a32ffcb5
commit 93ee1864ad

View file

@ -708,6 +708,25 @@ def _build_model_param_to_info_mapping(model_list: list) -> dict:
return model_param_to_info
def _get_model_infos_for_endpoint(model_param_to_info: dict, endpoint: dict) -> list:
"""
Return the model infos a health check endpoint result belongs to.
`model_param_to_info` is keyed by `litellm_params.model`, which is not unique:
several deployments (different api_base, or model-group aliases) can share it.
Each endpoint result carries the `model_id` of the deployment it was produced
for, so prefer an exact match on that and fall back to the model-param match
only when no id is available.
"""
model_infos = model_param_to_info.get(endpoint.get("model"), [])
endpoint_model_id = endpoint.get("model_id")
if endpoint_model_id:
matching = [info for info in model_infos if info.get("model_id") == endpoint_model_id]
if matching:
return matching
return model_infos
def _aggregate_health_check_results(
model_param_to_info: dict,
healthy_endpoints: list,
@ -732,7 +751,7 @@ def _aggregate_health_check_results(
for endpoint in healthy_endpoints:
model_param = endpoint.get("model")
if model_param and model_param in model_param_to_info:
for model_info in model_param_to_info[model_param]:
for model_info in _get_model_infos_for_endpoint(model_param_to_info, endpoint):
key = (model_info["model_id"], model_info["model_name"])
if key not in model_results:
model_results[key] = {
@ -749,7 +768,7 @@ def _aggregate_health_check_results(
model_param = endpoint.get("model")
error_message = endpoint.get("error")
if model_param and model_param in model_param_to_info:
for model_info in model_param_to_info[model_param]:
for model_info in _get_model_infos_for_endpoint(model_param_to_info, endpoint):
key = (model_info["model_id"], model_info["model_name"])
if key not in model_results:
model_results[key] = {