mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat: health check max tokens
This commit is contained in:
parent
c58aea4888
commit
2553698da5
3 changed files with 82 additions and 1 deletions
|
|
@ -44,7 +44,9 @@ class HealthCheckHelpers:
|
|||
model_params["model"] = cheapest_models[0]
|
||||
model_params["litellm_logging_obj"] = litellm_logging_obj
|
||||
model_params["fallbacks"] = fallback_models
|
||||
model_params["max_tokens"] = 10 # gpt-5-nano throws errors for max_tokens=1
|
||||
model_params["max_tokens"] = model_params.get(
|
||||
"max_tokens", 10
|
||||
) # gpt-5-nano throws errors for max_tokens=1
|
||||
await acompletion(**model_params)
|
||||
return {}
|
||||
|
||||
|
|
|
|||
|
|
@ -234,6 +234,14 @@ def _update_litellm_params_for_health_check(
|
|||
- for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID
|
||||
"""
|
||||
litellm_params["messages"] = _get_random_llm_message()
|
||||
_health_check_max_tokens = model_info.get("health_check_max_tokens", None)
|
||||
if _health_check_max_tokens is not None:
|
||||
litellm_params["max_tokens"] = _health_check_max_tokens
|
||||
elif "*" not in (
|
||||
model_info.get("health_check_model") or litellm_params.get("model") or ""
|
||||
):
|
||||
litellm_params["max_tokens"] = 1
|
||||
|
||||
_health_check_model = model_info.get("health_check_model", None)
|
||||
if _health_check_model is not None:
|
||||
litellm_params["model"] = _health_check_model
|
||||
|
|
|
|||
71
tests/test_litellm/proxy/test_health_check_max_tokens.py
Normal file
71
tests/test_litellm/proxy/test_health_check_max_tokens.py
Normal file
|
|
@ -0,0 +1,71 @@
|
|||
|
||||
import pytest
|
||||
from litellm.proxy.health_check import _update_litellm_params_for_health_check
|
||||
from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_litellm_params_max_tokens_default():
|
||||
"""
|
||||
Test that max_tokens defaults to 1 for non-wildcard models.
|
||||
"""
|
||||
model_info = {}
|
||||
litellm_params = {"model": "gpt-4"}
|
||||
|
||||
updated_params = _update_litellm_params_for_health_check(model_info, litellm_params)
|
||||
|
||||
assert updated_params["max_tokens"] == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_litellm_params_max_tokens_custom():
|
||||
"""
|
||||
Test that max_tokens respects health_check_max_tokens from model_info.
|
||||
"""
|
||||
model_info = {"health_check_max_tokens": 5}
|
||||
litellm_params = {"model": "gpt-4"}
|
||||
|
||||
updated_params = _update_litellm_params_for_health_check(model_info, litellm_params)
|
||||
|
||||
assert updated_params["max_tokens"] == 5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_litellm_params_max_tokens_wildcard():
|
||||
"""
|
||||
Test that max_tokens does NOT default to 1 for wildcard models.
|
||||
"""
|
||||
model_info = {}
|
||||
litellm_params = {"model": "openai/*"}
|
||||
|
||||
updated_params = _update_litellm_params_for_health_check(model_info, litellm_params)
|
||||
|
||||
# Should not be set to 1
|
||||
assert "max_tokens" not in updated_params or updated_params["max_tokens"] != 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ahealth_check_wildcard_models_respects_max_tokens():
|
||||
"""
|
||||
Test that ahealth_check_wildcard_models respects max_tokens if passed,
|
||||
otherwise defaults to 10.
|
||||
"""
|
||||
with patch("litellm.litellm_core_utils.llm_request_utils.pick_cheapest_chat_models_from_llm_provider", return_value=["gpt-4o-mini"]), \
|
||||
patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
|
||||
|
||||
# Test Case 1: No max_tokens passed, should default to 10
|
||||
model_params = {}
|
||||
await HealthCheckHelpers.ahealth_check_wildcard_models(
|
||||
model="openai/*",
|
||||
custom_llm_provider="openai",
|
||||
model_params=model_params,
|
||||
litellm_logging_obj=MagicMock()
|
||||
)
|
||||
assert model_params["max_tokens"] == 10
|
||||
|
||||
# Test Case 2: Custom health_check_max_tokens passed via model_params, should be respected
|
||||
model_params = {"max_tokens": 3}
|
||||
await HealthCheckHelpers.ahealth_check_wildcard_models(
|
||||
model="openai/*",
|
||||
custom_llm_provider="openai",
|
||||
model_params=model_params,
|
||||
litellm_logging_obj=MagicMock()
|
||||
)
|
||||
assert model_params["max_tokens"] == 3
|
||||
Loading…
Add table
Reference in a new issue