diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 00f0ef1756f..4ce4be01bda 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -730,6 +730,31 @@ 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: Mapping[str, list[Mapping[str, Any]]], + endpoint: Mapping[str, Any], +) -> list[Mapping[str, Any]]: + """ + 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. When an + id is present but matches nothing locally, the result is for a deployment + this proxy does not know about, so return an empty list instead of falling + back to the broad model-name match (which would reintroduce the original + cross-attribution bug). Fall back to the model-param mapping only when no + id is available at all. + """ + 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] + return matching + return model_infos + + def _aggregate_health_check_results( model_param_to_info: dict, healthy_endpoints: list, @@ -754,7 +779,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] = { @@ -771,7 +796,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] = {