diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index bf76f99db69..11903094ece 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -42,20 +42,63 @@ def get_provider_models( return None +def _resolve_nested_groups( + group_name: str, + model_access_groups: Dict[str, List[str]], + group_memberships: Dict[str, List[str]], + visited: Set[str], +) -> List[str]: + """ + Expand a group name to the full list of model names it transitively includes, + following parent -> child edges in `group_memberships`. + + Runs on the auth path, so cyclic edges are logged and skipped rather than raised: + a malformed row must not 500 the proxy. + """ + if group_name in visited: + verbose_proxy_logger.warning( + "access group cycle detected at '%s' - skipping cyclic edge", + group_name, + ) + return [] + visited.add(group_name) + + resolved: List[str] = list(model_access_groups.get(group_name, [])) + for child in group_memberships.get(group_name, []): + resolved.extend( + _resolve_nested_groups( + group_name=child, + model_access_groups=model_access_groups, + group_memberships=group_memberships, + visited=visited, + ) + ) + return resolved + + def _get_models_from_access_groups( model_access_groups: Dict[str, List[str]], all_models: List[str], include_model_access_groups: Optional[bool] = False, + group_memberships: Optional[Dict[str, List[str]]] = None, ) -> List[str]: + memberships = group_memberships or {} idx_to_remove = [] new_models = [] for idx, model in enumerate(all_models): - if model in model_access_groups: + if model in model_access_groups or model in memberships: if ( not include_model_access_groups ): # remove access group, unless requested - e.g. when creating a key idx_to_remove.append(idx) - new_models.extend(model_access_groups[model]) + new_models.extend( + _resolve_nested_groups( + group_name=model, + model_access_groups=model_access_groups, + group_memberships=memberships, + visited=set(), + ) + ) for idx in sorted(idx_to_remove, reverse=True): all_models.pop(idx) @@ -96,6 +139,7 @@ def get_key_models( model_access_groups: Dict[str, List[str]], include_model_access_groups: Optional[bool] = False, only_model_access_groups: Optional[bool] = False, + group_memberships: Optional[Dict[str, List[str]]] = None, ) -> List[str]: """ Returns: @@ -104,6 +148,8 @@ def get_key_models( - If model_access_groups is provided, only return models that are in the access groups - If include_model_access_groups is True, it includes the 'keys' of the model_access_groups in the response - {"beta-models": ["gpt-4", "claude-v1"]} -> returns 'beta-models' + - If group_memberships is provided, expands nested groups transitively + (parent -> child edges); cyclic edges are skipped """ all_models: List[str] = [] if len(user_api_key_dict.models) > 0: @@ -123,6 +169,7 @@ def get_key_models( model_access_groups=model_access_groups, all_models=all_models, include_model_access_groups=include_model_access_groups, + group_memberships=group_memberships, ) # deduplicate while preserving order @@ -137,12 +184,14 @@ def get_team_models( proxy_model_list: List[str], model_access_groups: Dict[str, List[str]], include_model_access_groups: Optional[bool] = False, + group_memberships: Optional[Dict[str, List[str]]] = None, ) -> 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 + - If group_memberships is provided, expands nested groups transitively """ all_models_set: Set[str] = set() if len(team_models) > 0: @@ -158,6 +207,7 @@ def get_team_models( model_access_groups=model_access_groups, all_models=list(all_models_set), include_model_access_groups=include_model_access_groups, + group_memberships=group_memberships, ) # deduplicate while preserving order