feat(proxy): wire write path + endpoint integration for nested groups

validate_models_exist accepts known access group names so a group's
members can themselves be groups. _classify_member_names splits the
caller's list into real models vs child groups, and each is routed to
its own write path.

create_model_group / update_access_group: classify -> dual-write
(deployments tagged for real models, parent->child edges for child groups).
update_access_group also clears prior edges where this group is parent
before re-inserting.

delete_access_group: also deletes membership edges where this group
appears on either side, to avoid dangling references.

AccessGroupInfo gains parent_groups + child_groups (always present,
empty by default). get_all_access_groups_from_db expands model_names
transitively via _resolve_nested_groups, surfaces pure-composition
groups that exist only in the membership table, and reports parent/
child relationships.

Self-references (parent == child) are rejected at write time with 400.
Multi-hop cycles are caught by the read-path DFS guard from the prior
commit.

Refs #28032
This commit is contained in:
Ashwin Upadhyay 2026-05-16 23:50:54 +05:30
parent d9d0d94b55
commit f66864c05a
2 changed files with 271 additions and 58 deletions

View file

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

View file

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