diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index 7ea9f9496d9..30e9b020245 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -44,17 +44,18 @@ def _strip_known_provider_prefix(model: str, known_providers: frozenset[str]) -> return model_suffix if sep and leading in known_providers else model -def _wildcard_health_check_models(wildcard_model: str, cheapest_models: Sequence[str]) -> tuple[str, ...]: - provider_prefix, sep, wildcard_suffix = wildcard_model.partition("/") - if not sep or wildcard_suffix == "*": +def _wildcard_health_check_models( + wildcard_suffix: str, custom_llm_provider: str, cheapest_models: Sequence[str] +) -> tuple[str, ...]: + if wildcard_suffix == "*": return tuple(cheapest_models) known_providers: Final = frozenset(provider.value for provider in LlmProviders) literal_prefix: Final = wildcard_suffix.replace("*", "") stripped_ids: Final = tuple(_strip_known_provider_prefix(model, known_providers) for model in cheapest_models) return tuple( - f"{provider_prefix}/{stripped}" + f"{custom_llm_provider}/{stripped}" if stripped.startswith(literal_prefix) - else f"{provider_prefix}/{wildcard_suffix.replace('*', stripped, 1)}" + else f"{custom_llm_provider}/{wildcard_suffix.replace('*', stripped, 1)}" for stripped in stripped_ids ) @@ -78,7 +79,9 @@ class HealthCheckHelpers: raise Exception( f"Unable to health check wildcard model for provider {custom_llm_provider}. Add a model on your config.yaml or contribute here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json" ) - candidates: Final = _wildcard_health_check_models(wildcard_model=model, cheapest_models=cheapest_models) + candidates: Final = _wildcard_health_check_models( + wildcard_suffix=model, custom_llm_provider=custom_llm_provider, cheapest_models=cheapest_models + ) fallback_models: Final = list(candidates[1:]) or None model_params["model"] = candidates[0] model_params["litellm_logging_obj"] = litellm_logging_obj diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index 2e010bdf293..1087d8be3ff 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -553,7 +553,8 @@ def test_ocr_health_check_document_raises_without_the_extension(): def test_wildcard_health_check_models_partial_prefix_substitutes_stripped_id(): candidates: Final = _wildcard_health_check_models( - "databricks/system.ai.*", + "system.ai.*", + "databricks", ["databricks/databricks-gemini-3-1-flash-image", "databricks/databricks-gpt-5"], ) assert candidates == ( @@ -562,15 +563,36 @@ def test_wildcard_health_check_models_partial_prefix_substitutes_stripped_id(): ) -def test_wildcard_health_check_models_bare_provider_and_star_wildcard_unchanged(): +def test_wildcard_health_check_models_star_wildcard_unchanged(): candidates: Final = ("databricks/databricks-gpt-5",) - assert _wildcard_health_check_models("databricks/*", candidates) == candidates - assert _wildcard_health_check_models("*", candidates) == candidates + assert _wildcard_health_check_models("*", "databricks", candidates) == candidates def test_wildcard_health_check_models_partial_prefix_matching_literal_keeps_stripped_id(): - assert _wildcard_health_check_models("openai/gpt-4*", ["gpt-4o-mini"]) == ("openai/gpt-4o-mini",) + assert _wildcard_health_check_models("gpt-4*", "openai", ["gpt-4o-mini"]) == ("openai/gpt-4o-mini",) def test_wildcard_health_check_models_partial_prefix_splices_suffix_around_star(): - assert _wildcard_health_check_models("openai/ft:*", ["gpt-4o-mini"]) == ("openai/ft:gpt-4o-mini",) + assert _wildcard_health_check_models("ft:*", "openai", ["gpt-4o-mini"]) == ("openai/ft:gpt-4o-mini",) + + +@pytest.mark.asyncio +async def test_ahealth_check_partial_prefix_wildcard_probes_spliced_model_id(): + from litellm.litellm_core_utils.llm_request_utils import pick_cheapest_chat_models_from_llm_provider + + cheapest: Final = pick_cheapest_chat_models_from_llm_provider(custom_llm_provider="databricks", n=3) + stripped: Final = tuple(model_id.partition("/")[-1] for model_id in cheapest) + + captured: dict[str, object] = {} + + async def _capture(**kwargs): + captured.update(kwargs) + + with patch("litellm.acompletion", side_effect=_capture): + await ahealth_check( + model_params={"model": "databricks/system.ai.*", "api_key": "x", "api_base": "https://example.invalid"}, + mode="chat", + ) + + assert captured["model"] == f"databricks/system.ai.{stripped[0]}" + assert captured["fallbacks"] == [f"databricks/system.ai.{model_id}" for model_id in stripped[1:]]