diff --git a/litellm/proxy/_experimental/out/404.html b/litellm/proxy/_experimental/out/404.html
deleted file mode 100644
index eb587481b37..00000000000
--- a/litellm/proxy/_experimental/out/404.html
+++ /dev/null
@@ -1 +0,0 @@
-
404: This page could not be found.LiteLLM Dashboard404
This page could not be found.
\ No newline at end of file
diff --git a/litellm/proxy/_experimental/out/model_hub.html b/litellm/proxy/_experimental/out/model_hub.html
deleted file mode 100644
index 5b1276d2bce..00000000000
--- a/litellm/proxy/_experimental/out/model_hub.html
+++ /dev/null
@@ -1 +0,0 @@
-LiteLLM Dashboard
\ No newline at end of file
diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding.html
deleted file mode 100644
index b9619b97b41..00000000000
--- a/litellm/proxy/_experimental/out/onboarding.html
+++ /dev/null
@@ -1 +0,0 @@
-LiteLLM Dashboard
\ No newline at end of file
diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml
index 86946a65dac..e2b562bfc15 100644
--- a/litellm/proxy/_new_secret_config.yaml
+++ b/litellm/proxy/_new_secret_config.yaml
@@ -5,6 +5,8 @@ model_list:
api_key: os.environ/AZURE_API_KEY
api_base: os.environ/AZURE_API_BASE
temperature: 0.2
+ model_info:
+ access_groups: ["default"]
- model_name: gpt-4o
litellm_params:
model: openai/gpt-4o
diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py
index 45e8c844c4f..2127dfb5090 100644
--- a/litellm/proxy/auth/auth_checks.py
+++ b/litellm/proxy/auth/auth_checks.py
@@ -12,7 +12,7 @@ Run checks for:
import time
import traceback
from datetime import datetime
-from typing import TYPE_CHECKING, Any, List, Literal, Optional
+from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional
import httpx
from pydantic import BaseModel
@@ -869,7 +869,7 @@ async def can_key_call_model(
)
from collections import defaultdict
- access_groups = defaultdict(list)
+ access_groups: Dict[str, List[str]] = defaultdict(list)
if llm_router:
access_groups = llm_router.get_model_access_groups(model_name=model)
if (
diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py
index dd22a7f47d0..c568d887a96 100644
--- a/litellm/proxy/auth/model_checks.py
+++ b/litellm/proxy/auth/model_checks.py
@@ -1,6 +1,6 @@
# What is this?
## Common checks for /v1/models and `/model/info`
-from typing import List, Optional
+from typing import Dict, List, Optional
import litellm
from litellm._logging import verbose_proxy_logger
@@ -38,15 +38,36 @@ def get_provider_models(provider: str) -> Optional[List[str]]:
return None
+def _get_models_from_access_groups(
+ model_access_groups: Dict[str, List[str]],
+ all_models: List[str],
+) -> List[str]:
+ idx_to_remove = []
+ new_models = []
+ for idx, model in enumerate(all_models):
+ if model in model_access_groups:
+ idx_to_remove.append(idx)
+ new_models.extend(model_access_groups[model])
+
+ for idx in sorted(idx_to_remove, reverse=True):
+ all_models.pop(idx)
+
+ all_models.extend(new_models)
+ return all_models
+
+
def get_key_models(
- user_api_key_dict: UserAPIKeyAuth, proxy_model_list: List[str]
+ user_api_key_dict: UserAPIKeyAuth,
+ proxy_model_list: List[str],
+ model_access_groups: Dict[str, List[str]],
) -> List[str]:
"""
Returns:
- List of model name strings
- Empty list if no models set
+ - If model_access_groups is provided, only return models that are in the access groups
"""
- all_models = []
+ all_models: List[str] = []
if len(user_api_key_dict.models) > 0:
all_models = user_api_key_dict.models
if SpecialModelNames.all_team_models.value in all_models:
@@ -54,17 +75,24 @@ def get_key_models(
if SpecialModelNames.all_proxy_models.value in all_models:
all_models = proxy_model_list
+ all_models = _get_models_from_access_groups(
+ model_access_groups=model_access_groups, all_models=all_models
+ )
+
verbose_proxy_logger.debug("ALL KEY MODELS - {}".format(len(all_models)))
return all_models
def get_team_models(
- user_api_key_dict: UserAPIKeyAuth, proxy_model_list: List[str]
+ user_api_key_dict: UserAPIKeyAuth,
+ proxy_model_list: List[str],
+ model_access_groups: Dict[str, List[str]],
) -> List[str]:
"""
Returns:
- List of model name strings
- Empty list if no models set
+ - If model_access_groups is provided, only return models that are in the access groups
"""
all_models = []
if len(user_api_key_dict.team_models) > 0:
@@ -74,6 +102,10 @@ def get_team_models(
if SpecialModelNames.all_proxy_models.value in all_models:
all_models = proxy_model_list
+ all_models = _get_models_from_access_groups(
+ model_access_groups=model_access_groups, all_models=all_models
+ )
+
verbose_proxy_logger.debug("ALL TEAM MODELS - {}".format(len(all_models)))
return all_models
@@ -93,9 +125,7 @@ def get_complete_model_list(
If list contains wildcard -> return known provider models
"""
-
unique_models = set()
-
if key_models:
unique_models.update(key_models)
elif team_models:
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index 4210d6035c5..2052082c4bc 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -101,7 +101,7 @@ def generate_feedback_box():
print() # noqa
-import pydantic
+from collections import defaultdict
import litellm
from litellm import (
@@ -3207,16 +3207,23 @@ async def model_list(
"""
global llm_model_list, general_settings, llm_router
all_models = []
+ model_access_groups: Dict[str, List[str]] = defaultdict(list)
## CHECK IF MODEL RESTRICTIONS ARE SET AT KEY/TEAM LEVEL ##
if llm_router is None:
proxy_model_list = []
else:
proxy_model_list = llm_router.get_model_names()
+ model_access_groups = llm_router.get_model_access_groups()
key_models = get_key_models(
- user_api_key_dict=user_api_key_dict, proxy_model_list=proxy_model_list
+ user_api_key_dict=user_api_key_dict,
+ proxy_model_list=proxy_model_list,
+ model_access_groups=model_access_groups,
)
+
team_models = get_team_models(
- user_api_key_dict=user_api_key_dict, proxy_model_list=proxy_model_list
+ user_api_key_dict=user_api_key_dict,
+ proxy_model_list=proxy_model_list,
+ model_access_groups=model_access_groups,
)
all_models = get_complete_model_list(
key_models=key_models,
@@ -7136,17 +7143,22 @@ async def model_info_v1( # noqa: PLR0915
return {"data": _deployment_info_dict}
all_models: List[dict] = []
+ model_access_groups: Dict[str, List[str]] = defaultdict(list)
## CHECK IF MODEL RESTRICTIONS ARE SET AT KEY/TEAM LEVEL ##
if llm_router is None:
proxy_model_list = []
else:
proxy_model_list = llm_router.get_model_names()
-
+ model_access_groups = llm_router.get_model_access_groups()
key_models = get_key_models(
- user_api_key_dict=user_api_key_dict, proxy_model_list=proxy_model_list
+ user_api_key_dict=user_api_key_dict,
+ proxy_model_list=proxy_model_list,
+ model_access_groups=model_access_groups,
)
team_models = get_team_models(
- user_api_key_dict=user_api_key_dict, proxy_model_list=proxy_model_list
+ user_api_key_dict=user_api_key_dict,
+ proxy_model_list=proxy_model_list,
+ model_access_groups=model_access_groups,
)
all_models_str = get_complete_model_list(
key_models=key_models,
@@ -7358,16 +7370,22 @@ async def model_group_info(
status_code=500, detail={"error": "LLM Router is not loaded in"}
)
## CHECK IF MODEL RESTRICTIONS ARE SET AT KEY/TEAM LEVEL ##
+ model_access_groups: Dict[str, List[str]] = defaultdict(list)
if llm_router is None:
proxy_model_list = []
else:
proxy_model_list = llm_router.get_model_names()
+ model_access_groups = llm_router.get_model_access_groups()
key_models = get_key_models(
- user_api_key_dict=user_api_key_dict, proxy_model_list=proxy_model_list
+ user_api_key_dict=user_api_key_dict,
+ proxy_model_list=proxy_model_list,
+ model_access_groups=model_access_groups,
)
team_models = get_team_models(
- user_api_key_dict=user_api_key_dict, proxy_model_list=proxy_model_list
+ user_api_key_dict=user_api_key_dict,
+ proxy_model_list=proxy_model_list,
+ model_access_groups=model_access_groups,
)
all_models_str = get_complete_model_list(
key_models=key_models,
diff --git a/litellm/router.py b/litellm/router.py
index ece6a97598b..368be65775d 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -4653,7 +4653,9 @@ class Router:
return returned_models
return None
- def get_model_access_groups(self, model_name: Optional[str] = None):
+ def get_model_access_groups(
+ self, model_name: Optional[str] = None
+ ) -> Dict[str, List[str]]:
"""
If model_name is provided, only return access groups for that model.
"""
diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py
index 5d2979de134..5ebff141a97 100644
--- a/tests/proxy_unit_tests/test_proxy_utils.py
+++ b/tests/proxy_unit_tests/test_proxy_utils.py
@@ -885,3 +885,52 @@ def test_enforced_params_check(
user_api_key_dict=user_api_key_dict,
premium_user=True,
)
+
+
+def test_get_key_models():
+ from litellm.proxy.auth.model_checks import get_key_models
+ from collections import defaultdict
+
+ user_api_key_dict = UserAPIKeyAuth(
+ api_key="test_api_key",
+ user_id="test_user_id",
+ org_id="test_org_id",
+ models=["default"],
+ )
+ proxy_model_list = ["gpt-4o", "gpt-3.5-turbo"]
+ model_access_groups = defaultdict(list)
+ model_access_groups["default"].extend(["gpt-4o", "gpt-3.5-turbo"])
+ model_access_groups["default"].extend(["gpt-4o-mini"])
+ model_access_groups["team2"].extend(["gpt-3.5-turbo"])
+
+ result = get_key_models(
+ user_api_key_dict=user_api_key_dict,
+ proxy_model_list=proxy_model_list,
+ model_access_groups=model_access_groups,
+ )
+ assert result == ["gpt-4o", "gpt-3.5-turbo", "gpt-4o-mini"]
+
+
+def test_get_team_models():
+ from litellm.proxy.auth.model_checks import get_team_models
+ from collections import defaultdict
+
+ user_api_key_dict = UserAPIKeyAuth(
+ api_key="test_api_key",
+ user_id="test_user_id",
+ org_id="test_org_id",
+ models=[],
+ team_models=["default"],
+ )
+ proxy_model_list = ["gpt-4o", "gpt-3.5-turbo"]
+ model_access_groups = defaultdict(list)
+ model_access_groups["default"].extend(["gpt-4o", "gpt-3.5-turbo"])
+ model_access_groups["default"].extend(["gpt-4o-mini"])
+ model_access_groups["team2"].extend(["gpt-3.5-turbo"])
+
+ result = get_team_models(
+ user_api_key_dict=user_api_key_dict,
+ proxy_model_list=proxy_model_list,
+ model_access_groups=model_access_groups,
+ )
+ assert result == ["gpt-4o", "gpt-3.5-turbo", "gpt-4o-mini"]