diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 7d67750c78f..62721799662 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -366,10 +366,22 @@ def _update_litellm_params_for_health_check( - updates the `voice` param with the `health_check_voice` for `audio_speech` mode if it exists Doc: https://docs.litellm.ai/docs/proxy/health#text-to-speech-models - for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID """ + _NON_CHAT_MODES = { + "image_generation", + "video_generation", + "embedding", + "rerank", + "transcription", + "audio_speech", + } + _mode = model_info.get("mode", None) litellm_params["messages"] = _get_random_llm_message() - _resolved_max_tokens = _resolve_health_check_max_tokens(model_info, litellm_params) - if _resolved_max_tokens is not None: - litellm_params["max_tokens"] = _resolved_max_tokens + if _mode not in _NON_CHAT_MODES: + _resolved_max_tokens = _resolve_health_check_max_tokens( + model_info, litellm_params + ) + if _resolved_max_tokens is not None: + litellm_params["max_tokens"] = _resolved_max_tokens _health_check_model = model_info.get("health_check_model", None) if _health_check_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 09211b72c3e..2d2b4eac9fe 100644 --- a/tests/test_litellm/proxy/test_health_check_max_tokens.py +++ b/tests/test_litellm/proxy/test_health_check_max_tokens.py @@ -86,6 +86,37 @@ async def test_ahealth_check_wildcard_models_respects_max_tokens(): assert model_params["max_tokens"] == 3 +@pytest.mark.parametrize( + "mode", + [ + "image_generation", + "video_generation", + "embedding", + "rerank", + "transcription", + "audio_speech", + ], +) +@pytest.mark.asyncio +async def test_update_litellm_params_no_max_tokens_for_non_chat_modes( + monkeypatch, mode +): + """ + Non-chat modes (image_generation, embedding, etc.) must not receive max_tokens + because those endpoints reject it with a 400 error. + """ + monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", None) + monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING", None) + model_info = {"mode": mode} + litellm_params = {"model": "openai/dall-e-3"} + + updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert ( + "max_tokens" not in updated_params + ), f"max_tokens should not be set for mode={mode!r}, got {updated_params.get('max_tokens')}" + + @pytest.mark.asyncio async def test_background_health_check_max_tokens_env_var(monkeypatch): """