From 99636dff2b661a965689c998587c481d137c18c8 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 28 Sep 2026 23:50:24 +0000 Subject: [PATCH] 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> --- .../health_check_helpers.py | 23 ++- litellm/proxy/auth/model_checks.py | 51 +++--- .../test_wildcard_partial_prefix_expansion.py | 149 ++++++++++++++++++ .../proxy/auth/test_model_checks.py | 64 +++++--- .../test_health_check_helpers.py | 18 +++ 5 files changed, 260 insertions(+), 45 deletions(-) create mode 100644 tests/e2e/test_wildcard_partial_prefix_expansion.py diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index a0f027cd58f..16342728774 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -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 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/e2e/test_wildcard_partial_prefix_expansion.py b/tests/e2e/test_wildcard_partial_prefix_expansion.py new file mode 100644 index 00000000000..48b265c2c62 --- /dev/null +++ b/tests/e2e/test_wildcard_partial_prefix_expansion.py @@ -0,0 +1,149 @@ +"""Regression: wildcard expansion of a partial provider prefix like +`databricks/system.ai.*` must produce `databricks/system.ai.` ids, +not `databricks/system.ai.databricks/databricks-`. + +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}" diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index 13171a42cda..9204a69a89e 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 @@ -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"] 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..4925bcb0b5d 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,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