mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
fix(proxy): hydrate wildcard discovery credentials
This commit is contained in:
parent
cff3e0b75e
commit
12eeece17a
3 changed files with 168 additions and 2 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue