mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
eae8ed7f3c
commit
99636dff2b
5 changed files with 260 additions and 45 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
|
|
|||
149
tests/e2e/test_wildcard_partial_prefix_expansion.py
Normal file
149
tests/e2e/test_wildcard_partial_prefix_expansion.py
Normal 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}"
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue