From 9f85fd62fe171dca1abc4828a8797a8c3d07a537 Mon Sep 17 00:00:00 2001 From: Ashwin Upadhyay Date: Sun, 17 May 2026 00:00:27 +0530 Subject: [PATCH] feat(proxy): wire nested group resolution into request-path callers Makes the feature actually active end-to-end. get_available_models_for_user (utils.py) and model_info_v1 (proxy_server.py) now fetch group_memberships once per call and pass them to get_key_models / get_team_models. When prisma_client is None (SDK-only mode) the map stays empty and behavior is identical to today. Also: - Rename _resolve_nested_groups -> resolve_nested_groups since it's used across modules; leading underscore was misleading. - Update NewModelGroupResponse.models_updated field comment to reflect the new semantics (deployment tags + membership edges, not just models). Refs #28032 --- litellm/proxy/auth/model_checks.py | 6 +++--- .../model_access_group_management_endpoints.py | 4 ++-- litellm/proxy/proxy_server.py | 12 ++++++++++++ litellm/proxy/utils.py | 13 +++++++++++++ .../model_management_endpoints.py | 2 +- 5 files changed, 31 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index 11903094ece..a513beecabf 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -42,7 +42,7 @@ def get_provider_models( return None -def _resolve_nested_groups( +def resolve_nested_groups( group_name: str, model_access_groups: Dict[str, List[str]], group_memberships: Dict[str, List[str]], @@ -66,7 +66,7 @@ def _resolve_nested_groups( resolved: List[str] = list(model_access_groups.get(group_name, [])) for child in group_memberships.get(group_name, []): resolved.extend( - _resolve_nested_groups( + resolve_nested_groups( group_name=child, model_access_groups=model_access_groups, group_memberships=group_memberships, @@ -92,7 +92,7 @@ def _get_models_from_access_groups( ): # remove access group, unless requested - e.g. when creating a key idx_to_remove.append(idx) new_models.extend( - _resolve_nested_groups( + resolve_nested_groups( group_name=model, model_access_groups=model_access_groups, group_memberships=memberships, diff --git a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py index e951e531a6d..60b7306d388 100644 --- a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py @@ -12,7 +12,7 @@ from fastapi import APIRouter, Depends, HTTPException from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.model_checks import _resolve_nested_groups +from litellm.proxy.auth.model_checks import resolve_nested_groups from litellm.proxy.auth.user_api_key_auth import user_api_key_auth # Clear cache and reload models to pick up the access group changes @@ -370,7 +370,7 @@ async def get_all_access_groups_from_db( result: Dict[str, AccessGroupInfo] = {} for group in all_groups: - expanded = _resolve_nested_groups( + expanded = resolve_nested_groups( group_name=group, model_access_groups=flat_group_models, group_memberships=group_memberships, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0114774cc2d..6cbc7a9216c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -367,6 +367,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import ( router as key_management_router, ) from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + get_group_memberships_from_db, router as model_access_group_management_router, ) from litellm.proxy.management_endpoints.model_management_endpoints import ( @@ -11976,15 +11977,26 @@ async def model_info_v1( # noqa: PLR0915 else: proxy_model_list = llm_router.get_model_names() model_access_groups = llm_router.get_model_access_groups() + + # Parent->child edges for nested access groups. Empty when no DB is + # configured, preserving today's flat behavior. + group_memberships: Dict[str, List[str]] = {} + if prisma_client is not None: + group_memberships = await get_group_memberships_from_db( + prisma_client=prisma_client + ) + key_models = get_key_models( user_api_key_dict=user_api_key_dict, proxy_model_list=proxy_model_list, model_access_groups=model_access_groups, + group_memberships=group_memberships, ) team_models = get_team_models( team_models=user_api_key_dict.team_models, proxy_model_list=proxy_model_list, model_access_groups=model_access_groups, + group_memberships=group_memberships, ) all_models_str = get_complete_model_list( key_models=key_models, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 559d5c99b9d..8c6bdb01d5c 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -5847,6 +5847,9 @@ async def get_available_models_for_user( get_key_models, get_team_models, ) + from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + get_group_memberships_from_db, + ) from litellm.proxy.management_endpoints.team_endpoints import validate_membership # Get proxy model list and access groups @@ -5857,12 +5860,21 @@ async def get_available_models_for_user( proxy_model_list = llm_router.get_model_names() model_access_groups = llm_router.get_model_access_groups() + # Parent->child edges for nested access groups. Empty when no DB is + # configured (e.g. SDK-only mode), preserving today's flat behavior. + group_memberships: Dict[str, List[str]] = {} + if prisma_client is not None: + group_memberships = await get_group_memberships_from_db( + prisma_client=prisma_client + ) + # Get key models key_models = get_key_models( user_api_key_dict=user_api_key_dict, proxy_model_list=proxy_model_list, model_access_groups=model_access_groups, include_model_access_groups=include_model_access_groups, + group_memberships=group_memberships, ) # Get team models @@ -5887,6 +5899,7 @@ async def get_available_models_for_user( proxy_model_list=proxy_model_list, model_access_groups=model_access_groups, include_model_access_groups=include_model_access_groups, + group_memberships=group_memberships, ) # Get complete model list diff --git a/litellm/types/proxy/management_endpoints/model_management_endpoints.py b/litellm/types/proxy/management_endpoints/model_management_endpoints.py index 0065cfabf89..ec7a209e981 100644 --- a/litellm/types/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/types/proxy/management_endpoints/model_management_endpoints.py @@ -33,7 +33,7 @@ class NewModelGroupResponse(BaseModel): access_group: str model_names: Optional[List[str]] = None model_ids: Optional[List[str]] = None - models_updated: int # Number of models updated + models_updated: int # Number of writes performed (deployment tags + membership edges for nested groups) class UpdateModelGroupRequest(BaseModel):