fix(proxy): hydrate wildcard discovery credentials

This commit is contained in:
Dibyo Mukherjee 2026-05-19 14:52:06 -04:00
parent cff3e0b75e
commit 12eeece17a
3 changed files with 168 additions and 2 deletions

View file

@ -4,6 +4,7 @@ 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
@ -178,6 +179,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 +224,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 +232,25 @@ 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():
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 +269,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 +307,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 +320,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,141 @@ 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,
}
@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