mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
d9d0d94b55
commit
f66864c05a
2 changed files with 271 additions and 58 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue