fix(proxy): hydrate wildcard discovery credentials (#28284) (#28419)

* fix(proxy): hydrate wildcard discovery credentials

* fix(proxy): constrain wildcard credential hydration

Co-authored-by: Dibyo Mukherjee <dibyo@adobe.com>
This commit is contained in:
yuneng-jiang 2026-05-20 20:03:05 -07:00 committed by GitHub
parent 79a5a7abad
commit 37ef8d9059
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 276 additions and 3 deletions

View File

@ -4,13 +4,17 @@ from typing import Dict, List, Optional, Set
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth
from litellm.router import Router
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
from litellm.types.router import LiteLLM_Params
from litellm.types.router import CredentialLiteLLMParams, LiteLLM_Params
from litellm.utils import get_valid_models
_CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields)
def _check_wildcard_routing(model: str) -> bool:
"""
Returns True if a model is a provider wildcard.
@ -178,6 +182,7 @@ def get_complete_model_list(
model_access_groups: Dict[str, List[str]] = {},
include_model_access_groups: Optional[bool] = False,
only_model_access_groups: Optional[bool] = False,
team_id: Optional[str] = None,
) -> List[str]:
"""Logic for returning complete model list for a given key + team pair"""
@ -222,6 +227,7 @@ def get_complete_model_list(
unique_models=unique_models,
return_wildcard_routes=return_wildcard_routes,
llm_router=llm_router,
team_id=team_id,
)
complete_model_list = unique_models + all_wildcard_models
@ -229,6 +235,29 @@ def get_complete_model_list(
return complete_model_list
def _hydrate_litellm_credential_name(
litellm_params: Optional[LiteLLM_Params],
) -> Optional[LiteLLM_Params]:
if litellm_params is None or litellm_params.litellm_credential_name is None:
return litellm_params
credential_values = CredentialAccessor.get_credential_values(
litellm_params.litellm_credential_name
)
if not credential_values:
return litellm_params
litellm_params = litellm_params.model_copy()
for key, value in credential_values.items():
if (
key in _CREDENTIAL_LITELLM_PARAM_FIELDS
and getattr(litellm_params, key, None) is None
):
setattr(litellm_params, key, value)
litellm_params.litellm_credential_name = None
return litellm_params
def get_known_models_from_wildcard(
wildcard_model: str, litellm_params: Optional[LiteLLM_Params] = None
) -> List[str]:
@ -247,7 +276,7 @@ def get_known_models_from_wildcard(
else:
provider = wildcard_provider_prefix
# get all known provider models
litellm_params = _hydrate_litellm_credential_name(litellm_params)
wildcard_models = get_provider_models(
provider=provider, litellm_params=litellm_params
@ -285,6 +314,7 @@ def _get_wildcard_models(
unique_models: List[str],
return_wildcard_routes: Optional[bool] = False,
llm_router: Optional[Router] = None,
team_id: Optional[str] = None,
) -> List[str]:
models_to_remove = set()
all_wildcard_models = []
@ -297,7 +327,9 @@ def _get_wildcard_models(
## get litellm params from model
if llm_router is not None:
model_list = llm_router.get_model_list(model_name=model)
model_list = llm_router.get_model_list(
model_name=model, team_id=team_id
)
if model_list:
for router_model in model_list:
wildcard_models = get_known_models_from_wildcard(

View File

@ -6068,6 +6068,8 @@ async def get_available_models_for_user(
include_model_access_groups=include_model_access_groups,
)
effective_team_id = team_id or user_api_key_dict.team_id
# Get complete model list
all_models = get_complete_model_list(
key_models=key_models,
@ -6080,6 +6082,7 @@ async def get_available_models_for_user(
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
only_model_access_groups=only_model_access_groups,
team_id=effective_team_id,
)
return all_models

View File

@ -249,3 +249,241 @@ def test_get_complete_model_list_byok_wildcard_expansion():
assert len(result) > 0
assert all(m.startswith("openai/") for m in result)
assert "openai/*" not in result
def test_get_complete_model_list_expands_team_scoped_wildcard_with_stored_credential(
monkeypatch,
):
"""
Team-scoped BYOK wildcard deployments are stored under an internal model_name,
with the public wildcard name in model_info.team_public_model_name.
"""
import litellm
from litellm import Router
from litellm.proxy.auth import model_checks
from litellm.proxy.auth.model_checks import get_complete_model_list
from litellm.types.utils import CredentialItem
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="openai-credential",
credential_info={"provider": "openai"},
credential_values={
"api_key": "stored-openai-key",
"api_base": "https://example.openai.test/v1",
},
)
],
)
captured_params = {}
def fake_get_provider_models(provider, litellm_params=None):
captured_params["provider"] = provider
captured_params["api_key"] = litellm_params.api_key
captured_params["api_base"] = litellm_params.api_base
captured_params["credential_name"] = litellm_params.litellm_credential_name
return ["gpt-4o"]
monkeypatch.setattr(model_checks, "get_provider_models", fake_get_provider_models)
router = Router(
model_list=[
{
"model_name": "model_name_team-1_generated",
"litellm_params": {
"model": "openai/*",
"custom_llm_provider": "openai",
"litellm_credential_name": "openai-credential",
},
"model_info": {
"team_id": "team-1",
"team_public_model_name": "openai/*",
},
}
]
)
result = get_complete_model_list(
key_models=[],
team_models=["openai/*"],
proxy_model_list=[],
user_model=None,
infer_model_from_keys=False,
llm_router=router,
team_id="team-1",
)
assert "openai/gpt-4o" in result
assert captured_params == {
"provider": "openai",
"api_key": "stored-openai-key",
"api_base": "https://example.openai.test/v1",
"credential_name": None,
}
def test_wildcard_credential_hydration_preserves_deployment_params(
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
from litellm.types.utils import CredentialItem
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="openai-credential",
credential_info={"provider": "openai"},
credential_values={
"api_key": "stored-openai-key",
"api_version": "credential-version",
"model": "openai/wrong-model",
"unexpected_field": "unexpected-value",
},
)
],
)
captured_params = {}
def fake_get_provider_models(provider, litellm_params=None):
captured_params["provider"] = provider
captured_params["model"] = litellm_params.model
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"
)
return ["gpt-4o"]
monkeypatch.setattr(model_checks, "get_provider_models", fake_get_provider_models)
result = get_known_models_from_wildcard(
wildcard_model="openai/*",
litellm_params=LiteLLM_Params(
model="openai/*",
custom_llm_provider="openai",
api_version="deployment-version",
litellm_credential_name="openai-credential",
),
)
assert result == ["openai/gpt-4o"]
assert captured_params == {
"provider": "openai",
"model": "openai/*",
"api_key": "stored-openai-key",
"api_version": "deployment-version",
"credential_name": None,
"has_unexpected_field": False,
}
def test_wildcard_credential_hydration_preserves_missing_credential_name(
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
monkeypatch.setattr(litellm, "credential_list", [])
captured_params = {}
def fake_get_provider_models(provider, litellm_params=None):
captured_params["provider"] = provider
captured_params["api_key"] = litellm_params.api_key
captured_params["credential_name"] = litellm_params.litellm_credential_name
return ["gpt-4o"]
monkeypatch.setattr(model_checks, "get_provider_models", fake_get_provider_models)
result = get_known_models_from_wildcard(
wildcard_model="openai/*",
litellm_params=LiteLLM_Params(
model="openai/*",
custom_llm_provider="openai",
api_key=None,
litellm_credential_name="missing-credential",
),
)
assert result == ["openai/gpt-4o"]
assert captured_params == {
"provider": "openai",
"api_key": None,
"credential_name": "missing-credential",
}
@pytest.mark.asyncio
async def test_get_available_models_for_user_expands_query_team_wildcard(
monkeypatch,
):
import litellm
from litellm import Router
from litellm.proxy.auth import model_checks
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import get_available_models_for_user
from litellm.types.utils import CredentialItem
monkeypatch.setattr(
litellm,
"credential_list",
[
CredentialItem(
credential_name="openai-credential",
credential_info={"provider": "openai"},
credential_values={"api_key": "stored-openai-key"},
)
],
)
def fake_get_provider_models(provider, litellm_params=None):
assert litellm_params.api_key == "stored-openai-key"
assert litellm_params.litellm_credential_name is None
return ["gpt-4o-mini"]
monkeypatch.setattr(model_checks, "get_provider_models", fake_get_provider_models)
router = Router(
model_list=[
{
"model_name": "model_name_team-1_generated",
"litellm_params": {
"model": "openai/*",
"custom_llm_provider": "openai",
"litellm_credential_name": "openai-credential",
},
"model_info": {
"team_id": "team-1",
"team_public_model_name": "openai/*",
},
}
]
)
result = await get_available_models_for_user(
user_api_key_dict=UserAPIKeyAuth(
api_key="sk-test",
models=[],
team_id="team-1",
team_models=["openai/*"],
),
llm_router=router,
general_settings={},
user_model=None,
team_id="team-1",
)
assert "openai/gpt-4o-mini" in result