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
This commit is contained in:
Ashwin Upadhyay 2026-05-16 23:28:06 +05:30
parent 92729d8d64
commit d9d0d94b55

View file

@ -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