diff --git a/docs/my-website/docs/proxy/health.md b/docs/my-website/docs/proxy/health.md index 6f98265e40a..2764a6f0d4f 100644 --- a/docs/my-website/docs/proxy/health.md +++ b/docs/my-website/docs/proxy/health.md @@ -330,6 +330,22 @@ model_list: health_check_timeout: 10 # 👈 OVERRIDE HEALTH CHECK TIMEOUT ``` +## Health Check Max Tokens + +By default, health checks use `max_tokens=1` to minimize cost and latency. For wildcard models, the default is `max_tokens=10`. + +You can override this per-model by setting `health_check_max_tokens` in the `model_info` section of your config.yaml. + +```yaml +model_list: + - model_name: openai/gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY + model_info: + health_check_max_tokens: 5 # 👈 OVERRIDE HEALTH CHECK MAX TOKENS +``` + ## `/health/readiness` Unprotected endpoint for checking if proxy is ready to accept requests diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index 9fb0c4126b3..9e972f1910b 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -14,7 +14,6 @@ TEST_PDF_URL = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9U class HealthCheckHelpers: - @staticmethod async def ahealth_check_wildcard_models( model: str, @@ -132,7 +131,7 @@ class HealthCheckHelpers: Callable, ]: """ - Returns a dictionary of mode handlers for health check calls. + Returns a dictionary of mode handlers for health check calls. Mode Handlers are Callables that need to be run for execution of the health check call. @@ -217,4 +216,4 @@ class HealthCheckHelpers: "document_url": TEST_PDF_URL, }, ), - } \ No newline at end of file + } diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 341ea4bd9e2..a8d0e3e9af2 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -329,7 +329,9 @@ async def perform_health_check( # Filter by model_id first so a single deployment is checked when id is specified if model_id is not None: - _by_id = [x for x in model_list if (x.get("model_info") or {}).get("id") == model_id] + _by_id = [ + x for x in model_list if (x.get("model_info") or {}).get("id") == model_id + ] if _by_id: model_list = _by_id elif model is not None: diff --git a/tests/test_litellm/proxy/test_health_check_max_tokens.py b/tests/test_litellm/proxy/test_health_check_max_tokens.py index da2bb21abea..e26f7fb9f20 100644 --- a/tests/test_litellm/proxy/test_health_check_max_tokens.py +++ b/tests/test_litellm/proxy/test_health_check_max_tokens.py @@ -1,9 +1,9 @@ - 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(): """ @@ -11,11 +11,12 @@ async def test_update_litellm_params_max_tokens_default(): """ 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(): """ @@ -23,11 +24,12 @@ async def test_update_litellm_params_max_tokens_custom(): """ 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(): """ @@ -35,37 +37,39 @@ async def test_update_litellm_params_max_tokens_wildcard(): """ 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: - + 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): # 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() + 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() + litellm_logging_obj=MagicMock(), ) assert model_params["max_tokens"] == 3