diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index 47a27c8ef5b..9fb0c4126b3 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -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 {} diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index d228bdb2129..341ea4bd9e2 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -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 diff --git a/tests/test_litellm/proxy/test_health_check_max_tokens.py b/tests/test_litellm/proxy/test_health_check_max_tokens.py new file mode 100644 index 00000000000..da2bb21abea --- /dev/null +++ b/tests/test_litellm/proxy/test_health_check_max_tokens.py @@ -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