This commit is contained in:
Chase Cai 2026-09-28 19:23:28 -04:00 • committed by GitHub
commit 5613c7c487
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 81 additions and 2 deletions

View file

@ -251,7 +251,9 @@ class HealthCheckHelpers:
),
"responses": lambda: litellm.aresponses(
**_filter_model_params(model_params=model_params),
input=prompt or "test",
input=[ # mutable-ok: Responses API requires a JSON message list
{"role": "user", "content": prompt or "test"}, # mutable-ok: JSON message object
],
),
"ocr": lambda: litellm.aocr(
**_filter_model_params(model_params=model_params),

View file

@ -69,6 +69,9 @@ class BaseConfig(ABC):
def __init__(self):
pass
def get_health_check_mode(self) -> str | None:
return None
@classmethod
def get_config(cls):
return {

View file

@ -26,6 +26,9 @@ class ChatGPTConfig(OpenAIConfig):
def api_base_without_login(self) -> str:
return self.authenticator.get_api_base()
def get_health_check_mode(self) -> str:
return "responses"
def _get_openai_compatible_provider_info(
self,
model: str,

View file

@ -8751,6 +8751,15 @@ async def ahealth_check(
api_base=api_base_from_params,
api_key=api_key_from_params,
)
if mode is None and custom_llm_provider in { # mutable-ok: short-lived provider-value membership set
provider.value for provider in LlmProviders
}:
provider_config = ProviderConfigManager.get_provider_chat_config(
model=model,
provider=LlmProviders(custom_llm_provider),
)
if provider_config is not None:
mode = provider_config.get_health_check_mode()
if model in litellm.model_cost and mode is None:
mode = litellm.model_cost[model].get("mode")

View file

@ -140,6 +140,68 @@ async def test_ahealth_check_supports_image_edit_mode():
assert "error" not in result
assert "Mode image_edit not supported" not in str(result)
@pytest.mark.asyncio
async def test_ahealth_check_uses_provider_mode_and_responses_message_input():
mock_response = MagicMock()
mock_response._hidden_params = {}
mock_aresponses = AsyncMock(return_value=mock_response)
mock_acompletion = AsyncMock()
with (
patch( # test-quality-ok: isolate provider and handler calls for focused health-check routing
"litellm.main.get_llm_provider",
return_value=("gpt-5.6-sol", "chatgpt", None, None),
),
patch.object( # test-quality-ok: isolate health-check metadata plumbing for routing assertion
HealthCheckHelpers,
"_update_model_params_with_health_check_tracking_information",
side_effect=lambda model_params: model_params,
),
patch("litellm.aresponses", new=mock_aresponses), # test-quality-ok: verify provider-mode dispatch
patch("litellm.acompletion", new=mock_acompletion), # test-quality-ok: verify chat fallback dispatch
):
result = await ahealth_check(
model_params={"model": "chatgpt/gpt-5.6-sol"},
mode=None,
prompt="health check",
)
assert "error" not in result
mock_aresponses.assert_awaited_once()
assert mock_aresponses.await_args.kwargs["input"] == [
{"role": "user", "content": "health check"}
]
mock_acompletion.assert_not_awaited()
@pytest.mark.asyncio
async def test_ahealth_check_provider_without_default_mode_keeps_chat():
mock_response = MagicMock()
mock_response._hidden_params = {}
mock_aresponses = AsyncMock()
mock_acompletion = AsyncMock(return_value=mock_response)
with (
patch( # test-quality-ok: isolate provider and handler calls for focused health-check routing
"litellm.main.get_llm_provider",
return_value=("custom-model", "custom", None, None),
),
patch.object( # test-quality-ok: isolate health-check metadata plumbing for routing assertion
HealthCheckHelpers,
"_update_model_params_with_health_check_tracking_information",
side_effect=lambda model_params: model_params,
),
patch("litellm.aresponses", new=mock_aresponses), # test-quality-ok: verify provider-mode dispatch
patch("litellm.acompletion", new=mock_acompletion), # test-quality-ok: verify chat fallback dispatch
):
result = await ahealth_check(
model_params={"model": "custom/custom-model"},
mode=None,
)
assert "error" not in result
mock_acompletion.assert_awaited_once()
mock_aresponses.assert_not_awaited()
def test_update_model_params_with_health_check_tracking_information():
@ -439,7 +501,7 @@ async def test_realtime_health_check_uses_model_level_vertex_params():
"websockets.connect",
lambda url, **kwargs: _FakeWebsocketConnect(connect_calls, url, **kwargs),
),
patch.object(
patch.object( # test-quality-ok: isolate health-check metadata plumbing for routing assertion
HealthCheckHelpers,
"_update_model_params_with_health_check_tracking_information",
staticmethod(lambda model_params: model_params),