mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(health): apply model_info.health_check_params to health check probes
This commit is contained in:
parent
f005afa146
commit
d0dd24ed6d
3 changed files with 160 additions and 4 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue