added access groups permission checks
This commit is contained in:
parent
1828e68515
commit
cd148dcb82
@ -1342,6 +1342,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 = [
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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,7 +197,10 @@ 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.
|
||||
|
||||
1. First checks native team-level agent permissions (object_permission)
|
||||
2. Also includes agents from team's access_group_ids (unified access groups)
|
||||
|
||||
Note: object_permission is already loaded by get_team_object() in main auth flow.
|
||||
"""
|
||||
@ -193,26 +211,54 @@ class AgentRequestHandler:
|
||||
return []
|
||||
|
||||
try:
|
||||
# Get team object permission (already loaded in main auth flow)
|
||||
all_agents: List[str] = []
|
||||
|
||||
# 1. Get agents from object_permission (native permissions)
|
||||
object_permissions = await AgentRequestHandler._get_team_object_permission(
|
||||
user_api_key_auth
|
||||
)
|
||||
|
||||
if object_permissions is None:
|
||||
return []
|
||||
if object_permissions is not None:
|
||||
# Get direct agents
|
||||
direct_agents = object_permissions.agents or []
|
||||
|
||||
# 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 []
|
||||
# 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. Fallback: get agent IDs from team's access_group_ids (unified access groups)
|
||||
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,
|
||||
)
|
||||
|
||||
# Combine both lists
|
||||
all_agents = direct_agents + access_group_agents
|
||||
if prisma_client is not None:
|
||||
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,
|
||||
)
|
||||
if team_obj is not None:
|
||||
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)}")
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -3,7 +3,16 @@ 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,
|
||||
_delete_cache_access_object,
|
||||
)
|
||||
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 +52,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 +135,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 +218,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)
|
||||
|
||||
|
||||
@ -215,6 +271,10 @@ 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)
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
||||
@ -1598,7 +1598,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,
|
||||
@ -1826,7 +1826,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,
|
||||
@ -2060,7 +2060,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,
|
||||
@ -2077,7 +2077,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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -1356,14 +1356,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 +1390,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 +1408,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,
|
||||
|
||||
Loading…
Reference in New Issue
Block a user