fix(health): apply model_info.health_check_params to health check probes

This commit is contained in:
mateo-berri 2026-08-24 11:01:34 -07:00
parent f005afa146
commit d0dd24ed6d
3 changed files with 160 additions and 4 deletions

View file

@ -445,6 +445,9 @@ def _update_litellm_params_for_health_check(model_info: dict, litellm_params: di
"""
Update the litellm params for health check.
- merges `model_info.health_check_params` into the probe request, so a deployment whose provider
requires a payload field litellm does not synthesize (e.g. `mediaSource` for Bedrock TwelveLabs
Pegasus) can supply it. The dedicated knobs below are applied afterwards and win on conflict.
- gets a short `messages` param for health check
- adds a bounded `max_tokens` when the deployment is a chat-style mode
(`chat`, `completion`, `responses`) or the operator explicitly opts in
@ -459,6 +462,16 @@ def _update_litellm_params_for_health_check(model_info: dict, litellm_params: di
model_info,
litellm_params, # any-ok: untyped router config dict
)
_health_check_params: Final = model_info.get("health_check_params", None)
if isinstance(_health_check_params, dict):
litellm_params.update(_health_check_params)
elif _health_check_params is not None:
logger.warning(
"health_check_params for model %s is a %s, expected a dict. Ignoring it.",
litellm_params.get("model"),
type(_health_check_params).__name__,
)
litellm_params["messages"] = _get_random_llm_message()
if _should_inject_health_check_max_tokens(
model_info,

View file

@ -1888,6 +1888,8 @@ async def test_model_connection(
# already resolved before reaching this endpoint; any remaining
# reference must have come from the request body.
_reject_os_environ_references(request_litellm_params)
if model_info:
_reject_os_environ_references(model_info)
model_name: Final = request_litellm_params.get("model")
# Look up model configuration from router if model name is provided
@ -1951,20 +1953,19 @@ async def test_model_connection(
}
## Auth check
auth_model_info: Final = loaded_model_info if loaded_model_info is not None else model_info
resolved_model_info: Final = loaded_model_info if loaded_model_info is not None else model_info
await ModelManagementAuthChecks.can_user_make_model_call(
model_params=Deployment(
model_name="test_model",
litellm_params=LiteLLM_Params(**litellm_params),
model_info=auth_model_info,
model_info=resolved_model_info,
),
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
)
# Include health_check_params if provided
litellm_params = _update_litellm_params_for_health_check(
model_info={},
model_info=resolved_model_info or {},
litellm_params=litellm_params,
)
mode = mode or litellm_params.pop("mode", None)

View file

@ -1,7 +1,11 @@
import json
import logging
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import respx
import litellm
from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers
from litellm.proxy import health_check as hc_module
from litellm.proxy.health_check import (
@ -543,3 +547,141 @@ async def test_run_model_health_check_skips_auto_router_deployment():
fake_ahealth_check.assert_not_called()
assert result == {}
# ---------------------------------------------------------------------------
# model_info.health_check_params
#
# Some providers require a payload field litellm does not synthesize for a
# probe. Bedrock TwelveLabs Pegasus rejects any Invoke body without a top-level
# `mediaSource`, so every health check on such a deployment failed with
# "Invalid JSON: $: required property 'mediaSource' not found". The config key
# was accepted and then never read, so operators had no way to supply it.
# ---------------------------------------------------------------------------
def test_health_check_params_merge_into_probe_params():
"""health_check_params reach the probe request for the deployment that declares them."""
media_source = {"s3Location": {"uri": "s3://my-bucket/clip.mp4"}}
updated = _update_litellm_params_for_health_check(
{"mode": "chat", "health_check_params": {"mediaSource": media_source}},
{"model": "bedrock/us.twelvelabs.pegasus-1-2-v1:0"},
)
assert updated["mediaSource"] == media_source
assert updated["model"] == "us.twelvelabs.pegasus-1-2-v1:0"
assert updated["custom_llm_provider"] == "bedrock"
def test_health_check_params_lose_to_dedicated_health_check_knobs():
"""The dedicated knobs are applied after the merge, so they win on conflict."""
model_info = {
"mode": "chat",
"health_check_params": {
"max_tokens": 4096,
"model": "openai/expensive-model",
"messages": [{"role": "user", "content": "from health_check_params"}],
"reasoning_effort": "high",
},
"health_check_max_tokens": 5,
"health_check_model": "openai/cheap-model",
"health_check_reasoning_effort": "none",
}
updated = _update_litellm_params_for_health_check(model_info, {"model": "openai/dummy"})
assert updated["max_tokens"] == 5
assert updated["model"] == "openai/cheap-model"
assert updated["reasoning_effort"] == "none"
assert updated["messages"] != model_info["health_check_params"]["messages"]
def test_health_check_params_lose_to_the_audio_speech_voice_knob():
"""health_check_voice still wins for audio_speech deployments."""
updated = _update_litellm_params_for_health_check(
{
"mode": "audio_speech",
"health_check_params": {"voice": "sage", "response_format": "wav"},
"health_check_voice": "shimmer",
},
{"model": "openai/tts-1"},
)
assert updated["voice"] == "shimmer"
assert updated["response_format"] == "wav"
@pytest.mark.parametrize(
"bad_value",
["mediaSource", ["mediaSource"], 5, True],
)
def test_health_check_params_ignored_when_not_a_dict(bad_value, caplog):
"""A misconfigured health_check_params is skipped with a warning instead of breaking the probe."""
with caplog.at_level(logging.WARNING, logger="litellm.proxy.health_check"):
updated = _update_litellm_params_for_health_check(
{"mode": "chat", "health_check_params": bad_value},
{"model": "openai/dummy"},
)
assert updated["model"] == "openai/dummy"
assert updated["max_tokens"] == 16
assert "health_check_params" in caplog.text
def test_health_check_params_apply_to_non_chat_modes():
"""Non-chat probes get health_check_params too, and still no max_tokens."""
updated = _update_litellm_params_for_health_check(
{"mode": "embedding", "health_check_params": {"dimensions": 8}},
{"model": "bedrock/amazon.titan-embed-text-v2:0"},
)
assert updated["dimensions"] == 8
assert "max_tokens" not in updated
async def _pegasus_health_check_request_body(model_info: dict, monkeypatch) -> dict:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
litellm_params = _update_litellm_params_for_health_check(
model_info,
{
"model": "bedrock/us.twelvelabs.pegasus-1-2-v1:0",
"aws_access_key_id": "fake-access-key",
"aws_secret_access_key": "fake-secret-key",
"aws_region_name": "us-east-1",
},
)
with respx.mock(assert_all_called=True) as respx_mock:
invoke_route = respx_mock.post(
host="bedrock-runtime.us-east-1.amazonaws.com",
path__regex=r"/model/.+/invoke",
).respond(json={"message": "a person walks a dog", "finishReason": "stop"})
result = await litellm.ahealth_check(litellm_params, mode="chat")
assert "error" not in result, result
return json.loads(invoke_route.calls.last.request.content)
@pytest.mark.asyncio
async def test_health_check_params_reach_the_bedrock_invoke_body(monkeypatch):
"""The probe Bedrock actually receives carries mediaSource, which is what unblocks Pegasus."""
media_source = {"s3Location": {"uri": "s3://my-bucket/clip.mp4"}}
body = await _pegasus_health_check_request_body(
{"mode": "chat", "health_check_params": {"mediaSource": media_source}}, monkeypatch
)
assert body["mediaSource"] == media_source
assert body["maxOutputTokens"] == 16
assert body["inputPrompt"]
@pytest.mark.asyncio
async def test_bedrock_invoke_body_has_no_media_source_without_health_check_params(monkeypatch):
"""Negative control: the field only appears because the deployment asked for it."""
body = await _pegasus_health_check_request_body({"mode": "chat"}, monkeypatch)
assert "mediaSource" not in body