diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index 4493bb103d7..22c1a3e0678 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 Dict, List, Optional +from typing import Dict, List, Optional, Set import litellm from litellm._logging import verbose_proxy_logger @@ -11,6 +11,11 @@ from litellm.utils import get_valid_models def _check_wildcard_routing(model: str) -> bool: """ Returns True if a model is a provider wildcard. + + eg: + - anthropic/* + - openai/* + - * """ if model == "*": return True @@ -119,6 +124,7 @@ def get_complete_model_list( proxy_model_list: List[str], user_model: Optional[str], infer_model_from_keys: Optional[bool], + return_wildcard_routes: Optional[bool] = False, ) -> List[str]: """Logic for returning complete model list for a given key + team pair""" @@ -128,7 +134,7 @@ def get_complete_model_list( If list contains wildcard -> return known provider models """ - unique_models = set() + unique_models: Set[str] = set() if key_models: unique_models.update(key_models) elif team_models: @@ -143,10 +149,26 @@ def get_complete_model_list( valid_models = get_valid_models() unique_models.update(valid_models) + all_wildcard_models = _get_wildcard_models( + unique_models=unique_models, return_wildcard_routes=return_wildcard_routes + ) + + return list(unique_models) + all_wildcard_models + + +def _get_wildcard_models( + unique_models: Set[str], return_wildcard_routes: Optional[bool] = False +) -> List[str]: models_to_remove = set() all_wildcard_models = [] for model in unique_models: if _check_wildcard_routing(model=model): + + if ( + return_wildcard_routes is True + ): # will add the wildcard route to the list eg: anthropic/*. + all_wildcard_models.append(model) + provider = model.split("/")[0] # get all known provider models wildcard_models = get_provider_models(provider=provider) @@ -157,4 +179,4 @@ def get_complete_model_list( for model in models_to_remove: unique_models.remove(model) - return list(unique_models) + all_wildcard_models + return all_wildcard_models diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3a96205155b..957025ab3d8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -3322,6 +3322,7 @@ class ProxyStartupEvent: ) # if project requires model list async def model_list( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + return_wildcard_routes: Optional[bool] = False, ): """ Use `/model/info` - to get detailed model information, example - pricing, mode, etc. @@ -3354,6 +3355,7 @@ async def model_list( proxy_model_list=proxy_model_list, user_model=user_model, infer_model_from_keys=general_settings.get("infer_model_from_keys", False), + return_wildcard_routes=return_wildcard_routes, ) return dict( data=[