fix(proxy): expand partial wildcard prefixes against provider-stripped model ids

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-28 23:50:24 +00:00
parent eae8ed7f3c
commit 99636dff2b
5 changed files with 260 additions and 45 deletions

View file

@ -3,7 +3,7 @@ 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
@ -39,6 +39,19 @@ def _ocr_health_check_document(model: str, custom_llm_provider: str) -> Document
return native(model, custom_llm_provider)
def _wildcard_health_check_models(wildcard_model: str, cheapest_models: Sequence[str]) -> tuple[str, ...]:
"""Substitute each candidate id into the wildcard's dynamic segment so the health
check probes a name the wildcard route actually expands, instead of the bare
cost-map id (e.g. `databricks/system.ai.*` + `databricks/databricks-x` ->
`databricks/system.ai.databricks-x`)."""
provider_prefix, sep, wildcard_suffix = wildcard_model.partition("/")
if not sep or wildcard_suffix == "*":
return tuple(cheapest_models)
return tuple(
f"{provider_prefix}/{wildcard_suffix.replace('*', model.partition('/')[-1], 1)}" for model in cheapest_models
)
class HealthCheckHelpers:
@staticmethod
async def ahealth_check_wildcard_models(
@ -58,11 +71,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"
)
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_model=model, 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
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,149 @@
"""Regression: wildcard expansion of a partial provider prefix like
`databricks/system.ai.*` must produce `databricks/system.ai.<endpoint>` ids,
not `databricks/system.ai.databricks/databricks-<name>`.
On main, get_known_models_from_wildcard prepends the literal prefix before the
`*` ("system.ai.") to cost-map keys that already carry the `databricks/`
provider prefix, so /v1/models and /model/info list names that Databricks
rejects with "Invalid Unity Catalog name". See LIT-8910.
The test boots its own proxy (the shared stack's model_list is fixed), with two
wildcard deployments: `databricks/system.ai.*` under test and `databricks/*` as
a control whose expansion must stay `databricks/databricks-*`.
"""
import socket
import subprocess
import sys
import time
from collections.abc import Generator
from pathlib import Path
from typing import Final, cast
import pytest
from e2e_http import URL, AuthHeaders, NoBody, get, probe, unwrap
from models import ModelsListParams, ModelsListResponse
from pydantic import BaseModel, ConfigDict
REPO_ROOT: Final = Path(__file__).resolve().parents[2]
MASTER_KEY: Final = "sk-1234"
AUTH: Final = AuthHeaders(authorization=f"Bearer {MASTER_KEY}")
CONFIG_YAML: Final = """\
model_list:
- model_name: databricks/system.ai.*
litellm_params:
model: databricks/system.ai.*
api_key: fake-databricks-token
api_base: https://example.invalid
- model_name: databricks/*
litellm_params:
model: databricks/*
api_key: fake-databricks-token
api_base: https://example.invalid
general_settings:
master_key: sk-1234
"""
class _ModelInfoParams(BaseModel):
model_config = ConfigDict(protected_namespaces=(), extra="ignore")
model: str = ""
class _ModelInfoRow(BaseModel):
model_name: str
litellm_params: _ModelInfoParams = _ModelInfoParams()
class _ModelInfoResponse(BaseModel):
data: list[_ModelInfoRow] = []
def _free_port() -> int:
with socket.socket() as sock:
sock.bind(("127.0.0.1", 0))
return cast("tuple[str, int]", sock.getsockname())[1]
@pytest.fixture(scope="module")
def wildcard_proxy(tmp_path_factory: pytest.TempPathFactory) -> Generator[str]:
config_path = tmp_path_factory.mktemp("wildcard") / "config.yaml"
config_path.write_text(CONFIG_YAML)
port = _free_port()
proc = subprocess.Popen(
[
sys.executable,
str(REPO_ROOT / "litellm/proxy/proxy_cli.py"),
"--config",
str(config_path),
"--port",
str(port),
],
cwd=REPO_ROOT,
env={
"PATH": "/usr/bin:/bin:/usr/local/bin",
"HOME": str(Path.home()),
"PYTHONPATH": str(REPO_ROOT),
"LITELLM_LOCAL_MODEL_COST_MAP": "true",
"LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY": "true",
"LITELLM_MASTER_KEY": MASTER_KEY,
},
)
base = f"http://127.0.0.1:{port}"
try:
deadline = time.time() + 180
while True:
alive = probe(URL(f"{base}/health/liveliness"), headers=NoBody(), params=NoBody(), timeout=5.0)
if alive.healthy:
break
if proc.poll() is not None:
pytest.fail(f"proxy exited during boot with {proc.returncode}")
assert time.time() < deadline, "proxy did not answer /health/liveliness within 180s"
time.sleep(2)
yield base
finally:
proc.terminate()
proc.wait(timeout=30)
def _list_model_ids(base: str) -> list[str]:
listed = get(
URL(f"{base}/v1/models"),
headers=AUTH,
params=ModelsListParams(return_wildcard_routes=False),
response_type=ModelsListResponse,
)
return [entry.id for entry in unwrap(listed).data]
def test_partial_prefix_wildcard_expands_to_unity_catalog_names(wildcard_proxy: str) -> None:
ids = _list_model_ids(wildcard_proxy)
system_ai_ids = [model_id for model_id in ids if model_id.startswith("databricks/system.ai.")]
assert system_ai_ids, f"no databricks/system.ai.* expansion in {ids}"
malformed = [model_id for model_id in system_ai_ids if "/" in model_id.removeprefix("databricks/system.ai.")]
assert not malformed, f"expanded ids embed the provider-prefixed cost-map key after 'system.ai.': {malformed}"
def test_provider_wildcard_still_expands_to_cost_map_names(wildcard_proxy: str) -> None:
ids = _list_model_ids(wildcard_proxy)
databricks_ids = [model_id for model_id in ids if model_id.startswith("databricks/databricks-")]
assert databricks_ids, f"databricks/* did not expand to cost-map names: {ids}"
def test_model_info_keeps_provider_model_for_expanded_deployments(wildcard_proxy: str) -> None:
info = unwrap(
get(
URL(f"{wildcard_proxy}/model/info"),
headers=AUTH,
params=NoBody(),
response_type=_ModelInfoResponse,
)
)
bad = [row.model_name for row in info.data if "system.ai.databricks/" in 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
@ -11,9 +12,7 @@ def test_get_team_models_for_all_models_and_team_only_models():
model_access_groups = {}
include_model_access_groups = False
result = get_team_models(
team_models, proxy_model_list, model_access_groups, include_model_access_groups
)
result = get_team_models(team_models, proxy_model_list, model_access_groups, include_model_access_groups)
combined_models = team_models + proxy_model_list
assert set(result) == set(combined_models)
@ -246,9 +245,7 @@ def test_get_key_models_does_not_mutate_input():
),
],
)
def test_get_complete_model_list_order(
key_models, team_models, proxy_model_list, model_list, expected
):
def test_get_complete_model_list_order(key_models, team_models, proxy_model_list, model_list, expected):
"""
Test that get_complete_model_list preserves order
"""
@ -401,9 +398,7 @@ def test_wildcard_credential_hydration_preserves_deployment_params(
captured_params["api_key"] = litellm_params.api_key
captured_params["api_version"] = litellm_params.api_version
captured_params["credential_name"] = litellm_params.litellm_credential_name
captured_params["has_unexpected_field"] = hasattr(
litellm_params, "unexpected_field"
)
captured_params["has_unexpected_field"] = hasattr(litellm_params, "unexpected_field")
return ["gpt-4o"]
monkeypatch.setattr(model_checks, "get_provider_models", fake_get_provider_models)
@ -448,9 +443,7 @@ def test_wildcard_custom_prefix_does_not_stack_provider_prefix(monkeypatch):
result = get_known_models_from_wildcard(
wildcard_model="ollama_server1/*",
litellm_params=LiteLLM_Params(
model="ollama_chat/*", custom_llm_provider="ollama_chat"
),
litellm_params=LiteLLM_Params(model="ollama_chat/*", custom_llm_provider="ollama_chat"),
)
assert result == ["ollama_server1/gemma3:1b", "ollama_server1/llama3:8b"]
@ -477,14 +470,47 @@ def test_wildcard_custom_prefix_keeps_org_segment_for_non_provider_first_segment
result = get_known_models_from_wildcard(
wildcard_model="my_hf/*",
litellm_params=LiteLLM_Params(
model="huggingface/*", custom_llm_provider="huggingface"
),
litellm_params=LiteLLM_Params(model="huggingface/*", custom_llm_provider="huggingface"),
)
assert result == ["my_hf/meta-llama/Llama-3-8B"]
def test_partial_prefix_wildcard_expands_against_stripped_provider_ids(monkeypatch):
"""A partial wildcard like ``databricks/system.ai.*`` must splice the prefix before
``*`` onto the provider-stripped model id, not onto the whole cost-map key. Prepending
to the provider-prefixed key yielded ``databricks/system.ai.databricks/databricks-x``,
an id no upstream accepts.
"""
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,
):
@ -927,9 +953,7 @@ def test_add_known_models_refreshes_models_by_provider_for_wildcard_expansion():
assert fake_model not in litellm.models_by_provider["vertex_ai"]
try:
litellm.add_known_models(
model_cost_map={
fake_model: {"litellm_provider": "vertex_ai-language-models", "mode": "chat"}
}
model_cost_map={fake_model: {"litellm_provider": "vertex_ai-language-models", "mode": "chat"}}
)
assert fake_model in litellm.models_by_provider["vertex_ai"]
assert litellm.models_by_provider is captured_reference
@ -977,6 +1001,4 @@ def test_transcribe_is_a_known_provider_for_wildcard_expansion():
assert "transcribe" in litellm.models_by_provider
assert "transcribe/StartTranscriptionJob" in litellm.models_by_provider["transcribe"]
assert get_provider_models("transcribe") == ["transcribe/StartTranscriptionJob"]
assert get_known_models_from_wildcard("transcribe/*") == [
"transcribe/StartTranscriptionJob"
]
assert get_known_models_from_wildcard("transcribe/*") == ["transcribe/StartTranscriptionJob"]

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,20 @@ 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(
"databricks/system.ai.*",
["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_bare_provider_and_star_wildcard_unchanged():
candidates: Final = ("databricks/databricks-gpt-5",)
assert _wildcard_health_check_models("databricks/*", candidates) == candidates
assert _wildcard_health_check_models("*", candidates) == candidates