From d9d0d94b55a98b05d7b55b9df6d02699437cdac7 Mon Sep 17 00:00:00 2001 From: Ashwin Upadhyay Date: Sat, 16 May 2026 23:28:06 +0530 Subject: [PATCH] feat(proxy): resolve nested access groups via DFS with cycle detection _get_models_from_access_groups takes a new optional group_memberships map (parent -> [child_groups]) and expands recursively via _resolve_nested_groups. Cycles are logged and skipped, not raised - this runs on the auth path and a malformed row must not 500 the proxy. Backwards compatible: callers that don't pass group_memberships get identical behavior to today. get_key_models, get_team_models thread the new kwarg through. Refs #28032 --- litellm/proxy/auth/model_checks.py | 54 ++++++++++++++++++++++++++++-- 1 file changed, 52 insertions(+), 2 deletions(-) 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