From 0fe8bfe87ac9dcc53efb726d28de2853a7038270 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Wed, 18 Dec 2024 09:33:33 -0800 Subject: [PATCH] =?UTF-8?q?fix(proxy=5Fserver.py):=20pass=20model=20access?= =?UTF-8?q?=20groups=20to=20get=5Fkey/get=5Fteam=20mo=E2=80=A6=20(#7281)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(proxy_server.py): pass model access groups to get_key/get_team models allows end user to see actual models they have access to, instead of default models * fix(auth_checks.py): fix linting errors * fix: fix linting errors --- litellm/proxy/_experimental/out/404.html | 1 - .../proxy/_experimental/out/model_hub.html | 1 - .../proxy/_experimental/out/onboarding.html | 1 - litellm/proxy/_new_secret_config.yaml | 2 + litellm/proxy/auth/auth_checks.py | 4 +- litellm/proxy/auth/model_checks.py | 42 +++++++++++++--- litellm/proxy/proxy_server.py | 34 ++++++++++--- litellm/router.py | 4 +- tests/proxy_unit_tests/test_proxy_utils.py | 49 +++++++++++++++++++ 9 files changed, 118 insertions(+), 20 deletions(-) delete mode 100644 litellm/proxy/_experimental/out/404.html delete mode 100644 litellm/proxy/_experimental/out/model_hub.html delete mode 100644 litellm/proxy/_experimental/out/onboarding.html 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 Dashboard

404

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"]