mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge f5735591d7 into f285229b51
This commit is contained in:
commit
72c8eb44ef
5 changed files with 220 additions and 26 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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:]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue