This commit is contained in:
devin-ai-integration[bot] 2026-09-30 16:56:32 -04:00 • committed by GitHub
commit 72c8eb44ef
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 220 additions and 26 deletions

View file

@ -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

View file

@ -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 []

View file

@ -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}"

View file

@ -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,
):

View file

@ -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:]]