mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
92729d8d64
commit
d9d0d94b55
1 changed files with 52 additions and 2 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue