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 b05cfef5760..e951e531a6d 100644 --- a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py @@ -6,12 +6,13 @@ Endpoints here: """ import json -from typing import Any, Dict, List, Tuple +from typing import Any, Dict, List, Optional, Set, Tuple 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.user_api_key_auth import user_api_key_auth # Clear cache and reload models to pick up the access group changes @@ -31,22 +32,129 @@ from litellm.types.proxy.management_endpoints.model_management_endpoints import router = APIRouter() -def validate_models_exist(model_names: List[str], llm_router) -> Tuple[bool, List[str]]: +def validate_models_exist( + model_names: List[str], + llm_router, + known_access_groups: Optional[Set[str]] = None, +) -> Tuple[bool, List[str]]: """ - Validate that all requested model names exist in the router. - Checks only exact model name matches. + Validate that all requested member names exist as either a router model name + or a known access group name (for nested groups). Returns: - Tuple[bool, List[str]]: (all_valid, missing_models) + Tuple[bool, List[str]]: (all_valid, missing_names) """ if llm_router is None: return False, model_names router_model_names = set(llm_router.get_model_names()) - missing = [m for m in model_names if m not in router_model_names] + known_groups = known_access_groups or set() + missing = [ + m + for m in model_names + if m not in router_model_names and m not in known_groups + ] return (len(missing) == 0, missing) +def _classify_member_names( + names: List[str], + router_model_names: Set[str], + known_access_groups: Set[str], +) -> Tuple[List[str], List[str], List[str]]: + """ + Split a list of member names into (real_models, child_groups, unknown). + + Names registered both as a router model and as a group are classified + as a model - matches the existing precedence in `_get_models_from_access_groups` + where direct model_access_groups expansion runs before any other lookup. + """ + real_models: List[str] = [] + child_groups: List[str] = [] + unknown: List[str] = [] + for name in names: + if name in router_model_names: + real_models.append(name) + elif name in known_access_groups: + child_groups.append(name) + else: + unknown.append(name) + return real_models, child_groups, unknown + + +async def get_group_memberships_from_db( + prisma_client: PrismaClient, +) -> Dict[str, List[str]]: + """ + Build parent_group -> [child_groups] map from the membership table. + Single query, in-memory bucketing - no N+1. + """ + rows = await prisma_client.db.litellm_accessgroupmembership.find_many() + memberships: Dict[str, List[str]] = {} + for row in rows: + memberships.setdefault(row.parent_group, []).append(row.child_group) + return memberships + + +async def upsert_group_memberships( + parent_group: str, + child_groups: List[str], + prisma_client: PrismaClient, +) -> int: + """ + Insert parent_group -> child_group edges into LiteLLM_AccessGroupMembership. + Skips duplicates via the unique constraint. Rejects self-references eagerly. + + Returns: + int: number of new edges inserted. + """ + if not child_groups: + return 0 + + if parent_group in child_groups: + raise HTTPException( + status_code=400, + detail={ + "error": ( + f"Access group '{parent_group}' cannot include itself " + "as a member." + ) + }, + ) + + rows = [ + {"parent_group": parent_group, "child_group": child} + for child in child_groups + ] + result = await prisma_client.db.litellm_accessgroupmembership.create_many( + data=rows, + skip_duplicates=True, + ) + return result + + +async def delete_group_membership_edges( + access_group: str, + prisma_client: PrismaClient, +) -> int: + """ + Delete every membership row where `access_group` appears as parent or child. + Used during group deletion to avoid dangling references. + + Returns: + int: number of edges deleted. + """ + result = await prisma_client.db.litellm_accessgroupmembership.delete_many( + where={ + "OR": [ + {"parent_group": access_group}, + {"child_group": access_group}, + ] + } + ) + return result + + def add_access_group_to_deployment( model_info: Dict[str, Any], access_group: str ) -> Tuple[Dict[str, Any], bool]: @@ -209,16 +317,24 @@ async def get_all_access_groups_from_db( prisma_client: PrismaClient, ) -> Dict[str, AccessGroupInfo]: """ - Get all access groups from the database. + Get all access groups from the database, including nested-group structure. + + Builds the direct group -> {models, count} map by scanning deployments, + then layers in parent/child edges from LiteLLM_AccessGroupMembership. + Pure-composition groups (those that exist only as a parent in the + membership table, with no deployment tag) are surfaced too. + + `model_names` is expanded transitively via DFS - cyclic edges are skipped. + `deployment_count` remains the direct-tag count, not transitive. Returns: - Dict[str, AccessGroupInfo]: Dictionary mapping access_group name to info + Dict[str, AccessGroupInfo]: name -> info, including parent/child groups. """ - # Get all deployments + # Direct group membership from deployment tags deployments = await prisma_client.db.litellm_proxymodeltable.find_many() - # Build access group map - access_group_map: Dict[str, Dict[str, Any]] = {} + direct_models: Dict[str, Set[str]] = {} + deployment_count: Dict[str, int] = {} for deployment in deployments: model_info = deployment.model_info or {} @@ -226,22 +342,49 @@ async def get_all_access_groups_from_db( model_name = deployment.model_name for access_group in access_groups: - if access_group not in access_group_map: - access_group_map[access_group] = { - "model_names": set(), - "deployment_count": 0, - } + direct_models.setdefault(access_group, set()).add(model_name) + deployment_count[access_group] = ( + deployment_count.get(access_group, 0) + 1 + ) - access_group_map[access_group]["model_names"].add(model_name) - access_group_map[access_group]["deployment_count"] += 1 + # Group-to-group edges + group_memberships = await get_group_memberships_from_db(prisma_client) - # Convert to AccessGroupInfo objects - result = {} - for access_group, data in access_group_map.items(): - result[access_group] = AccessGroupInfo( - access_group=access_group, - model_names=sorted(list(data["model_names"])), - deployment_count=data["deployment_count"], + # Pure-composition groups exist only in the membership table. + # Surface them so they appear in /access_group/list and existence checks. + all_groups: Set[str] = set(direct_models.keys()) + for parent, children in group_memberships.items(): + all_groups.add(parent) + all_groups.update(children) + + # Reverse index: child -> [parents] + parents_of: Dict[str, List[str]] = {} + for parent, children in group_memberships.items(): + for child in children: + parents_of.setdefault(child, []).append(parent) + + # Build the flat group -> [direct model names] map for resolution + flat_group_models: Dict[str, List[str]] = { + g: sorted(list(direct_models.get(g, set()))) for g in all_groups + } + + result: Dict[str, AccessGroupInfo] = {} + for group in all_groups: + expanded = _resolve_nested_groups( + group_name=group, + model_access_groups=flat_group_models, + group_memberships=group_memberships, + visited=set(), + ) + # Dedup while preserving order, then sort for stable response shape + expanded_unique = sorted(set(expanded)) + + result[group] = AccessGroupInfo( + access_group=group, + model_names=expanded_unique, + deployment_count=deployment_count.get(group, 0), + parent_groups=sorted(parents_of.get(group, [])), + child_groups=sorted(group_memberships.get(group, [])), ) return result @@ -316,20 +459,6 @@ async def create_model_group( # If model_ids is provided, use it (more precise targeting) use_model_ids = has_model_ids - # Validate model_names exist in router (only if using model_names path) - if not use_model_ids and has_model_names: - assert data.model_names is not None - all_valid, missing_models = validate_models_exist( - model_names=data.model_names, - llm_router=llm_router, - ) - - if not all_valid: - raise HTTPException( - status_code=400, - detail={"error": f"Model(s) not found: {', '.join(missing_models)}"}, - ) - # Check if database is connected if prisma_client is None: raise HTTPException( @@ -338,7 +467,8 @@ async def create_model_group( ) try: - # Check if access group already exists + # Check if access group already exists. Done before validation so we + # know the set of known groups (needed to accept nested-group members). existing_access_groups = await get_all_access_groups_from_db( prisma_client=prisma_client ) @@ -351,7 +481,30 @@ async def create_model_group( }, ) - # Update deployments using the appropriate method + # Validate model_names exist in router or as a known access group + # (only if using model_names path). model_ids targets specific deployments + # so nesting doesn't apply. + if not use_model_ids and has_model_names: + assert data.model_names is not None + known_groups = set(existing_access_groups.keys()) + all_valid, missing = validate_models_exist( + model_names=data.model_names, + llm_router=llm_router, + known_access_groups=known_groups, + ) + + if not all_valid: + raise HTTPException( + status_code=400, + detail={ + "error": f"Model(s) or access group(s) not found: {', '.join(missing)}" + }, + ) + + # Write path. model_ids -> existing per-deployment tagging. + # model_names -> classify into real models vs. nested child groups and + # route each to the appropriate table. + models_updated = 0 if use_model_ids: assert data.model_ids is not None models_updated = await update_specific_deployments_with_access_group( @@ -361,16 +514,32 @@ async def create_model_group( ) else: assert data.model_names is not None - models_updated = await update_deployments_with_access_group( - model_names=data.model_names, - access_group=data.access_group, - prisma_client=prisma_client, + router_model_names = ( + set(llm_router.get_model_names()) if llm_router is not None else set() ) + real_models, child_groups, _ = _classify_member_names( + names=data.model_names, + router_model_names=router_model_names, + known_access_groups=set(existing_access_groups.keys()), + ) + + if real_models: + models_updated += await update_deployments_with_access_group( + model_names=real_models, + access_group=data.access_group, + prisma_client=prisma_client, + ) + if child_groups: + models_updated += await upsert_group_memberships( + parent_group=data.access_group, + child_groups=child_groups, + prisma_client=prisma_client, + ) await clear_cache() verbose_proxy_logger.info( - f"Successfully created access group '{data.access_group}' with {models_updated} models updated" + f"Successfully created access group '{data.access_group}' with {models_updated} writes" ) return NewModelGroupResponse( @@ -588,22 +757,27 @@ async def update_access_group( detail={"error": f"Failed to check access group existence: {str(e)}"}, ) - # Validation: Check if all new models exist (only if using model_names path) + # Validation: Check if all new members exist as router models or known groups. + # model_ids path targets specific deployments, so nesting doesn't apply. if not use_model_ids and has_model_names: assert data.model_names is not None - all_valid, missing_models = validate_models_exist( + known_groups = set(access_groups_map.keys()) + all_valid, missing = validate_models_exist( model_names=data.model_names, llm_router=llm_router, + known_access_groups=known_groups, ) if not all_valid: raise HTTPException( status_code=400, - detail={"error": f"Model(s) not found: {', '.join(missing_models)}"}, + detail={ + "error": f"Model(s) or access group(s) not found: {', '.join(missing)}" + }, ) try: - # Step 1: Remove access group from ALL DB deployments (skip config models) + # Step 1a: Clear deployment tags from ALL DB deployments all_deployments = await prisma_client.db.litellm_proxymodeltable.find_many() for deployment in all_deployments: @@ -620,7 +794,15 @@ async def update_access_group( data={"model_info": json.dumps(updated_model_info)}, ) - # Step 2: Add access group using the appropriate method + # Step 1b: Clear existing parent->child edges where this group is parent. + # Edges where it appears as a child are preserved (other groups still + # reference this one and that's outside the scope of this update). + await prisma_client.db.litellm_accessgroupmembership.delete_many( + where={"parent_group": access_group} + ) + + # Step 2: Re-add membership using the appropriate write path + models_updated = 0 if use_model_ids: assert data.model_ids is not None models_updated = await update_specific_deployments_with_access_group( @@ -630,17 +812,33 @@ async def update_access_group( ) else: assert data.model_names is not None - models_updated = await update_deployments_with_access_group( - model_names=data.model_names, - access_group=access_group, - prisma_client=prisma_client, + router_model_names = ( + set(llm_router.get_model_names()) if llm_router is not None else set() ) + real_models, child_groups, _ = _classify_member_names( + names=data.model_names, + router_model_names=router_model_names, + known_access_groups=set(access_groups_map.keys()), + ) + + if real_models: + models_updated += await update_deployments_with_access_group( + model_names=real_models, + access_group=access_group, + prisma_client=prisma_client, + ) + if child_groups: + models_updated += await upsert_group_memberships( + parent_group=access_group, + child_groups=child_groups, + prisma_client=prisma_client, + ) # Clear cache and reload models to pick up the access group changes await clear_cache() verbose_proxy_logger.info( - f"Successfully updated access group '{access_group}' with {models_updated} models updated" + f"Successfully updated access group '{access_group}' with {models_updated} writes" ) return NewModelGroupResponse( @@ -740,6 +938,13 @@ async def delete_access_group( ) models_updated += 1 + # Clean up parent/child edges where this group appears on either side + # to avoid dangling references in the membership table. + await delete_group_membership_edges( + access_group=access_group, + prisma_client=prisma_client, + ) + # Clear cache and reload models to pick up the access group changes await clear_cache() diff --git a/litellm/types/proxy/management_endpoints/model_management_endpoints.py b/litellm/types/proxy/management_endpoints/model_management_endpoints.py index bbbfc0de9f8..0065cfabf89 100644 --- a/litellm/types/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/types/proxy/management_endpoints/model_management_endpoints.py @@ -53,8 +53,16 @@ class DeleteModelGroupResponse(BaseModel): class AccessGroupInfo(BaseModel): access_group: str - model_names: List[str] # List of model names in this access group - deployment_count: int # Total number of deployments with this access group + model_names: List[str] # Transitively-resolved model names (includes nested groups' models) + deployment_count: int # Direct-tag deployment count for this access group only + parent_groups: List[str] = Field( + default_factory=list, + description="Access groups that include THIS group as a member (nested groups).", + ) + child_groups: List[str] = Field( + default_factory=list, + description="Access groups that THIS group includes as members (nested groups).", + ) class ListAccessGroupsResponse(BaseModel):