Merge pull request #21190 from BerriAI/litellm_access_groups_inte

[Feature] Access Group Checks
This commit is contained in:
yuneng-jiang 2026-02-14 16:28:09 -08:00 committed by GitHub
commit 55f169aa69
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
24 changed files with 1164 additions and 134 deletions

View File

@ -1360,6 +1360,9 @@ SPECIAL_LITELLM_AUTH_TOKEN = ["ui-token"]
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(
os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60)
)
DEFAULT_ACCESS_GROUP_CACHE_TTL = int(
os.getenv("DEFAULT_ACCESS_GROUP_CACHE_TTL", 600)
)
# Sentry Scrubbing Configuration
SENTRY_DENYLIST = [

View File

@ -2563,6 +2563,21 @@ class LiteLLM_TagTable(LiteLLMPydanticObjectBase):
model_config = ConfigDict(protected_namespaces=())
class LiteLLM_AccessGroupTable(LiteLLMPydanticObjectBase):
access_group_id: str
access_group_name: str
description: Optional[str] = None
access_model_names: List[str] = []
access_mcp_server_ids: List[str] = []
access_agent_ids: List[str] = []
assigned_team_ids: List[str] = []
assigned_key_ids: List[str] = []
created_at: Optional[datetime] = None
created_by: Optional[str] = None
updated_at: Optional[datetime] = None
updated_by: Optional[str] = None
class LiteLLM_SpendLogs(LiteLLMPydanticObjectBase):
request_id: str
api_key: str

View File

@ -145,7 +145,10 @@ class AgentRequestHandler:
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[str]:
"""
Get allowed agents for a key from its object_permission.
Get allowed agents for a key.
1. First checks native key-level agent permissions (object_permission)
2. Also includes agents from key's access_group_ids (unified access groups)
Note: object_permission is already loaded by get_key_object() in main auth flow.
"""
@ -153,25 +156,37 @@ class AgentRequestHandler:
return []
try:
# Get key object permission (already loaded in main auth flow)
all_agents: List[str] = []
# 1. Get agents from object_permission (native permissions)
key_object_permission = AgentRequestHandler._get_key_object_permission(
user_api_key_auth
)
if key_object_permission is None:
return []
if key_object_permission is not None:
# Get direct agents
direct_agents = key_object_permission.agents or []
# Get direct agents
direct_agents = key_object_permission.agents or []
# Get agents from access groups
access_group_agents = (
await AgentRequestHandler._get_agents_from_access_groups(
key_object_permission.agent_access_groups or []
# Get agents from access groups
access_group_agents = (
await AgentRequestHandler._get_agents_from_access_groups(
key_object_permission.agent_access_groups or []
)
)
)
# Combine both lists
all_agents = direct_agents + access_group_agents
all_agents = direct_agents + access_group_agents
# 2. Fallback: get agent IDs from key's access_group_ids (unified access groups)
key_access_group_ids = user_api_key_auth.access_group_ids or []
if key_access_group_ids:
from litellm.proxy.auth.auth_checks import (
_get_agent_ids_from_access_groups,
)
unified_agents = await _get_agent_ids_from_access_groups(
access_group_ids=key_access_group_ids,
)
all_agents.extend(unified_agents)
return list(set(all_agents))
except Exception as e:
verbose_logger.warning(f"Failed to get allowed agents for key: {str(e)}")
@ -182,9 +197,12 @@ class AgentRequestHandler:
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[str]:
"""
Get allowed agents for a team from its object_permission.
Get allowed agents for a team.
Note: object_permission is already loaded by get_team_object() in main auth flow.
1. First checks native team-level agent permissions (object_permission)
2. Also includes agents from team's access_group_ids (unified access groups)
Fetches the team object once and reuses it for both permission sources.
"""
if user_api_key_auth is None:
return []
@ -193,26 +211,57 @@ class AgentRequestHandler:
return []
try:
# Get team object permission (already loaded in main auth flow)
object_permissions = await AgentRequestHandler._get_team_object_permission(
user_api_key_auth
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if object_permissions is None:
if not prisma_client:
return []
# Get direct agents
direct_agents = object_permissions.agents or []
# Get agents from access groups
access_group_agents = (
await AgentRequestHandler._get_agents_from_access_groups(
object_permissions.agent_access_groups or []
)
# Fetch the team object once for both permission sources
team_obj = await get_team_object(
team_id=user_api_key_auth.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
# Combine both lists
all_agents = direct_agents + access_group_agents
if team_obj is None:
return []
all_agents: List[str] = []
# 1. Get agents from object_permission (native permissions)
object_permissions = team_obj.object_permission
if object_permissions is not None:
# Get direct agents
direct_agents = object_permissions.agents or []
# Get agents from access groups
access_group_agents = (
await AgentRequestHandler._get_agents_from_access_groups(
object_permissions.agent_access_groups or []
)
)
all_agents = direct_agents + access_group_agents
# 2. Also include agents from team's access_group_ids (unified access groups)
team_access_group_ids = team_obj.access_group_ids or []
if team_access_group_ids:
from litellm.proxy.auth.auth_checks import (
_get_agent_ids_from_access_groups,
)
unified_agents = await _get_agent_ids_from_access_groups(
access_group_ids=team_access_group_ids,
)
all_agents.extend(unified_agents)
return list(set(all_agents))
except Exception as e:
verbose_logger.warning(f"Failed to get allowed agents for team: {str(e)}")

View File

@ -23,6 +23,7 @@ from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.constants import (
CLI_JWT_EXPIRATION_HOURS,
CLI_JWT_TOKEN_NAME,
DEFAULT_ACCESS_GROUP_CACHE_TTL,
DEFAULT_IN_MEMORY_TTL,
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
DEFAULT_MAX_RECURSE_DEPTH,
@ -32,6 +33,7 @@ from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.proxy._types import (
RBAC_ROLES,
CallInfo,
LiteLLM_AccessGroupTable,
LiteLLM_BudgetTable,
LiteLLM_EndUserTable,
Litellm_EntityType,
@ -210,7 +212,7 @@ async def common_checks(
# 2. If team can call model
if _model and team_object:
if not can_team_access_model(
if not await can_team_access_model(
model=_model,
team_object=team_object,
llm_router=llm_router,
@ -1499,6 +1501,110 @@ async def get_team_object(
)
async def _cache_access_object(
access_group_id: str,
access_group_table: LiteLLM_AccessGroupTable,
user_api_key_cache: DualCache,
proxy_logging_obj: Optional[ProxyLogging] = None,
):
key = "access_group_id:{}".format(access_group_id)
await user_api_key_cache.async_set_cache(
key=key,
value=access_group_table,
ttl=DEFAULT_ACCESS_GROUP_CACHE_TTL,
)
async def _delete_cache_access_object(
access_group_id: str,
user_api_key_cache: DualCache,
proxy_logging_obj: Optional[ProxyLogging] = None,
):
key = "access_group_id:{}".format(access_group_id)
user_api_key_cache.delete_cache(key=key)
## UPDATE REDIS CACHE ##
if proxy_logging_obj is not None:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(
key=key
)
@log_db_metrics
async def get_access_object(
access_group_id: str,
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> LiteLLM_AccessGroupTable:
"""
- Check if access_group_id in proxy AccessGroupTable
- Always checks cache first, then DB only when not found in cache
- if valid, return LiteLLM_AccessGroupTable object
- if not, then raise an error
Unlike get_team_object, this has no check_cache_only or check_db_only flags;
it always follows cache-first-then-db semantics.
Raises:
- HTTPException: If access group doesn't exist in db or cache (status_code=404)
"""
if prisma_client is None:
raise Exception(
"No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys"
)
key = "access_group_id:{}".format(access_group_id)
# Always check cache first
cached_access_obj = await user_api_key_cache.async_get_cache(key=key)
if cached_access_obj is not None:
if isinstance(cached_access_obj, dict):
return LiteLLM_AccessGroupTable(**cached_access_obj)
elif isinstance(cached_access_obj, LiteLLM_AccessGroupTable):
return cached_access_obj
# Not in cache - fetch from DB
try:
response = await prisma_client.db.litellm_accessgrouptable.find_unique(
where={"access_group_id": access_group_id}
)
if response is None:
raise HTTPException(
status_code=404,
detail={
"error": f"Access group doesn't exist in db. Access group={access_group_id}."
},
)
_response = LiteLLM_AccessGroupTable(**response.dict())
# Save to cache
await _cache_access_object(
access_group_id=access_group_id,
access_group_table=_response,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
return _response
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception(
"Error getting access group for access_group_id: %s",
access_group_id,
)
raise HTTPException(
status_code=404,
detail={
"error": f"Access group doesn't exist in db. Access group={access_group_id}. Error: {e}"
},
)
@log_db_metrics
async def get_team_object_by_alias(
team_alias: str,
@ -2013,6 +2119,126 @@ async def get_org_object(
)
async def _get_resources_from_access_groups(
access_group_ids: List[str],
resource_field: Literal[
"access_model_names", "access_mcp_server_ids", "access_agent_ids"
],
prisma_client: Optional[PrismaClient] = None,
user_api_key_cache: Optional[DualCache] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> List[str]:
"""
Fetch access groups by their IDs (from cache or DB) and collect
the specified resource field across all of them.
Args:
access_group_ids: List of access group IDs to fetch
resource_field: Which resource list to extract from each access group
- "access_model_names": model names (for model access checks)
- "access_mcp_server_ids": MCP server IDs (for MCP access checks)
- "access_agent_ids": agent IDs (for agent access checks)
prisma_client: Optional PrismaClient (lazy-imported from proxy_server if None)
user_api_key_cache: Optional DualCache (lazy-imported from proxy_server if None)
proxy_logging_obj: Optional ProxyLogging (lazy-imported from proxy_server if None)
Returns:
Deduplicated list of resource identifiers from all resolved access groups.
"""
if not access_group_ids:
return []
# Lazy import to avoid circular imports
if prisma_client is None or user_api_key_cache is None:
from litellm.proxy.proxy_server import (
prisma_client as _prisma_client,
proxy_logging_obj as _proxy_logging_obj,
user_api_key_cache as _user_api_key_cache,
)
prisma_client = prisma_client or _prisma_client
user_api_key_cache = user_api_key_cache or _user_api_key_cache
proxy_logging_obj = proxy_logging_obj or _proxy_logging_obj
if user_api_key_cache is None:
return []
resources: List[str] = []
for ag_id in access_group_ids:
try:
ag = await get_access_object(
access_group_id=ag_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
resources.extend(getattr(ag, resource_field, []))
except Exception:
verbose_proxy_logger.debug(
"Could not fetch access group %s for resource field %s",
ag_id,
resource_field,
)
return list(set(resources))
async def _get_models_from_access_groups(
access_group_ids: List[str],
prisma_client: Optional[PrismaClient] = None,
user_api_key_cache: Optional[DualCache] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> List[str]:
"""
Collect model names from unified access groups.
Models are matched by model name for backwards compatibility.
"""
return await _get_resources_from_access_groups(
access_group_ids=access_group_ids,
resource_field="access_model_names",
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
async def _get_mcp_server_ids_from_access_groups(
access_group_ids: List[str],
prisma_client: Optional[PrismaClient] = None,
user_api_key_cache: Optional[DualCache] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> List[str]:
"""
Collect MCP server IDs from unified access groups.
MCPs are matched by server ID.
"""
return await _get_resources_from_access_groups(
access_group_ids=access_group_ids,
resource_field="access_mcp_server_ids",
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
async def _get_agent_ids_from_access_groups(
access_group_ids: List[str],
prisma_client: Optional[PrismaClient] = None,
user_api_key_cache: Optional[DualCache] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> List[str]:
"""
Collect agent IDs from unified access groups.
Agents are matched by agent ID.
"""
return await _get_resources_from_access_groups(
access_group_ids=access_group_ids,
resource_field="access_agent_ids",
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
def _check_model_access_helper(
model: str,
llm_router: Optional[Router],
@ -2165,20 +2391,41 @@ async def can_key_call_model(
"""
Checks if token can call a given model
1. First checks native key-level model permissions (current implementation)
2. If not allowed natively, falls back to access_group_ids on the key
Returns:
- True: if token allowed to call model
Raises:
- Exception: If token not allowed to call model
"""
return _can_object_call_model(
model=model,
llm_router=llm_router,
models=valid_token.models,
team_model_aliases=valid_token.team_model_aliases,
team_id=valid_token.team_id,
object_type="key",
)
try:
return _can_object_call_model(
model=model,
llm_router=llm_router,
models=valid_token.models,
team_model_aliases=valid_token.team_model_aliases,
team_id=valid_token.team_id,
object_type="key",
)
except ProxyException:
# Fallback: check key's access_group_ids
key_access_group_ids = valid_token.access_group_ids or []
if key_access_group_ids:
models_from_groups = await _get_models_from_access_groups(
access_group_ids=key_access_group_ids,
)
if models_from_groups:
return _can_object_call_model(
model=model,
llm_router=llm_router,
models=models_from_groups,
team_model_aliases=valid_token.team_model_aliases,
team_id=valid_token.team_id,
object_type="key",
)
raise
def can_org_access_model(
@ -2200,7 +2447,7 @@ def can_org_access_model(
)
def can_team_access_model(
async def can_team_access_model(
model: Union[str, List[str]],
team_object: Optional[LiteLLM_TeamTable],
llm_router: Optional[Router],
@ -2209,15 +2456,37 @@ def can_team_access_model(
"""
Returns True if the team can access a specific model.
1. First checks native team-level model permissions (current implementation)
2. If not allowed natively, falls back to access_group_ids on the team
"""
return _can_object_call_model(
model=model,
llm_router=llm_router,
models=team_object.models if team_object else [],
team_model_aliases=team_model_aliases,
team_id=team_object.team_id if team_object else None,
object_type="team",
)
try:
return _can_object_call_model(
model=model,
llm_router=llm_router,
models=team_object.models if team_object else [],
team_model_aliases=team_model_aliases,
team_id=team_object.team_id if team_object else None,
object_type="team",
)
except ProxyException:
# Fallback: check team's access_group_ids
team_access_group_ids = (
(team_object.access_group_ids or []) if team_object else []
)
if team_access_group_ids:
models_from_groups = await _get_models_from_access_groups(
access_group_ids=team_access_group_ids,
)
if models_from_groups:
return _can_object_call_model(
model=model,
llm_router=llm_router,
models=models_from_groups,
team_model_aliases=team_model_aliases,
team_id=team_object.team_id if team_object else None,
object_type="team",
)
raise
async def can_user_call_model(

View File

@ -964,7 +964,7 @@ class JWTAuthManager:
team_models = team_object.models
if isinstance(team_models, list) and (
not requested_model
or can_team_access_model(
or await can_team_access_model(
model=requested_model,
team_object=team_object,
llm_router=llm_router,

View File

@ -1187,18 +1187,27 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# Check 6: Additional Common Checks across jwt + key auth
if valid_token.team_id is not None:
_team_obj: Optional[LiteLLM_TeamTable] = LiteLLM_TeamTable(
team_id=valid_token.team_id,
max_budget=valid_token.team_max_budget,
soft_budget=valid_token.team_soft_budget,
spend=valid_token.team_spend,
tpm_limit=valid_token.team_tpm_limit,
rpm_limit=valid_token.team_rpm_limit,
blocked=valid_token.team_blocked,
models=valid_token.team_models,
metadata=valid_token.team_metadata,
object_permission_id=valid_token.team_object_permission_id,
)
try:
_team_obj = await get_team_object(
team_id=valid_token.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except HTTPException:
_team_obj = LiteLLM_TeamTableCachedObj(
team_id=valid_token.team_id,
max_budget=valid_token.team_max_budget,
soft_budget=valid_token.team_soft_budget,
spend=valid_token.team_spend,
tpm_limit=valid_token.team_tpm_limit,
rpm_limit=valid_token.team_rpm_limit,
blocked=valid_token.team_blocked,
models=valid_token.team_models,
metadata=valid_token.team_metadata,
object_permission_id=valid_token.team_object_permission_id,
)
else:
_team_obj = None

View File

@ -3,7 +3,19 @@ from typing import List
from fastapi import APIRouter, Depends, HTTPException, status
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import (
CommonProxyErrors,
LiteLLM_AccessGroupTable,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import (
_cache_access_object,
_cache_key_object,
_cache_team_object,
_delete_cache_access_object,
_get_team_object_from_cache,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.utils import get_prisma_client_or_throw
@ -43,6 +55,45 @@ def _record_to_response(record) -> AccessGroupResponse:
)
def _record_to_access_group_table(record) -> LiteLLM_AccessGroupTable:
"""Convert a Prisma record to a LiteLLM_AccessGroupTable pydantic object for caching."""
return LiteLLM_AccessGroupTable(**record.dict())
async def _cache_access_group_record(record) -> None:
"""
Cache an access group Prisma record in the user_api_key_cache.
Uses a lazy import of user_api_key_cache and proxy_logging_obj from proxy_server
to avoid circular imports, following the same pattern as key_management_endpoints.
"""
from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
access_group_table = _record_to_access_group_table(record)
await _cache_access_object(
access_group_id=record.access_group_id,
access_group_table=access_group_table,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
async def _invalidate_cache_access_group(access_group_id: str) -> None:
"""
Invalidate (delete) an access group entry from both in-memory and Redis caches.
Uses a lazy import of user_api_key_cache and proxy_logging_obj from proxy_server
to avoid circular imports, following the same pattern as key_management_endpoints.
"""
from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
await _delete_cache_access_object(
access_group_id=access_group_id,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
@router.post(
"/v1/access_group",
response_model=AccessGroupResponse,
@ -87,6 +138,10 @@ async def create_access_group(
detail=f"Access group '{data.access_group_name}' already exists",
)
raise
# Cache the newly created access group for read-heavy access patterns
await _cache_access_group_record(record)
return _record_to_response(record)
@ -166,6 +221,10 @@ async def update_access_group(
detail=f"Access group '{update_data.get('access_group_name', '')}' already exists",
)
raise
# Write the updated record into cache (same key, overwrites stale entry)
await _cache_access_group_record(record)
return _record_to_response(record)
@ -181,6 +240,10 @@ async def delete_access_group(
prisma_client = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
try:
# Track affected team IDs and key tokens for cache invalidation
affected_team_ids: list = []
affected_key_tokens: list = []
async with prisma_client.db.tx() as tx:
existing = await tx.litellm_accessgrouptable.find_unique(
where={"access_group_id": access_group_id}
@ -196,6 +259,7 @@ async def delete_access_group(
where={"access_group_ids": {"hasSome": [access_group_id]}}
)
for team in teams_with_group:
affected_team_ids.append(team.team_id)
updated_ids = [tid for tid in (team.access_group_ids or []) if tid != access_group_id]
await tx.litellm_teamtable.update(
where={"team_id": team.team_id},
@ -206,6 +270,7 @@ async def delete_access_group(
where={"access_group_ids": {"hasSome": [access_group_id]}}
)
for key in keys_with_group:
affected_key_tokens.append(key.token)
updated_ids = [kid for kid in (key.access_group_ids or []) if kid != access_group_id]
await tx.litellm_verificationtoken.update(
where={"token": key.token},
@ -215,6 +280,48 @@ async def delete_access_group(
await tx.litellm_accessgrouptable.delete(
where={"access_group_id": access_group_id}
)
# Invalidate the deleted access group from cache
await _invalidate_cache_access_group(access_group_id)
# Patch cached team and key objects to remove the deleted access_group_id
# instead of fully invalidating them (keeps cache warm, avoids DB re-fetch)
from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
for team_id in affected_team_ids:
cached_team = await _get_team_object_from_cache(
key="team_id:{}".format(team_id),
proxy_logging_obj=proxy_logging_obj,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
)
if cached_team is not None and cached_team.access_group_ids:
cached_team.access_group_ids = [
ag_id for ag_id in cached_team.access_group_ids if ag_id != access_group_id
]
await _cache_team_object(
team_id=team_id,
team_table=cached_team,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
for token in affected_key_tokens:
cached_key = await user_api_key_cache.async_get_cache(key=token)
if cached_key is not None:
if isinstance(cached_key, dict):
cached_key = UserAPIKeyAuth(**cached_key)
if isinstance(cached_key, UserAPIKeyAuth) and cached_key.access_group_ids:
cached_key.access_group_ids = [
ag_id for ag_id in cached_key.access_group_ids if ag_id != access_group_id
]
await _cache_key_object(
hashed_token=token,
user_api_key_obj=cached_key,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
except HTTPException:
raise
except Exception as e:

View File

@ -1599,7 +1599,7 @@ async def _process_single_key_update(
status_code=500,
detail={"error": "Team object not found for team change validation"},
)
validate_key_team_change(
await validate_key_team_change(
key=existing_key_row,
team=team_obj,
change_initiated_by=user_api_key_dict,
@ -1828,7 +1828,7 @@ async def update_key_fn(
"error": "Team object not found for team change validation"
},
)
validate_key_team_change(
await validate_key_team_change(
key=existing_key_row,
team=team_obj,
change_initiated_by=user_api_key_dict,
@ -2062,7 +2062,7 @@ async def bulk_update_keys(
)
def validate_key_team_change(
async def validate_key_team_change(
key: LiteLLM_VerificationToken,
team: LiteLLM_TeamTable,
change_initiated_by: UserAPIKeyAuth,
@ -2079,7 +2079,7 @@ def validate_key_team_change(
# Check if the team has access to the key's models
if len(key.models) > 0:
for model in key.models:
can_team_access_model(
await can_team_access_model(
model=model,
team_object=team,
llm_router=llm_router,
@ -2479,6 +2479,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
auto_rotate: Optional[bool] = None,
rotation_interval: Optional[str] = None,
router_settings: Optional[dict] = None,
access_group_ids: Optional[list] = None,
):
from litellm.proxy.proxy_server import premium_user, prisma_client
@ -2595,6 +2596,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
"allowed_routes": allowed_routes or [],
"object_permission_id": object_permission_id,
"router_settings": router_settings_json,
"access_group_ids": access_group_ids or [],
}
# Add rotation fields if auto_rotate is enabled

View File

@ -413,7 +413,7 @@ async def test_can_team_access_model(model, team_models, expect_to_work):
team_id="test-team",
models=team_models,
)
result = can_team_access_model(
result = await can_team_access_model(
model=model,
team_object=team_object,
llm_router=None,
@ -754,3 +754,225 @@ async def test_can_key_call_model_with_aliases(model, alias_map, expect_to_work)
valid_token=user_api_key_object,
llm_router=router,
)
# ---------------------------------------------------------------------------
# Access group cache helpers (_cache_access_object, _delete_cache_access_object)
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_cache_access_object():
"""Test _cache_access_object stores access group in cache with correct key."""
from litellm.proxy.auth.auth_checks import _cache_access_object
from litellm.proxy._types import LiteLLM_AccessGroupTable
cache = DualCache()
ag_id = "ag-test-123"
ag_table = LiteLLM_AccessGroupTable(
access_group_id=ag_id,
access_group_name="test-group",
access_model_names=["gpt-4"],
)
await _cache_access_object(
access_group_id=ag_id,
access_group_table=ag_table,
user_api_key_cache=cache,
)
cached = await cache.async_get_cache(key=f"access_group_id:{ag_id}")
assert cached is not None
if isinstance(cached, dict):
assert cached.get("access_group_id") == ag_id
assert cached.get("access_group_name") == "test-group"
else:
assert cached.access_group_id == ag_id
assert cached.access_group_name == "test-group"
@pytest.mark.asyncio
async def test_delete_cache_access_object():
"""Test _delete_cache_access_object removes access group from in-memory cache."""
from litellm.proxy.auth.auth_checks import _delete_cache_access_object
from litellm.proxy._types import LiteLLM_AccessGroupTable
cache = DualCache()
ag_id = "ag-delete-test"
ag_table = LiteLLM_AccessGroupTable(
access_group_id=ag_id,
access_group_name="to-delete",
)
await cache.async_set_cache(key=f"access_group_id:{ag_id}", value=ag_table, ttl=60)
await _delete_cache_access_object(access_group_id=ag_id, user_api_key_cache=cache)
cached = await cache.async_get_cache(key=f"access_group_id:{ag_id}")
assert cached is None
# ---------------------------------------------------------------------------
# Access group resource fetchers (_get_models_from_access_groups, _get_agent_ids_from_access_groups)
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
"resource_field, access_group_data, expected",
[
(
"access_model_names",
{"access_group_id": "ag-1", "access_model_names": ["gpt-4", "claude-3"]},
["gpt-4", "claude-3"],
),
(
"access_agent_ids",
{"access_group_id": "ag-2", "access_agent_ids": ["agent-a", "agent-b"]},
["agent-a", "agent-b"],
),
(
"access_model_names",
{"access_group_id": "ag-3", "access_model_names": []},
[],
),
],
)
@pytest.mark.asyncio
async def test_get_resources_from_access_groups(resource_field, access_group_data, expected):
"""Test _get_resources_from_access_groups returns correct resource list from access groups."""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import LiteLLM_AccessGroupTable
from litellm.proxy.auth.auth_checks import (
_get_agent_ids_from_access_groups,
_get_models_from_access_groups,
)
ag_table = LiteLLM_AccessGroupTable(
access_group_id=access_group_data["access_group_id"],
access_group_name="test",
access_model_names=access_group_data.get("access_model_names", []),
access_agent_ids=access_group_data.get("access_agent_ids", []),
)
with patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=ag_table,
):
if resource_field == "access_model_names":
result = await _get_models_from_access_groups(
access_group_ids=[access_group_data["access_group_id"]],
prisma_client=MagicMock(),
user_api_key_cache=DualCache(),
)
else:
result = await _get_agent_ids_from_access_groups(
access_group_ids=[access_group_data["access_group_id"]],
prisma_client=MagicMock(),
user_api_key_cache=DualCache(),
)
assert sorted(result) == sorted(expected)
@pytest.mark.asyncio
async def test_get_models_from_access_groups_empty_ids():
"""Test _get_models_from_access_groups returns empty list when access_group_ids is empty."""
from litellm.proxy.auth.auth_checks import _get_models_from_access_groups
result = await _get_models_from_access_groups(access_group_ids=[])
assert result == []
# ---------------------------------------------------------------------------
# can_team_access_model with access_group_ids fallback
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_can_team_access_model_via_access_group_ids():
"""Test can_team_access_model allows access when team has access_group_ids granting model access."""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import can_team_access_model
team_object = LiteLLM_TeamTable(
team_id="test-team",
models=[],
access_group_ids=["ag-with-gpt4"],
)
with patch(
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
new_callable=AsyncMock,
return_value=["gpt-4"],
):
result = await can_team_access_model(
model="gpt-4",
team_object=team_object,
llm_router=None,
team_model_aliases=None,
)
assert result is True
@pytest.mark.asyncio
async def test_can_team_access_model_access_group_ids_denied():
"""Test can_team_access_model denies when neither team models nor access_group_ids grant access."""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import can_team_access_model
from litellm.proxy._types import ProxyException
team_object = LiteLLM_TeamTable(
team_id="test-team",
models=["gpt-3.5-turbo"],
access_group_ids=["ag-other"],
)
with patch(
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
new_callable=AsyncMock,
return_value=["claude-3"],
):
with pytest.raises(ProxyException):
await can_team_access_model(
model="gpt-4",
team_object=team_object,
llm_router=None,
team_model_aliases=None,
)
# ---------------------------------------------------------------------------
# can_key_call_model with access_group_ids fallback
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_can_key_call_model_via_access_group_ids():
"""Test can_key_call_model allows access when key has access_group_ids granting model access."""
from unittest.mock import AsyncMock, patch
from litellm.proxy.auth.auth_checks import can_key_call_model
user_api_key_object = UserAPIKeyAuth(
token="test-token",
models=[],
access_group_ids=["ag-with-gpt4"],
)
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "openai/gpt-4", "api_key": "test"},
}
]
)
with patch(
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
new_callable=AsyncMock,
return_value=["gpt-4"],
):
await can_key_call_model(
model="gpt-4",
llm_model_list=[],
valid_token=user_api_key_object,
llm_router=router,
)

View File

@ -4,7 +4,7 @@ Unit tests for AgentRequestHandler - Agent permission management for keys and te
import os
import sys
from unittest.mock import patch
from unittest.mock import AsyncMock, patch
import pytest
@ -111,3 +111,57 @@ class TestAgentRequestHandler:
result = await AgentRequestHandler.get_allowed_agents(user_api_key_auth=mock_user_auth)
assert result == []
async def test_get_allowed_agents_for_key_via_access_group_ids(self):
"""
Test that _get_allowed_agents_for_key includes agents from key's access_group_ids
(unified access groups) when key has no native object_permission.
"""
mock_user_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
access_group_ids=["ag-with-agents"],
)
with patch.object(
AgentRequestHandler, "_get_key_object_permission", return_value=None
):
with patch(
"litellm.proxy.auth.auth_checks._get_agent_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["agent-from-ag-1", "agent-from-ag-2"],
):
result = await AgentRequestHandler._get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
)
assert sorted(result) == ["agent-from-ag-1", "agent-from-ag-2"]
async def test_get_allowed_agents_for_key_combines_native_and_access_groups(self):
"""
Test that _get_allowed_agents_for_key combines agents from native object_permission
and key's access_group_ids (unified access groups).
"""
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
mock_permission = LiteLLM_ObjectPermissionTable(
object_permission_id="obj-1",
agents=["native-agent-1"],
agent_access_groups=[],
)
mock_user_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
access_group_ids=["ag-1"],
)
# Attach object_permission so _get_key_object_permission returns it
mock_user_auth.object_permission = mock_permission
with patch(
"litellm.proxy.auth.auth_checks._get_agent_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["agent-from-ag"],
):
result = await AgentRequestHandler._get_allowed_agents_for_key(
user_api_key_auth=mock_user_auth
)
assert sorted(result) == ["agent-from-ag", "native-agent-1"]

View File

@ -37,19 +37,27 @@ def _make_access_group_record(
updated_by: str | None = "admin-user",
created_at: datetime | None = None,
):
created_at_val = created_at or datetime.now()
updated_at_val = datetime.now()
data = {
"access_group_id": access_group_id,
"access_group_name": access_group_name,
"description": description,
"access_model_names": access_model_names or [],
"access_mcp_server_ids": access_mcp_server_ids or [],
"access_agent_ids": access_agent_ids or [],
"assigned_team_ids": assigned_team_ids or [],
"assigned_key_ids": assigned_key_ids or [],
"created_at": created_at_val,
"created_by": created_by,
"updated_at": updated_at_val,
"updated_by": updated_by,
}
record = MagicMock()
record.access_group_id = access_group_id
record.access_group_name = access_group_name
record.description = description
record.access_model_names = access_model_names or []
record.access_mcp_server_ids = access_mcp_server_ids or []
record.access_agent_ids = access_agent_ids or []
record.assigned_team_ids = assigned_team_ids or []
record.assigned_key_ids = assigned_key_ids or []
record.created_at = created_at or datetime.now()
record.created_by = created_by
record.updated_at = datetime.now()
record.updated_by = updated_by
for k, v in data.items():
setattr(record, k, v)
record.dict = lambda: data
record.model_dump = lambda: data
return record
@ -116,6 +124,27 @@ def client_and_mocks(monkeypatch):
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
# Mock user_api_key_cache and proxy_logging_obj for cache operations (create/update/delete)
mock_cache = MagicMock()
mock_cache.async_set_cache = AsyncMock(return_value=None)
mock_cache.async_get_cache = AsyncMock(return_value=None)
mock_cache.delete_cache = MagicMock(return_value=None)
monkeypatch.setattr(ps, "user_api_key_cache", mock_cache)
mock_proxy_logging = MagicMock()
mock_proxy_logging.internal_usage_cache = MagicMock()
mock_proxy_logging.internal_usage_cache.dual_cache = MagicMock()
mock_proxy_logging.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock(
return_value=None
)
mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache = AsyncMock(
return_value=None
)
mock_proxy_logging.internal_usage_cache.dual_cache.async_set_cache = AsyncMock(
return_value=None
)
monkeypatch.setattr(ps, "proxy_logging_obj", mock_proxy_logging)
admin_user = UserAPIKeyAuth(
user_id="admin_user",
user_role=LitellmUserRoles.PROXY_ADMIN,
@ -124,7 +153,7 @@ def client_and_mocks(monkeypatch):
client = TestClient(app)
yield client, mock_prisma, mock_access_group_table
yield client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging
app.dependency_overrides.clear()
monkeypatch.setattr(ps, "prisma_client", ps.prisma_client)
@ -155,7 +184,7 @@ ACCESS_GROUP_PATHS = ["/v1/access_group", "/v1/unified_access_group"]
)
def test_create_access_group_success(client_and_mocks, base_path, payload):
"""Create access group with various payloads returns 201."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
resp = client.post(base_path, json=payload)
assert resp.status_code == 201
@ -167,7 +196,7 @@ def test_create_access_group_success(client_and_mocks, base_path, payload):
def test_create_access_group_duplicate_name_conflict(client_and_mocks):
"""Create with duplicate name returns 409."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_name="existing-group")
mock_table.find_unique = AsyncMock(return_value=existing)
@ -187,7 +216,7 @@ def test_create_access_group_duplicate_name_conflict(client_and_mocks):
)
def test_create_access_group_race_condition_returns_409(client_and_mocks, error_message):
"""Create race condition: Prisma unique constraint surfaces as 409, not 500."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
mock_table.find_unique = AsyncMock(return_value=None)
mock_table.create = AsyncMock(side_effect=Exception(error_message))
@ -200,7 +229,7 @@ def test_create_access_group_race_condition_returns_409(client_and_mocks, error_
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
def test_create_access_group_forbidden_non_admin(client_and_mocks, user_role):
"""Non-admin users cannot create access groups."""
client, _, _ = client_and_mocks
client, *_ = client_and_mocks
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="regular_user",
@ -214,7 +243,7 @@ def test_create_access_group_forbidden_non_admin(client_and_mocks, user_role):
def test_create_access_group_validation_missing_name(client_and_mocks):
"""Create with missing access_group_name returns 422."""
client, _, _ = client_and_mocks
client, *_ = client_and_mocks
resp = client.post("/v1/access_group", json={})
assert resp.status_code == 422
@ -222,7 +251,7 @@ def test_create_access_group_validation_missing_name(client_and_mocks):
def test_create_access_group_500_on_non_constraint_prisma_error(client_and_mocks):
"""Create with non-unique-constraint Prisma error returns 500."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
mock_table.find_unique = AsyncMock(return_value=None)
mock_table.create = AsyncMock(side_effect=Exception("Some other database error"))
@ -241,7 +270,7 @@ def test_create_access_group_500_on_non_constraint_prisma_error(client_and_mocks
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_list_access_groups_success_empty(client_and_mocks, base_path):
"""List access groups returns empty list when none exist."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
resp = client.get(base_path)
assert resp.status_code == 200
@ -252,7 +281,7 @@ def test_list_access_groups_success_empty(client_and_mocks, base_path):
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_list_access_groups_success_with_items(client_and_mocks, base_path):
"""List access groups returns items when they exist."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
records = [
_make_access_group_record(access_group_id="ag-1", access_group_name="group-1"),
@ -271,7 +300,7 @@ def test_list_access_groups_success_with_items(client_and_mocks, base_path):
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_list_access_groups_ordered_by_created_at_desc(client_and_mocks, base_path):
"""List access groups calls find_many with created_at desc order."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
older = datetime(2025, 1, 1, 12, 0, 0)
newer = datetime(2025, 1, 2, 12, 0, 0)
@ -302,7 +331,7 @@ def test_list_access_groups_ordered_by_created_at_desc(client_and_mocks, base_pa
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
def test_list_access_groups_forbidden_non_admin(client_and_mocks, user_role):
"""Non-admin users cannot list access groups."""
client, _, _ = client_and_mocks
client, *_ = client_and_mocks
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="regular_user",
@ -323,7 +352,7 @@ def test_list_access_groups_forbidden_non_admin(client_and_mocks, user_role):
@pytest.mark.parametrize("access_group_id", ["ag-123", "ag-other-id"])
def test_get_access_group_success(client_and_mocks, base_path, access_group_id):
"""Get access group by id returns record when found."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
record = _make_access_group_record(access_group_id=access_group_id)
mock_table.find_unique = AsyncMock(return_value=record)
@ -335,7 +364,7 @@ def test_get_access_group_success(client_and_mocks, base_path, access_group_id):
def test_get_access_group_not_found(client_and_mocks):
"""Get access group returns 404 when not found."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
mock_table.find_unique = AsyncMock(return_value=None)
@ -347,7 +376,7 @@ def test_get_access_group_not_found(client_and_mocks):
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
def test_get_access_group_forbidden_non_admin(client_and_mocks, user_role):
"""Non-admin users cannot get access group."""
client, _, _ = client_and_mocks
client, *_ = client_and_mocks
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="regular_user",
@ -375,7 +404,7 @@ def test_get_access_group_forbidden_non_admin(client_and_mocks, user_role):
)
def test_update_access_group_success(client_and_mocks, base_path, update_payload):
"""Update access group with various payloads returns 200."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_id="ag-update")
mock_table.find_unique = AsyncMock(return_value=existing)
@ -387,7 +416,7 @@ def test_update_access_group_success(client_and_mocks, base_path, update_payload
def test_update_access_group_not_found(client_and_mocks):
"""Update access group returns 404 when not found."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
mock_table.find_unique = AsyncMock(return_value=None)
@ -403,7 +432,7 @@ def test_update_access_group_not_found(client_and_mocks):
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
def test_update_access_group_forbidden_non_admin(client_and_mocks, user_role):
"""Non-admin users cannot update access groups."""
client, _, _ = client_and_mocks
client, *_ = client_and_mocks
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="regular_user",
@ -417,7 +446,7 @@ def test_update_access_group_forbidden_non_admin(client_and_mocks, user_role):
def test_update_access_group_empty_body(client_and_mocks):
"""Update with empty body succeeds; only updated_by is set."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="unchanged")
mock_table.find_unique = AsyncMock(return_value=existing)
@ -433,7 +462,7 @@ def test_update_access_group_empty_body(client_and_mocks):
def test_update_access_group_name_success(client_and_mocks):
"""Update access_group_name succeeds when new name is unique."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="old-name")
mock_table.find_unique = AsyncMock(return_value=existing)
@ -447,7 +476,7 @@ def test_update_access_group_name_success(client_and_mocks):
def test_update_access_group_name_duplicate_conflict(client_and_mocks):
"""Update access_group_name to existing name returns 409 (unique constraint)."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="old-name")
mock_table.find_unique = AsyncMock(return_value=existing)
@ -471,7 +500,7 @@ def test_update_access_group_name_duplicate_conflict(client_and_mocks):
)
def test_update_access_group_name_unique_constraint_returns_409(client_and_mocks, error_message):
"""Update access_group_name: Prisma unique constraint surfaces as 409."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="old-name")
mock_table.find_unique = AsyncMock(return_value=existing)
@ -491,7 +520,7 @@ def test_update_access_group_name_unique_constraint_returns_409(client_and_mocks
@pytest.mark.parametrize("access_group_id", ["ag-123", "ag-delete-me"])
def test_delete_access_group_success(client_and_mocks, base_path, access_group_id):
"""Delete access group returns 204 when found."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_id=access_group_id)
mock_table.find_unique = AsyncMock(return_value=existing)
@ -503,7 +532,7 @@ def test_delete_access_group_success(client_and_mocks, base_path, access_group_i
def test_delete_access_group_not_found(client_and_mocks):
"""Delete access group returns 404 when not found."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
mock_table.find_unique = AsyncMock(return_value=None)
@ -516,7 +545,7 @@ def test_delete_access_group_not_found(client_and_mocks):
@pytest.mark.parametrize("user_role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
def test_delete_access_group_forbidden_non_admin(client_and_mocks, user_role):
"""Non-admin users cannot delete access groups."""
client, _, _ = client_and_mocks
client, *_ = client_and_mocks
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="regular_user",
@ -530,7 +559,7 @@ def test_delete_access_group_forbidden_non_admin(client_and_mocks, user_role):
def test_delete_access_group_cleans_up_teams_and_keys(client_and_mocks):
"""Delete removes access_group_id from teams and keys before deleting the group."""
client, mock_prisma, mock_access_group_table = client_and_mocks
client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
mock_key_table = mock_prisma.db.litellm_verificationtoken
@ -563,9 +592,208 @@ def test_delete_access_group_cleans_up_teams_and_keys(client_and_mocks):
)
@pytest.mark.parametrize(
"team_cache_group_ids,key_cache_group_ids,expected_team_ids_after,expected_key_ids_after",
[
# Team and key both cached with the deleted group
(
["ag-to-delete", "ag-keep"],
["ag-to-delete", "ag-stay"],
["ag-keep"],
["ag-stay"],
),
# Only team cached; key not in cache
(
["ag-to-delete"],
None,
[],
None,
),
# Only key cached; team not in cache
(
None,
["ag-to-delete"],
None,
[],
),
# Neither cached — nothing to patch
(
None,
None,
None,
None,
),
# Cached team has only the deleted group
(
["ag-to-delete"],
["ag-to-delete"],
[],
[],
),
# Cached objects have multiple groups, only the deleted one is removed
(
["ag-alpha", "ag-to-delete", "ag-beta"],
["ag-to-delete", "ag-gamma"],
["ag-alpha", "ag-beta"],
["ag-gamma"],
),
],
ids=[
"both_cached",
"only_team_cached",
"only_key_cached",
"neither_cached",
"single_group_removed",
"multi_group_partial_removal",
],
)
def test_delete_access_group_patches_cached_team_and_key(
client_and_mocks,
team_cache_group_ids,
key_cache_group_ids,
expected_team_ids_after,
expected_key_ids_after,
):
"""Delete patches cached team/key objects to remove the deleted access_group_id."""
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
mock_key_table = mock_prisma.db.litellm_verificationtoken
existing = _make_access_group_record(access_group_id="ag-to-delete")
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
# Set up a team and key in the DB that reference the group
team_with_group = MagicMock()
team_with_group.team_id = "team-1"
team_with_group.access_group_ids = ["ag-to-delete", "ag-keep"]
mock_team_table.find_many = AsyncMock(return_value=[team_with_group])
key_with_group = MagicMock()
key_with_group.token = "hashed-key-1"
key_with_group.access_group_ids = ["ag-to-delete"]
mock_key_table.find_many = AsyncMock(return_value=[key_with_group])
# Build cached team object (returned from proxy_logging dual cache)
if team_cache_group_ids is not None:
cached_team = LiteLLM_TeamTableCachedObj(
team_id="team-1",
access_group_ids=list(team_cache_group_ids),
)
mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache = AsyncMock(
return_value=cached_team
)
else:
mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache = AsyncMock(
return_value=None
)
# Build cached key object (returned from user_api_key_cache)
if key_cache_group_ids is not None:
cached_key = UserAPIKeyAuth(
token="hashed-key-1",
access_group_ids=list(key_cache_group_ids),
)
mock_cache.async_get_cache = AsyncMock(return_value=cached_key)
else:
mock_cache.async_get_cache = AsyncMock(return_value=None)
resp = client.delete("/v1/access_group/ag-to-delete")
assert resp.status_code == 204
# Verify DB cleanup always happens
mock_team_table.update.assert_awaited_once()
mock_key_table.update.assert_awaited_once()
# Verify cache patching
if expected_team_ids_after is not None:
# _cache_team_object writes via _cache_management_object -> async_set_cache
team_set_calls = [
c for c in mock_cache.async_set_cache.call_args_list
if c.kwargs.get("key", "") == "team_id:team-1"
or (len(c.args) >= 1 and c.args[0] == "team_id:team-1")
]
assert len(team_set_calls) >= 1, "Expected team cache to be patched"
# The cached team object should have the updated access_group_ids
written_team = team_set_calls[0].kwargs.get("value") or team_set_calls[0].args[1]
if isinstance(written_team, LiteLLM_TeamTableCachedObj):
assert written_team.access_group_ids == expected_team_ids_after
else:
# No team in cache — async_set_cache should not be called for team_id key
team_set_calls = [
c for c in mock_cache.async_set_cache.call_args_list
if c.kwargs.get("key", "") == "team_id:team-1"
or (len(c.args) >= 1 and c.args[0] == "team_id:team-1")
]
assert len(team_set_calls) == 0, "Should not patch team cache when not cached"
if expected_key_ids_after is not None:
key_set_calls = [
c for c in mock_cache.async_set_cache.call_args_list
if c.kwargs.get("key", "") == "hashed-key-1"
or (len(c.args) >= 1 and c.args[0] == "hashed-key-1")
]
assert len(key_set_calls) >= 1, "Expected key cache to be patched"
written_key = key_set_calls[0].kwargs.get("value") or key_set_calls[0].args[1]
if isinstance(written_key, UserAPIKeyAuth):
assert written_key.access_group_ids == expected_key_ids_after
else:
key_set_calls = [
c for c in mock_cache.async_set_cache.call_args_list
if c.kwargs.get("key", "") == "hashed-key-1"
or (len(c.args) >= 1 and c.args[0] == "hashed-key-1")
]
assert len(key_set_calls) == 0, "Should not patch key cache when not cached"
def test_delete_access_group_patches_key_cached_as_dict(client_and_mocks):
"""Delete correctly patches a key cached as a raw dict (not UserAPIKeyAuth)."""
client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
mock_key_table = mock_prisma.db.litellm_verificationtoken
existing = _make_access_group_record(access_group_id="ag-to-delete")
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
mock_team_table.find_many = AsyncMock(return_value=[])
key_with_group = MagicMock()
key_with_group.token = "hashed-key-dict"
key_with_group.access_group_ids = ["ag-to-delete", "ag-other"]
mock_key_table.find_many = AsyncMock(return_value=[key_with_group])
# No team in cache
mock_proxy_logging.internal_usage_cache.dual_cache.async_get_cache = AsyncMock(
return_value=None
)
# Key cached as a plain dict (as can happen with Redis serialization)
mock_cache.async_get_cache = AsyncMock(
return_value={
"token": "hashed-key-dict",
"access_group_ids": ["ag-to-delete", "ag-other"],
}
)
resp = client.delete("/v1/access_group/ag-to-delete")
assert resp.status_code == 204
# The key should have been re-cached with the deleted group removed
key_set_calls = [
c for c in mock_cache.async_set_cache.call_args_list
if c.kwargs.get("key", "") == "hashed-key-dict"
or (len(c.args) >= 1 and c.args[0] == "hashed-key-dict")
]
assert len(key_set_calls) >= 1, "Expected key cache to be patched"
written_key = key_set_calls[0].kwargs.get("value") or key_set_calls[0].args[1]
if isinstance(written_key, UserAPIKeyAuth):
assert written_key.access_group_ids == ["ag-other"]
def test_delete_access_group_503_on_db_connection_error(client_and_mocks):
"""Delete returns 503 when DB connection error occurs during transaction."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_id="ag-to-delete")
mock_table.find_unique = AsyncMock(return_value=existing)
@ -578,7 +806,7 @@ def test_delete_access_group_503_on_db_connection_error(client_and_mocks):
def test_delete_access_group_404_on_p2025_or_record_not_found(client_and_mocks):
"""Delete returns 404 when Prisma raises P2025 or record-not-found error."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_id="ag-to-delete")
mock_table.find_unique = AsyncMock(return_value=existing)
@ -591,7 +819,7 @@ def test_delete_access_group_404_on_p2025_or_record_not_found(client_and_mocks):
def test_delete_access_group_500_on_generic_exception(client_and_mocks):
"""Delete returns 500 when generic exception occurs during transaction."""
client, _, mock_table = client_and_mocks
client, _, mock_table, *_ = client_and_mocks
existing = _make_access_group_record(access_group_id="ag-to-delete")
mock_table.find_unique = AsyncMock(return_value=existing)
@ -625,10 +853,32 @@ def test_delete_access_group_500_on_generic_exception(client_and_mocks):
)
def test_access_group_endpoints_db_not_connected(client_and_mocks, monkeypatch, method, url, factory):
"""All endpoints return 500 when DB is not connected."""
client, _, _ = client_and_mocks
client, *_ = client_and_mocks
monkeypatch.setattr(ps, "prisma_client", None)
resp = getattr(client, method)(url, **factory())
assert resp.status_code == 500
assert resp.json()["detail"]["error"] == CommonProxyErrors.db_not_connected_error.value
# ---------------------------------------------------------------------------
# Unit tests for cache helpers (_record_to_access_group_table)
# ---------------------------------------------------------------------------
def test_record_to_access_group_table():
"""Test _record_to_access_group_table converts Prisma-like record to LiteLLM_AccessGroupTable."""
from litellm.proxy.management_endpoints.access_group_endpoints import _record_to_access_group_table
record = _make_access_group_record(
access_group_id="ag-unit-test",
access_group_name="unit-test-group",
access_model_names=["gpt-4", "claude-3"],
access_agent_ids=["agent-1"],
)
result = _record_to_access_group_table(record)
assert result.access_group_id == "ag-unit-test"
assert result.access_group_name == "unit-test-group"
assert result.access_model_names == ["gpt-4", "claude-3"]
assert result.access_agent_ids == ["agent-1"]

View File

@ -520,6 +520,51 @@ async def test_key_generation_with_object_permission(monkeypatch):
assert key_insert_calls[0]["data"].get("object_permission_id") == "objperm123"
@pytest.mark.asyncio
async def test_generate_key_helper_fn_with_access_group_ids(monkeypatch):
"""Ensure generate_key_helper_fn passes access_group_ids into the key insert payload."""
mock_prisma_client = AsyncMock()
mock_prisma_client.jsonify_object = lambda data: data # type: ignore
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_objectpermissiontable = MagicMock()
mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock(
return_value=MagicMock(object_permission_id=None)
)
captured_key_data = {}
async def _insert_data_side_effect(*args, **kwargs):
table_name = kwargs.get("table_name")
if table_name == "user":
return MagicMock(models=[], spend=0)
elif table_name == "key":
captured_key_data.update(kwargs.get("data", {}))
return MagicMock(
token="hashed_token_789",
litellm_budget_table=None,
object_permission=None,
created_at=None,
updated_at=None,
)
return MagicMock()
mock_prisma_client.insert_data = AsyncMock(side_effect=_insert_data_side_effect)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_helper_fn,
)
await generate_key_helper_fn(
request_type="key",
table_name="key",
user_id="test-user",
access_group_ids=["ag-1", "ag-2"],
)
assert captured_key_data.get("access_group_ids") == ["ag-1", "ag-2"]
@pytest.mark.asyncio
async def test_key_generation_with_mcp_tool_permissions(monkeypatch):
"""
@ -1356,14 +1401,15 @@ async def test_unblock_key_invalid_key_format(monkeypatch):
assert "Invalid key format" in str(exc_info.value.message)
def test_validate_key_team_change_with_member_permissions():
@pytest.mark.asyncio
async def test_validate_key_team_change_with_member_permissions():
"""
Test validate_key_team_change function with team member permissions.
This test covers the new logic that allows team members with specific
permissions to update keys, not just team admins.
"""
from unittest.mock import MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import KeyManagementRoutes
@ -1389,7 +1435,8 @@ def test_validate_key_team_change_with_member_permissions():
mock_member_object = MagicMock()
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.can_team_access_model"
"litellm.proxy.management_endpoints.key_management_endpoints.can_team_access_model",
new_callable=AsyncMock,
):
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._get_user_in_team"
@ -1406,7 +1453,7 @@ def test_validate_key_team_change_with_member_permissions():
mock_has_perms.return_value = True
# This should not raise an exception due to member permissions
validate_key_team_change(
await validate_key_team_change(
key=mock_key,
team=mock_team,
change_initiated_by=mock_change_initiator,

View File

@ -3,7 +3,7 @@
"version": "0.1.0",
"private": true,
"scripts": {
"dev": "next dev",
"dev": "next dev --webpack",
"build": "next build",
"start": "next start",
"lint": "next lint",

View File

@ -41,7 +41,7 @@ const mockAccessGroups: AccessGroupResponse[] = [
access_group_id: "ag-1",
access_group_name: "Group One",
description: "First group",
access_model_ids: [],
access_model_names: [],
access_mcp_server_ids: [],
access_agent_ids: [],
assigned_team_ids: [],

View File

@ -15,7 +15,7 @@ export interface AccessGroupResponse {
access_group_id: string;
access_group_name: string;
description: string | null;
access_model_ids: string[];
access_model_names: string[];
access_mcp_server_ids: string[];
access_agent_ids: string[];
assigned_team_ids: string[];

View File

@ -13,7 +13,7 @@ import { AccessGroupResponse, accessGroupKeys } from "./useAccessGroups";
export interface AccessGroupCreateParams {
access_group_name: string;
description?: string | null;
access_model_ids?: string[];
access_model_names?: string[];
access_mcp_server_ids?: string[];
access_agent_ids?: string[];
assigned_team_ids?: string[];

View File

@ -13,7 +13,7 @@ import { AccessGroupResponse, accessGroupKeys } from "./useAccessGroups";
export interface AccessGroupUpdateParams {
access_group_name?: string;
description?: string | null;
access_model_ids?: string[];
access_model_names?: string[];
access_mcp_server_ids?: string[];
access_agent_ids?: string[];
assigned_team_ids?: string[];

View File

@ -56,7 +56,7 @@ const createMockAccessGroup = (
access_group_id: "ag-1",
access_group_name: "Test Group",
description: "A test access group",
access_model_ids: ["model-1", "model-2"],
access_model_names: ["model-1", "model-2"],
access_mcp_server_ids: ["mcp-1"],
access_agent_ids: ["agent-1"],
assigned_team_ids: ["team-1"],
@ -319,7 +319,7 @@ describe("AccessGroupDetail", () => {
it("should show empty state in Models tab when no models assigned", () => {
mockUseAccessGroupDetails.mockReturnValue({
...baseMockReturnValue,
data: createMockAccessGroup({ access_model_ids: [] }),
data: createMockAccessGroup({ access_model_names: [] }),
} as ReturnType<typeof useAccessGroupDetails>);
renderWithProviders(

View File

@ -83,7 +83,7 @@ export function AccessGroupDetail({
);
}
const modelIds = accessGroup.access_model_ids ?? [];
const modelIds = accessGroup.access_model_names ?? [];
const mcpServerIds = accessGroup.access_mcp_server_ids ?? [];
const agentIds = accessGroup.access_agent_ids ?? [];
const keyIds = accessGroup.assigned_key_ids ?? [];

View File

@ -30,7 +30,7 @@ export function AccessGroupCreateModal({
const params: AccessGroupCreateParams = {
access_group_name: values.name,
description: values.description,
access_model_ids: values.modelIds,
access_model_names: values.modelIds,
access_mcp_server_ids: values.mcpServerIds,
access_agent_ids: values.agentIds,
};

View File

@ -32,7 +32,7 @@ export function AccessGroupEditModal({
form.setFieldsValue({
name: accessGroup.access_group_name,
description: accessGroup.description ?? "",
modelIds: accessGroup.access_model_ids ?? [],
modelIds: accessGroup.access_model_names ?? [],
mcpServerIds: accessGroup.access_mcp_server_ids ?? [],
agentIds: accessGroup.access_agent_ids ?? [],
});
@ -46,7 +46,7 @@ export function AccessGroupEditModal({
const params: AccessGroupUpdateParams = {
access_group_name: values.name,
description: values.description,
access_model_ids: values.modelIds,
access_model_names: values.modelIds,
access_mcp_server_ids: values.mcpServerIds,
access_agent_ids: values.agentIds,
};

View File

@ -9,7 +9,7 @@ const mockAccessGroups: AccessGroupResponse[] = [
access_group_id: "ag-1",
access_group_name: "Admin Group",
description: "Administrators with full access",
access_model_ids: ["m1", "m2"],
access_model_names: ["m1", "m2"],
access_mcp_server_ids: ["s1"],
access_agent_ids: ["a1"],
assigned_team_ids: [],
@ -23,7 +23,7 @@ const mockAccessGroups: AccessGroupResponse[] = [
access_group_id: "ag-2",
access_group_name: "Read Only",
description: "Read-only access to models",
access_model_ids: ["m1"],
access_model_names: ["m1"],
access_mcp_server_ids: [],
access_agent_ids: [],
assigned_team_ids: [],

View File

@ -59,7 +59,7 @@ function mapResponseToAccessGroup(r: AccessGroupResponse): AccessGroup {
id: r.access_group_id,
name: r.access_group_name,
description: r.description ?? "",
modelIds: r.access_model_ids,
modelIds: r.access_model_names,
mcpServerIds: r.access_mcp_server_ids,
agentIds: r.access_agent_ids,
keyIds: r.assigned_key_ids,
@ -199,29 +199,32 @@ export function AccessGroupsPage() {
enableSorting: false,
cell: ({ row }) => {
const record = row.original;
const modelIds = record.modelIds ?? [];
const mcpServerIds = record.mcpServerIds ?? [];
const agentIds = record.agentIds ?? [];
return (
<Flex gap={12} align="center">
<Tooltip title={`${record.modelIds.length} Models`}>
<Tooltip title={`${modelIds.length} Models`}>
<Tag color="blue" style={{ fontSize: 14, padding: "2px 8px", margin: 0 }}>
<Flex align="center" gap={6}>
<LayersIcon size={14} />
{record.modelIds.length}
{modelIds.length}
</Flex>
</Tag>
</Tooltip>
<Tooltip title={`${record.mcpServerIds.length} MCP Servers`}>
<Tooltip title={`${mcpServerIds.length} MCP Servers`}>
<Tag color="cyan" style={{ fontSize: 14, padding: "2px 8px", margin: 0 }}>
<Flex align="center" gap={6}>
<ServerIcon size={14} />
{record.mcpServerIds.length}
{mcpServerIds.length}
</Flex>
</Tag>
</Tooltip>
<Tooltip title={`${record.agentIds.length} Agents`}>
<Tooltip title={`${agentIds.length} Agents`}>
<Tag color="purple" style={{ fontSize: 14, padding: "2px 8px", margin: 0 }}>
<Flex align="center" gap={6}>
<BotIcon size={14} />
{record.agentIds.length}
{agentIds.length}
</Flex>
</Tag>
</Tooltip>

View File

@ -14,7 +14,7 @@
"moduleResolution": "bundler",
"resolveJsonModule": true,
"isolatedModules": true,
"jsx": "preserve",
"jsx": "react-jsx",
"incremental": true,
"plugins": [
{