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):