mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 54426c70b0 into 461a58c40a
This commit is contained in:
commit
cc90206dc7
1 changed files with 27 additions and 2 deletions
|
|
@ -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] = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue