* 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:
parent
79a5a7abad
commit
37ef8d9059
@ -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(
|
||||
|
||||
@ -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,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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user