mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
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
This commit is contained in:
parent
f66864c05a
commit
9f85fd62fe
5 changed files with 31 additions and 6 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue