diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index a0f027cd58f..692435b3910 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -3,13 +3,13 @@ Helper functions for health check calls. """ import base64 -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, Sequence from typing import TYPE_CHECKING, Final, Literal from litellm.llms.base_llm.ocr.transformation import DocumentType from litellm.rust_bridge import runtime from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR_HEALTH_CHECK_DOCUMENT -from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS +from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS, LlmProviders if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging @@ -39,6 +39,25 @@ def _ocr_health_check_document(model: str, custom_llm_provider: str) -> Document return native(model, custom_llm_provider) +def _strip_known_provider_prefix(model: str, known_providers: frozenset[str]) -> str: + leading, sep, model_suffix = model.partition("/") + return model_suffix if sep and leading in known_providers else model + + +def _wildcard_health_check_models( + wildcard_suffix: str, custom_llm_provider: str, candidate_models: Sequence[str] +) -> tuple[str, ...]: + if wildcard_suffix == "*": + return tuple(candidate_models[:3]) + 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 candidate_models) + matching: Final = [stripped for stripped in stripped_ids if stripped.startswith(literal_prefix)] + if matching: + return tuple(f"{custom_llm_provider}/{stripped}" for stripped in matching[:3]) + return tuple(f"{custom_llm_provider}/{wildcard_suffix.replace('*', stripped, 1)}" for stripped in stripped_ids[:3]) + + class HealthCheckHelpers: @staticmethod async def ahealth_check_wildcard_models( @@ -53,16 +72,16 @@ class HealthCheckHelpers: ) # this is a wildcard model, we need to pick a random model from the provider - cheapest_models = pick_cheapest_chat_models_from_llm_provider(custom_llm_provider=custom_llm_provider, n=3) + cheapest_models = pick_cheapest_chat_models_from_llm_provider(custom_llm_provider=custom_llm_provider, n=10_000) if len(cheapest_models) == 0: 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" ) - if len(cheapest_models) > 1: - fallback_models = cheapest_models[1:] # Pick the last 2 models from the shuffled list - else: - fallback_models = None - model_params["model"] = cheapest_models[0] + candidates: Final = _wildcard_health_check_models( + wildcard_suffix=model, custom_llm_provider=custom_llm_provider, candidate_models=cheapest_models + ) + fallback_models: Final = list(candidates[1:]) or None + model_params["model"] = candidates[0] model_params["litellm_logging_obj"] = litellm_logging_obj model_params["fallbacks"] = fallback_models model_params["max_tokens"] = model_params.get("max_tokens", 16) # GPT-5 models require max_output_tokens >= 16 diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index de2ca4762f1..944ed80280a 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -262,6 +262,24 @@ def _hydrate_litellm_credential_name( return litellm_params +def _strip_known_provider_prefix(model: str, known_providers: frozenset[str]) -> str: + leading, sep, model_suffix = model.partition("/") + return model_suffix if sep and leading in known_providers else model + + +def _apply_partial_wildcard_prefix( + provider_models: list[str], model_prefix: str, known_providers: frozenset[str] +) -> list[str]: + prefix_filtered: Final = [ + model + for model in provider_models + if _strip_known_provider_prefix(model, known_providers).startswith(model_prefix) + ] + if prefix_filtered: + return prefix_filtered + return [f"{model_prefix}{_strip_known_provider_prefix(model, known_providers)}" for model in provider_models] + + def get_known_models_from_wildcard(wildcard_model: str, litellm_params: LiteLLM_Params | None = None) -> list[str]: wildcard_model_to_expand: Final = ( litellm_params.model @@ -288,23 +306,24 @@ def get_known_models_from_wildcard(wildcard_model: str, litellm_params: LiteLLM_ litellm_params = _hydrate_litellm_credential_name(litellm_params) - wildcard_models = get_provider_models(provider=provider, litellm_params=litellm_params) + provider_models: Final = get_provider_models(provider=provider, litellm_params=litellm_params) - if wildcard_models is None: + if provider_models is None: return [] - if wildcard_suffix != "*": - ## CHECK IF PARTIAL FILTER e.g. `gemini-*` - model_prefix: Final = wildcard_suffix.replace("*", "") - is_partial_filter: Final = any(wc_model.startswith(model_prefix) for wc_model in wildcard_models) - if is_partial_filter: - filtered_wildcard_models = [wc_model for wc_model in wildcard_models if wc_model.startswith(model_prefix)] - wildcard_models = filtered_wildcard_models - else: - # add model prefix to wildcard models - wildcard_models = [f"{model_prefix}{model}" for model in wildcard_models] + known_providers: Final = frozenset(provider.value for provider in LlmProviders) + + ## CHECK IF PARTIAL FILTER e.g. `gemini-*` + wildcard_models: Final = ( + provider_models + if wildcard_suffix == "*" + else _apply_partial_wildcard_prefix( + provider_models=provider_models, + model_prefix=wildcard_suffix.replace("*", ""), + known_providers=known_providers, + ) + ) - known_providers: Final = {provider.value for provider in LlmProviders} suffix_appended_wildcard_models: Final = [] for model in wildcard_models: if not model.startswith(wildcard_provider_prefix): @@ -314,11 +333,7 @@ def get_known_models_from_wildcard(wildcard_model: str, litellm_params: LiteLLM_ # both, which would otherwise yield an uncallable "ollama_server1/ollama/gemma3:1b". # Only strip the leading segment when it is a known provider, so ids whose first # segment is an org rather than a provider (e.g. "meta-llama/Llama-3-8B") keep it. - leading, sep, model_suffix = model.partition("/") - if sep and leading in known_providers: - model = f"{wildcard_provider_prefix}/{model_suffix}" - else: - model = f"{wildcard_provider_prefix}/{model}" + model = f"{wildcard_provider_prefix}/{_strip_known_provider_prefix(model, known_providers)}" suffix_appended_wildcard_models.append(model) return suffix_appended_wildcard_models or [] diff --git a/tests/integration/routing/test_wildcard_partial_prefix_expansion.py b/tests/integration/routing/test_wildcard_partial_prefix_expansion.py new file mode 100644 index 00000000000..01206a2163a --- /dev/null +++ b/tests/integration/routing/test_wildcard_partial_prefix_expansion.py @@ -0,0 +1,71 @@ +"""Partial-prefix wildcards like ``databricks/system.ai.*`` must splice the literal +prefix onto provider-stripped cost-map ids, never onto the provider-prefixed key.""" + +from collections.abc import Iterator +from typing import Final + +import pytest + +from integration._support.client import Gateway, object_value, string_value + +PATTERNS: Final = ("databricks/system.ai.*", "databricks/*") + + +def _registered_model_id(gateway: Gateway, pattern: str) -> str: + created: Final = object_value( + gateway.post( + "/model/new", + { + "model_name": pattern, + "litellm_params": { + "model": pattern, + "api_key": "integration-wildcard-key", + "api_base": "https://example.invalid", + }, + }, + ) + ) + return string_value(object_value(created["model_info"])["id"]) + + +@pytest.fixture +def wildcard_deployments(gateway: Gateway) -> Iterator[None]: + with gateway.scenario() as scenario: + for pattern in PATTERNS: + scenario.cleanups.callback(scenario.delete_model, _registered_model_id(gateway, pattern)) + yield + + +def _listed_model_ids(gateway: Gateway) -> tuple[str, ...]: + data: Final = gateway.get("/v1/models")["data"] + assert isinstance(data, list) + return tuple(string_value(object_value(entry)["id"]) for entry in data) + + +def test_partial_prefix_wildcard_expands_to_unity_catalog_names(gateway: Gateway, wildcard_deployments: None) -> None: + prefixed: Final = tuple( + model_id for model_id in _listed_model_ids(gateway) if model_id.startswith("databricks/system.ai.") + ) + assert prefixed + assert all("/" not in model_id.removeprefix("databricks/system.ai.") for model_id in prefixed) + + +def test_provider_wildcard_still_expands_to_cost_map_names(gateway: Gateway, wildcard_deployments: None) -> None: + assert any(model_id.startswith("databricks/databricks-") for model_id in _listed_model_ids(gateway)) + + +def test_model_info_keeps_provider_model_for_expanded_deployments(gateway: Gateway, wildcard_deployments: None) -> None: + data: Final = gateway.get("/model/info")["data"] + assert isinstance(data, list) + expanded: Final = [ + row + for row in (object_value(entry) for entry in data) + if string_value(row["model_name"]).startswith("databricks/system.ai.") + ] + assert expanded, "model/info returned no expanded rows for databricks/system.ai.*" + bad: Final = [ + row["model_name"] + for row in expanded + if "system.ai.databricks/" in string_value(object_value(row["litellm_params"])["model"]) + ] + assert not bad, f"litellm_params.model carries the corrupted expanded name: {bad}" diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index 13171a42cda..866a8f42bdf 100644 --- a/tests/test_litellm/proxy/auth/test_model_checks.py +++ b/tests/test_litellm/proxy/auth/test_model_checks.py @@ -1,3 +1,4 @@ +from typing import Final from unittest.mock import patch import pytest @@ -485,6 +486,36 @@ def test_wildcard_custom_prefix_keeps_org_segment_for_non_provider_first_segment assert result == ["my_hf/meta-llama/Llama-3-8B"] +def test_partial_prefix_wildcard_expands_against_stripped_provider_ids(monkeypatch): + import litellm + from litellm.proxy.auth import model_checks + from litellm.proxy.auth.model_checks import get_known_models_from_wildcard + from litellm.types.router import LiteLLM_Params + + provider_ids: Final = list(litellm.models_by_provider["databricks"]) + monkeypatch.setattr( + model_checks, + "get_provider_models", + lambda provider, litellm_params=None: provider_ids, + ) + + result: Final = get_known_models_from_wildcard( + wildcard_model="databricks/system.ai.*", + litellm_params=LiteLLM_Params( + model="databricks/system.ai.*", + api_key="x", + api_base="https://example.invalid", + ), + ) + + expected: Final = [f"databricks/system.ai.{model_id.partition('/')[-1]}" for model_id in provider_ids] + assert result == expected + assert all( + model_id.startswith("databricks/system.ai.") and "/" not in model_id.removeprefix("databricks/system.ai.") + for model_id in result + ) + + def test_wildcard_credential_hydration_preserves_missing_credential_name( monkeypatch, ): 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 47c4576f91f..e5bd43369ec 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -14,6 +14,7 @@ from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME from litellm.litellm_core_utils.health_check_helpers import ( IMAGE_EDIT_HEALTH_CHECK_PROMPT, HealthCheckHelpers, + _wildcard_health_check_models, ) from litellm.main import ahealth_check from litellm.proxy._types import UserAPIKeyAuth @@ -548,3 +549,60 @@ def test_ocr_health_check_document_raises_without_the_extension(): _ocr_health_check_document(model="mistral/mistral-ocr-latest", custom_llm_provider="mistral") finally: NATIVE_OCR_HEALTH_CHECK_DOCUMENT.reset() + + +def test_wildcard_health_check_models_partial_prefix_substitutes_stripped_id(): + candidates: Final = _wildcard_health_check_models( + "system.ai.*", + "databricks", + ["databricks/databricks-gemini-3-1-flash-image", "databricks/databricks-gpt-5"], + ) + assert candidates == ( + "databricks/system.ai.databricks-gemini-3-1-flash-image", + "databricks/system.ai.databricks-gpt-5", + ) + + +def test_wildcard_health_check_models_star_wildcard_unchanged(): + candidates: Final = ("databricks/a", "databricks/b", "databricks/c", "databricks/d") + assert _wildcard_health_check_models("*", "databricks", candidates) == ( + "databricks/a", + "databricks/b", + "databricks/c", + ) + + +def test_wildcard_health_check_models_partial_prefix_matching_literal_keeps_stripped_id(): + assert _wildcard_health_check_models("gpt-4*", "openai", ["gpt-4o-mini"]) == ("openai/gpt-4o-mini",) + + +def test_wildcard_health_check_models_partial_prefix_filters_candidates_to_literal_prefix(): + assert _wildcard_health_check_models( + "gpt-4*", "openai", ["gpt-5-nano", "gpt-4o-mini", "gpt-4.1-mini", "gpt-4o"] + ) == ("openai/gpt-4o-mini", "openai/gpt-4.1-mini", "openai/gpt-4o") + + +def test_wildcard_health_check_models_partial_prefix_splices_suffix_around_star(): + 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:]]