Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_endpoints/access_group_endpoints.py: 76%
344 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import asyncio
2from collections.abc import Callable, Mapping, Sequence
3from dataclasses import dataclass
4from types import MappingProxyType
5from typing import Final, Protocol
7from fastapi import APIRouter, Depends, HTTPException, status
8from typing_extensions import ReadOnly, TypedDict
10from litellm._logging import verbose_proxy_logger
11from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
12from litellm.proxy._types import (
13 CommonProxyErrors,
14 LiteLLM_AccessGroupTable,
15 LitellmUserRoles,
16 UserAPIKeyAuth,
17)
18from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
19from litellm.proxy.auth.auth_checks import (
20 _cache_access_object,
21 _cache_key_object,
22 _cache_team_object,
23 _get_team_object_from_cache,
24)
25from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
26from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
27from litellm.proxy.management_helpers.access_group_team_sync import invalidate_access_group_cache
28from litellm.proxy.management_helpers.resource_display_names import (
29 agent_display_names,
30 key_display_names,
31 mcp_server_display_names,
32)
33from litellm.proxy.utils import PrismaClient, get_prisma_client_or_throw
34from litellm.repositories.table_repositories import AccessGroupRepository, TeamRepository
35from litellm.types.access_group import (
36 AccessGroupCreateRequest,
37 AccessGroupResource,
38 AccessGroupResponse,
39 AccessGroupUpdateRequest,
40)
42router: Final = APIRouter(
43 tags=["access group management"],
44)
47class _AccessGroupRecord(Protocol):
48 @property
49 def access_group_id(self) -> str: ... 49 ↛ exitline 49 didn't return from function 'access_group_id' because
51 @property
52 def access_mcp_server_ids(self) -> Sequence[str] | None: ... 52 ↛ exitline 52 didn't return from function 'access_mcp_server_ids' because
54 @property
55 def access_agent_ids(self) -> Sequence[str] | None: ... 55 ↛ exitline 55 didn't return from function 'access_agent_ids' because
57 @property
58 def assigned_team_ids(self) -> Sequence[str] | None: ... 58 ↛ exitline 58 didn't return from function 'assigned_team_ids' because
60 @property
61 def assigned_key_ids(self) -> Sequence[str] | None: ... 61 ↛ exitline 61 didn't return from function 'assigned_key_ids' because
63 def dict(self) -> Mapping[str, object]: ... 63 ↛ exitline 63 didn't return from function 'dict' because
66class _TeamRecord(Protocol):
67 @property
68 def team_id(self) -> str: ... 68 ↛ exitline 68 didn't return from function 'team_id' because
70 @property
71 def team_alias(self) -> str | None: ... 71 ↛ exitline 71 didn't return from function 'team_alias' because
73 @property
74 def access_group_ids(self) -> Sequence[str] | None: ... 74 ↛ exitline 74 didn't return from function 'access_group_ids' because
77class _KeyRecord(Protocol):
78 @property
79 def token(self) -> str: ... 79 ↛ exitline 79 didn't return from function 'token' because
81 @property
82 def access_group_ids(self) -> Sequence[str] | None: ... 82 ↛ exitline 82 didn't return from function 'access_group_ids' because
85class _AccessGroupTable(Protocol):
86 async def find_unique(self, where: Mapping[str, object]) -> _AccessGroupRecord | None: ... 86 ↛ exitline 86 didn't return from function 'find_unique' because
88 async def find_many(self, order: Mapping[str, object]) -> Sequence[_AccessGroupRecord]: ... 88 ↛ exitline 88 didn't return from function 'find_many' because
90 async def create(self, data: Mapping[str, object]) -> _AccessGroupRecord: ... 90 ↛ exitline 90 didn't return from function 'create' because
92 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _AccessGroupRecord: ... 92 ↛ exitline 92 didn't return from function 'update' because
94 async def delete(self, where: Mapping[str, object]) -> object: ... 94 ↛ exitline 94 didn't return from function 'delete' because
97class _TeamTable(Protocol):
98 async def find_unique(self, *, where: Mapping[str, object]) -> _TeamRecord | None: ... 98 ↛ exitline 98 didn't return from function 'find_unique' because
100 async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_TeamRecord]: ... 100 ↛ exitline 100 didn't return from function 'find_many' because
102 async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> object: ... 102 ↛ exitline 102 didn't return from function 'update' because
105class _KeyTable(Protocol):
106 async def find_unique(self, where: Mapping[str, object]) -> _KeyRecord | None: ... 106 ↛ exitline 106 didn't return from function 'find_unique' because
108 async def find_many(self, where: Mapping[str, object]) -> Sequence[_KeyRecord]: ... 108 ↛ exitline 108 didn't return from function 'find_many' because
110 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> object: ... 110 ↛ exitline 110 didn't return from function 'update' because
113class _AgentRecord(Protocol):
114 @property
115 def agent_id(self) -> str: ... 115 ↛ exitline 115 didn't return from function 'agent_id' because
117 @property
118 def access_group_ids(self) -> Sequence[str] | None: ... 118 ↛ exitline 118 didn't return from function 'access_group_ids' because
121class _AgentTable(Protocol):
122 async def find_many(self, where: Mapping[str, object]) -> Sequence[_AgentRecord]: ... 122 ↛ exitline 122 didn't return from function 'find_many' because
124 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> object: ... 124 ↛ exitline 124 didn't return from function 'update' because
127class _HasSomeFilter(TypedDict):
128 hasSome: ReadOnly[Sequence[str]]
131class _AgentAccessGroupsWhere(TypedDict):
132 access_group_ids: ReadOnly[_HasSomeFilter]
135class _AgentIdWhere(TypedDict):
136 agent_id: ReadOnly[str]
139class _AgentAccessGroupsData(TypedDict):
140 access_group_ids: ReadOnly[Sequence[str]]
143class _AccessGroupTx(Protocol):
144 @property
145 def litellm_accessgrouptable(self) -> _AccessGroupTable: ... 145 ↛ exitline 145 didn't return from function 'litellm_accessgrouptable' because
147 @property
148 def litellm_teamtable(self) -> _TeamTable: ... 148 ↛ exitline 148 didn't return from function 'litellm_teamtable' because
150 @property
151 def litellm_verificationtoken(self) -> _KeyTable: ... 151 ↛ exitline 151 didn't return from function 'litellm_verificationtoken' because
153 @property
154 def litellm_agentstable(self) -> _AgentTable: ... 154 ↛ exitline 154 didn't return from function 'litellm_agentstable' because
157def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None:
158 if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: 158 ↛ 159line 158 didn't jump to line 159 because the condition on line 158 was never true
159 raise HTTPException(
160 status_code=status.HTTP_403_FORBIDDEN,
161 detail={"error": CommonProxyErrors.not_allowed_access.value},
162 )
165def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None:
166 """Admin Viewer parity: PROXY_ADMIN or PROXY_ADMIN_VIEW_ONLY may read."""
167 from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
169 if not _user_has_admin_view(user_api_key_dict): 169 ↛ 170line 169 didn't jump to line 170 because the condition on line 169 was never true
170 raise HTTPException(
171 status_code=status.HTTP_403_FORBIDDEN,
172 detail={"error": CommonProxyErrors.not_allowed_access.value},
173 )
176@dataclass(frozen=True, slots=True)
177class _ResourceNames:
178 mcp_servers: Mapping[str, str]
179 agents: Mapping[str, str]
180 teams: Mapping[str, str | None]
181 keys: Mapping[str, str]
184def _label(ids: Sequence[str], names: Mapping[str, str | None]) -> tuple[AccessGroupResource, ...]:
185 return tuple(AccessGroupResource(id=resource_id, name=names.get(resource_id)) for resource_id in ids)
188def _record_to_response(
189 record: _AccessGroupRecord, *, assigned_team_ids: Sequence[str], names: _ResourceNames
190) -> AccessGroupResponse:
191 payload: Final = MappingProxyType(
192 {
193 **record.dict(),
194 "assigned_team_ids": assigned_team_ids,
195 "access_mcp_servers": _label(record.access_mcp_server_ids or (), names.mcp_servers),
196 "access_agents": _label(record.access_agent_ids or (), names.agents),
197 "assigned_teams": _label(assigned_team_ids, names.teams),
198 "assigned_keys": _label(record.assigned_key_ids or (), names.keys),
199 }
200 )
201 return AccessGroupResponse.model_validate(payload)
204def _ids_across(
205 records: Sequence[_AccessGroupRecord], pick: Callable[[_AccessGroupRecord], Sequence[str] | None]
206) -> tuple[str, ...]:
207 return tuple(dict.fromkeys(resource_id for record in records for resource_id in (pick(record) or ())))
210async def _responses_for(
211 prisma_client: PrismaClient, records: Sequence[_AccessGroupRecord]
212) -> tuple[AccessGroupResponse, ...]:
213 if not records:
214 return ()
215 teams: Final = await _teams_touching(TeamRepository(prisma_client).table, records)
216 mcp_servers, agents, keys = await asyncio.gather(
217 mcp_server_display_names(
218 prisma_client,
219 _ids_across(records, lambda record: record.access_mcp_server_ids),
220 global_mcp_server_manager.config_mcp_servers,
221 ),
222 agent_display_names(
223 prisma_client, _ids_across(records, lambda record: record.access_agent_ids), global_agent_registry
224 ),
225 key_display_names(prisma_client, _ids_across(records, lambda record: record.assigned_key_ids)),
226 )
227 names: Final = _ResourceNames(
228 mcp_servers=mcp_servers,
229 agents=agents,
230 teams=MappingProxyType({team.team_id: team.team_alias for team in teams}),
231 keys=keys,
232 )
233 attached: Final = _attached_team_ids_by_group(records, teams)
234 return tuple(
235 _record_to_response(record, assigned_team_ids=attached[record.access_group_id], names=names)
236 for record in records
237 )
240async def _response_for(prisma_client: PrismaClient, record: _AccessGroupRecord) -> AccessGroupResponse:
241 (response,) = await _responses_for(prisma_client, (record,))
242 return response
245def _attached_team_ids_by_group(
246 records: Sequence[_AccessGroupRecord], teams: Sequence[_TeamRecord]
247) -> Mapping[str, tuple[str, ...]]:
248 """Teams really attached to each group: the stored column minus ghosts, plus teams the mirror missed."""
249 real_team_ids: Final = frozenset(team.team_id for team in teams)
251 def attached(record: _AccessGroupRecord) -> tuple[str, ...]:
252 stored: Final = (team_id for team_id in (record.assigned_team_ids or ()) if team_id in real_team_ids)
253 carrying: Final = (team.team_id for team in teams if record.access_group_id in (team.access_group_ids or ()))
254 return tuple(dict.fromkeys((*stored, *carrying)))
256 return MappingProxyType({record.access_group_id: attached(record) for record in records})
259async def _teams_touching(team_table: _TeamTable, records: Sequence[_AccessGroupRecord]) -> Sequence[_TeamRecord]:
260 """Team rows listed on any of the groups or carrying any of them in access_group_ids."""
261 group_ids: Final = tuple(record.access_group_id for record in records)
262 stored_team_ids: Final = _ids_across(records, lambda record: record.assigned_team_ids)
263 carrying: Final = {"access_group_ids": {"hasSome": group_ids}} # mutable-ok: prisma where is a dict
264 listed: Final = {"team_id": {"in": stored_team_ids}} # mutable-ok: prisma where is a dict
265 return await team_table.find_many(where={"OR": (carrying, listed)}) # mutable-ok: prisma where is a dict
268async def _attached_team_ids_for(
269 team_table: _TeamTable, records: Sequence[_AccessGroupRecord]
270) -> Mapping[str, tuple[str, ...]]:
271 if not records: 271 ↛ 272line 271 didn't jump to line 272 because the condition on line 271 was never true
272 return MappingProxyType({})
273 return _attached_team_ids_by_group(records, await _teams_touching(team_table, records))
276async def _require_teams_exist(tx: _AccessGroupTx, team_ids: Sequence[str]) -> None:
277 if not team_ids:
278 return
279 where: Final = {"team_id": {"in": team_ids}} # mutable-ok: prisma where is a dict
280 found: Final = await tx.litellm_teamtable.find_many(where=where)
281 missing: Final = frozenset(team_ids) - frozenset(team.team_id for team in found)
282 if missing: 282 ↛ exitline 282 didn't return from function '_require_teams_exist' because the condition on line 282 was always true
283 raise HTTPException(
284 status_code=status.HTTP_400_BAD_REQUEST,
285 detail=f"Unknown team ids: {', '.join(sorted(missing))}",
286 )
289def _record_to_access_group_table(record: _AccessGroupRecord) -> LiteLLM_AccessGroupTable:
290 """Convert a Prisma record to a LiteLLM_AccessGroupTable pydantic object for caching."""
291 return LiteLLM_AccessGroupTable.model_validate(record.dict())
294async def _cache_access_group_record(record: _AccessGroupRecord) -> None:
295 """
296 Cache an access group Prisma record in the user_api_key_cache.
298 Uses a lazy import of user_api_key_cache and proxy_logging_obj from proxy_server
299 to avoid circular imports, following the same pattern as key_management_endpoints.
300 """
301 from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
303 access_group_table: Final = _record_to_access_group_table(record)
304 await _cache_access_object(
305 access_group_id=record.access_group_id,
306 access_group_table=access_group_table,
307 user_api_key_cache=user_api_key_cache,
308 proxy_logging_obj=proxy_logging_obj,
309 )
312# ---------------------------------------------------------------------------
313# DB sync helpers (called inside a Prisma transaction)
314# ---------------------------------------------------------------------------
317async def _sync_add_access_group_to_teams(tx: _AccessGroupTx, team_ids: list[str], access_group_id: str) -> None:
318 """Add access_group_id to each team's access_group_ids (idempotent)."""
319 for team_id in team_ids: 319 ↛ 320line 319 didn't jump to line 320 because the loop on line 319 never started
320 team = await tx.litellm_teamtable.find_unique(where={"team_id": team_id})
321 if team is not None and access_group_id not in (team.access_group_ids or []):
322 await tx.litellm_teamtable.update(
323 where={"team_id": team_id},
324 data={"access_group_ids": list(team.access_group_ids or []) + [access_group_id]},
325 )
328async def _sync_remove_access_group_from_teams(tx: _AccessGroupTx, team_ids: list[str], access_group_id: str) -> None:
329 """Remove access_group_id from each team's access_group_ids (idempotent)."""
330 for team_id in team_ids: 330 ↛ 331line 330 didn't jump to line 331 because the loop on line 330 never started
331 team = await tx.litellm_teamtable.find_unique(where={"team_id": team_id})
332 if team is not None and access_group_id in (team.access_group_ids or []):
333 await tx.litellm_teamtable.update(
334 where={"team_id": team_id},
335 data={"access_group_ids": [ag for ag in (team.access_group_ids or ()) if ag != access_group_id]},
336 )
339async def _sync_add_access_group_to_keys(tx: _AccessGroupTx, key_tokens: list[str], access_group_id: str) -> None:
340 """Add access_group_id to each key's access_group_ids (idempotent)."""
341 for token in key_tokens:
342 key = await tx.litellm_verificationtoken.find_unique(where={"token": token})
343 if key is not None and access_group_id not in (key.access_group_ids or []): 343 ↛ 344line 343 didn't jump to line 344 because the condition on line 343 was never true
344 await tx.litellm_verificationtoken.update(
345 where={"token": token},
346 data={"access_group_ids": list(key.access_group_ids or []) + [access_group_id]},
347 )
350async def _sync_remove_access_group_from_keys(tx: _AccessGroupTx, key_tokens: list[str], access_group_id: str) -> None:
351 """Remove access_group_id from each key's access_group_ids (idempotent)."""
352 for token in key_tokens:
353 key = await tx.litellm_verificationtoken.find_unique(where={"token": token})
354 if key is not None and access_group_id in (key.access_group_ids or []): 354 ↛ 355line 354 didn't jump to line 355 because the condition on line 354 was never true
355 await tx.litellm_verificationtoken.update(
356 where={"token": token},
357 data={"access_group_ids": [ag for ag in (key.access_group_ids or ()) if ag != access_group_id]},
358 )
361def _without_access_group(access_group_ids: Sequence[str] | None, access_group_id: str) -> tuple[str, ...]:
362 return tuple(ag for ag in (access_group_ids or ()) if ag != access_group_id)
365async def _detach_access_group_from_agents(tx: _AccessGroupTx, access_group_id: str) -> tuple[str, ...]:
366 agents_with_group: Final = await tx.litellm_agentstable.find_many(
367 where=_AgentAccessGroupsWhere(access_group_ids=_HasSomeFilter(hasSome=(access_group_id,)))
368 )
369 for agent in agents_with_group: 369 ↛ 370line 369 didn't jump to line 370 because the loop on line 369 never started
370 await tx.litellm_agentstable.update(
371 where=_AgentIdWhere(agent_id=agent.agent_id),
372 data=_AgentAccessGroupsData(
373 access_group_ids=_without_access_group(agent.access_group_ids, access_group_id)
374 ),
375 )
376 return tuple(agent.agent_id for agent in agents_with_group)
379def _detach_access_group_from_agent_registry(agent_ids: Sequence[str], access_group_id: str) -> None:
380 registered: Final = tuple(
381 agent
382 for agent in (global_agent_registry.get_agent_by_id(agent_id) for agent_id in agent_ids)
383 if agent is not None
384 )
385 for agent in registered: 385 ↛ 386line 385 didn't jump to line 386 because the loop on line 385 never started
386 global_agent_registry.deregister_agent(agent_name=agent.agent_name)
387 global_agent_registry.register_agent(
388 agent_config=agent.model_copy(
389 update=_AgentAccessGroupsData(
390 access_group_ids=_without_access_group(agent.access_group_ids, access_group_id)
391 )
392 )
393 )
396# ---------------------------------------------------------------------------
397# Cache patch helpers
398# ---------------------------------------------------------------------------
401async def _patch_team_caches_add_access_group(
402 team_ids: list[str],
403 access_group_id: str,
404 user_api_key_cache,
405 proxy_logging_obj,
406) -> None:
407 """Patch cached team objects to include access_group_id."""
408 for team_id in team_ids: 408 ↛ 409line 408 didn't jump to line 409 because the loop on line 408 never started
409 cached_team = await _get_team_object_from_cache(
410 key=f"team_id:{team_id}",
411 user_api_key_cache=user_api_key_cache,
412 parent_otel_span=None,
413 )
414 if cached_team is None:
415 continue
416 if cached_team.access_group_ids is None:
417 cached_team.access_group_ids = [access_group_id]
418 elif access_group_id not in cached_team.access_group_ids:
419 cached_team.access_group_ids = list(cached_team.access_group_ids) + [access_group_id]
420 else:
421 continue
422 await _cache_team_object(
423 team_id=team_id,
424 team_table=cached_team,
425 user_api_key_cache=user_api_key_cache,
426 proxy_logging_obj=proxy_logging_obj,
427 )
430async def _patch_team_caches_remove_access_group(
431 team_ids: list[str],
432 access_group_id: str,
433 user_api_key_cache,
434 proxy_logging_obj,
435) -> None:
436 """Patch cached team objects to remove access_group_id."""
437 for team_id in team_ids: 437 ↛ 438line 437 didn't jump to line 438 because the loop on line 437 never started
438 cached_team = await _get_team_object_from_cache(
439 key=f"team_id:{team_id}",
440 user_api_key_cache=user_api_key_cache,
441 parent_otel_span=None,
442 )
443 if cached_team is not None and cached_team.access_group_ids:
444 cached_team.access_group_ids = [ag for ag in cached_team.access_group_ids if ag != access_group_id]
445 await _cache_team_object(
446 team_id=team_id,
447 team_table=cached_team,
448 user_api_key_cache=user_api_key_cache,
449 proxy_logging_obj=proxy_logging_obj,
450 )
453async def _patch_key_caches_add_access_group(
454 key_tokens: list[str],
455 access_group_id: str,
456 user_api_key_cache,
457 proxy_logging_obj,
458) -> None:
459 """Patch cached key objects to include access_group_id."""
460 for token in key_tokens:
461 cached_key = await user_api_key_cache.async_get_cache(
462 key=token,
463 model_type=UserAPIKeyAuth,
464 )
465 if cached_key is None: 465 ↛ 467line 465 didn't jump to line 467 because the condition on line 465 was always true
466 continue
467 if cached_key.access_group_ids is None:
468 cached_key.access_group_ids = [access_group_id]
469 elif access_group_id not in cached_key.access_group_ids:
470 cached_key.access_group_ids = list(cached_key.access_group_ids) + [access_group_id]
471 else:
472 continue
473 await _cache_key_object(
474 hashed_token=token,
475 user_api_key_obj=cached_key,
476 user_api_key_cache=user_api_key_cache,
477 proxy_logging_obj=proxy_logging_obj,
478 )
481async def _patch_key_caches_remove_access_group(
482 key_tokens: list[str],
483 access_group_id: str,
484 user_api_key_cache,
485 proxy_logging_obj,
486) -> None:
487 """Patch cached key objects to remove access_group_id."""
488 for token in key_tokens:
489 cached_key = await user_api_key_cache.async_get_cache(
490 key=token,
491 model_type=UserAPIKeyAuth,
492 )
493 if cached_key is not None and cached_key.access_group_ids: 493 ↛ 494line 493 didn't jump to line 494 because the condition on line 493 was never true
494 cached_key.access_group_ids = [ag for ag in cached_key.access_group_ids if ag != access_group_id]
495 await _cache_key_object(
496 hashed_token=token,
497 user_api_key_obj=cached_key,
498 user_api_key_cache=user_api_key_cache,
499 proxy_logging_obj=proxy_logging_obj,
500 )
503# ---------------------------------------------------------------------------
504# CRUD endpoints
505# ---------------------------------------------------------------------------
508@router.post(
509 "/v1/access_group",
510 response_model=AccessGroupResponse,
511 status_code=status.HTTP_201_CREATED,
512)
513async def create_access_group(
514 data: AccessGroupCreateRequest,
515 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
516) -> AccessGroupResponse:
517 _require_proxy_admin(user_api_key_dict)
518 prisma_client: Final = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
520 try:
521 tx: _AccessGroupTx
522 async with prisma_client.db.tx() as tx:
523 existing: Final = await tx.litellm_accessgrouptable.find_unique(
524 where={"access_group_name": data.access_group_name}
525 )
526 if existing is not None:
527 raise HTTPException(
528 status_code=status.HTTP_409_CONFLICT,
529 detail=f"Access group '{data.access_group_name}' already exists",
530 )
531 await _require_teams_exist(tx, data.assigned_team_ids or ())
533 record: Final = await tx.litellm_accessgrouptable.create(
534 data={
535 "access_group_name": data.access_group_name,
536 "description": data.description,
537 "access_model_names": data.access_model_names or [],
538 "access_mcp_server_ids": data.access_mcp_server_ids or [],
539 "access_agent_ids": data.access_agent_ids or [],
540 "assigned_team_ids": data.assigned_team_ids or [],
541 "assigned_key_ids": data.assigned_key_ids or [],
542 "created_by": user_api_key_dict.user_id,
543 "updated_by": user_api_key_dict.user_id,
544 }
545 )
547 # Sync team and key tables to reference the new access group
548 await _sync_add_access_group_to_teams(tx, data.assigned_team_ids or [], record.access_group_id)
549 await _sync_add_access_group_to_keys(tx, data.assigned_key_ids or [], record.access_group_id)
550 except HTTPException:
551 raise
552 except Exception as e:
553 # Race condition: another request created the same name between find_unique and create.
554 if "unique constraint" in str(e).lower() or "P2002" in str(e):
555 raise HTTPException(
556 status_code=status.HTTP_409_CONFLICT,
557 detail=f"Access group '{data.access_group_name}' already exists",
558 )
559 raise
561 from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
563 await _cache_access_group_record(record)
564 await _patch_team_caches_add_access_group(
565 data.assigned_team_ids or [],
566 record.access_group_id,
567 user_api_key_cache,
568 proxy_logging_obj,
569 )
570 await _patch_key_caches_add_access_group(
571 data.assigned_key_ids or [],
572 record.access_group_id,
573 user_api_key_cache,
574 proxy_logging_obj,
575 )
577 return await _response_for(prisma_client, record)
580@router.get(
581 "/v1/access_group",
582 response_model=list[AccessGroupResponse],
583)
584async def list_access_groups(
585 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
586) -> Sequence[AccessGroupResponse]:
587 _require_admin_view(user_api_key_dict)
588 prisma_client: Final = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
590 table: Final = AccessGroupRepository(prisma_client).table
591 records: Final = await table.find_many(order={"created_at": "desc"})
592 return await _responses_for(prisma_client, records)
595@router.get(
596 "/v1/access_group/{access_group_id}",
597 response_model=AccessGroupResponse,
598)
599async def get_access_group(
600 access_group_id: str,
601 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
602) -> AccessGroupResponse:
603 _require_admin_view(user_api_key_dict)
604 prisma_client: Final = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
606 table: Final = AccessGroupRepository(prisma_client).table
607 record: Final = await table.find_unique(where={"access_group_id": access_group_id})
608 if record is None:
609 raise HTTPException(
610 status_code=status.HTTP_404_NOT_FOUND,
611 detail=f"Access group '{access_group_id}' not found",
612 )
613 return await _response_for(prisma_client, record)
616@router.put(
617 "/v1/access_group/{access_group_id}",
618 response_model=AccessGroupResponse,
619)
620async def update_access_group(
621 access_group_id: str,
622 data: AccessGroupUpdateRequest,
623 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
624) -> AccessGroupResponse:
625 _require_proxy_admin(user_api_key_dict)
626 prisma_client: Final = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
628 update_fields: Final = data.model_dump(exclude_unset=True)
629 update_data: Final[dict] = {"updated_by": user_api_key_dict.user_id}
630 for field, value in update_fields.items():
631 if (
632 field
633 in (
634 "assigned_team_ids",
635 "assigned_key_ids",
636 "access_model_names",
637 "access_mcp_server_ids",
638 "access_agent_ids",
639 )
640 and value is None
641 ):
642 value = []
643 update_data[field] = value
645 # Initialize delta lists before the try block so they remain accessible
646 # for cache updates after the transaction, even if an error path is added later.
647 teams_to_add: list[str] = []
648 teams_to_remove: list[str] = []
649 keys_to_add: list[str] = []
650 keys_to_remove: list[str] = []
652 try:
653 tx: _AccessGroupTx
654 async with prisma_client.db.tx() as tx:
655 # Read inside the transaction so delta computation is consistent with the write,
656 # avoiding a TOCTOU race where a concurrent update could make deltas stale.
657 existing: Final = await tx.litellm_accessgrouptable.find_unique(where={"access_group_id": access_group_id})
658 if existing is None:
659 raise HTTPException(
660 status_code=status.HTTP_404_NOT_FOUND,
661 detail=f"Access group '{access_group_id}' not found",
662 )
663 await _require_teams_exist(tx, data.assigned_team_ids or ())
665 attached: Final = await _attached_team_ids_for(tx.litellm_teamtable, (existing,))
666 old_team_ids: Final[set[str]] = set(attached[access_group_id])
667 old_key_ids: Final[set[str]] = set(existing.assigned_key_ids or [])
668 new_team_ids: Final[set[str]] = (
669 set(update_fields["assigned_team_ids"] or []) if "assigned_team_ids" in update_fields else old_team_ids
670 )
671 new_key_ids: Final[set[str]] = (
672 set(update_fields["assigned_key_ids"] or []) if "assigned_key_ids" in update_fields else old_key_ids
673 )
675 teams_to_add = list(new_team_ids - old_team_ids)
676 teams_to_remove = list(old_team_ids - new_team_ids)
677 keys_to_add = list(new_key_ids - old_key_ids)
678 keys_to_remove = list(old_key_ids - new_key_ids)
680 record: Final = await tx.litellm_accessgrouptable.update(
681 where={"access_group_id": access_group_id},
682 data=update_data,
683 )
685 await _sync_add_access_group_to_teams(tx, teams_to_add, access_group_id)
686 await _sync_remove_access_group_from_teams(tx, teams_to_remove, access_group_id)
687 await _sync_add_access_group_to_keys(tx, keys_to_add, access_group_id)
688 await _sync_remove_access_group_from_keys(tx, keys_to_remove, access_group_id)
689 except HTTPException:
690 raise
691 except Exception as e:
692 # Unique constraint violation (e.g. access_group_name already exists).
693 if "unique constraint" in str(e).lower() or "P2002" in str(e):
694 raise HTTPException(
695 status_code=status.HTTP_409_CONFLICT,
696 detail=f"Access group '{update_data.get('access_group_name', '')}' already exists",
697 )
698 raise
700 from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
702 await _cache_access_group_record(record)
703 await _patch_team_caches_add_access_group(teams_to_add, access_group_id, user_api_key_cache, proxy_logging_obj)
704 await _patch_team_caches_remove_access_group(
705 teams_to_remove, access_group_id, user_api_key_cache, proxy_logging_obj
706 )
707 await _patch_key_caches_add_access_group(keys_to_add, access_group_id, user_api_key_cache, proxy_logging_obj)
708 await _patch_key_caches_remove_access_group(keys_to_remove, access_group_id, user_api_key_cache, proxy_logging_obj)
710 return await _response_for(prisma_client, record)
713@router.delete(
714 "/v1/access_group/{access_group_id}",
715 status_code=status.HTTP_204_NO_CONTENT,
716)
717async def delete_access_group(
718 access_group_id: str,
719 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
720) -> None:
721 _require_proxy_admin(user_api_key_dict)
722 prisma_client: Final = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value)
724 try:
725 affected_team_ids: list[str] = []
726 affected_key_tokens: list[str] = []
728 tx: _AccessGroupTx
729 async with prisma_client.db.tx() as tx:
730 existing: Final = await tx.litellm_accessgrouptable.find_unique(where={"access_group_id": access_group_id})
731 if existing is None:
732 raise HTTPException(
733 status_code=status.HTTP_404_NOT_FOUND,
734 detail=f"Access group '{access_group_id}' not found",
735 )
737 # Union of: teams that have this access_group_id in their own access_group_ids
738 # AND teams listed in assigned_team_ids (handles out-of-sync data from before this sync was added)
739 teams_with_group: Final = await tx.litellm_teamtable.find_many(
740 where={"access_group_ids": {"hasSome": [access_group_id]}}
741 )
742 all_affected_team_ids: Final[set[str]] = {team.team_id for team in teams_with_group} | set(
743 existing.assigned_team_ids or []
744 )
745 affected_team_ids = list(all_affected_team_ids)
747 # Union of: keys that have this access_group_id in their own access_group_ids
748 # AND keys listed in assigned_key_ids (handles out-of-sync data)
749 keys_with_group: Final = await tx.litellm_verificationtoken.find_many(
750 where={"access_group_ids": {"hasSome": [access_group_id]}}
751 )
752 all_affected_key_tokens: Final[set[str]] = {key.token for key in keys_with_group} | set(
753 existing.assigned_key_ids or []
754 )
755 affected_key_tokens = list(all_affected_key_tokens)
757 # Update teams returned by find_many directly — we already have their data.
758 for team in teams_with_group: 758 ↛ 759line 758 didn't jump to line 759 because the loop on line 758 never started
759 await tx.litellm_teamtable.update(
760 where={"team_id": team.team_id},
761 data={"access_group_ids": [ag for ag in (team.access_group_ids or ()) if ag != access_group_id]},
762 )
763 # Use _sync_remove only for out-of-sync teams not found by the hasSome query.
764 out_of_sync_team_ids: Final = set(existing.assigned_team_ids or []) - {t.team_id for t in teams_with_group}
765 await _sync_remove_access_group_from_teams(tx, list(out_of_sync_team_ids), access_group_id)
767 # Update keys returned by find_many directly — we already have their data.
768 for key in keys_with_group: 768 ↛ 769line 768 didn't jump to line 769 because the loop on line 768 never started
769 await tx.litellm_verificationtoken.update(
770 where={"token": key.token},
771 data={"access_group_ids": [ag for ag in (key.access_group_ids or ()) if ag != access_group_id]},
772 )
773 # Use _sync_remove only for out-of-sync keys not found by the hasSome query.
774 out_of_sync_key_tokens: Final = set(existing.assigned_key_ids or []) - {k.token for k in keys_with_group}
775 await _sync_remove_access_group_from_keys(tx, list(out_of_sync_key_tokens), access_group_id)
777 detached_agent_ids: Final = await _detach_access_group_from_agents(tx, access_group_id)
779 await tx.litellm_accessgrouptable.delete(where={"access_group_id": access_group_id})
781 from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache
783 await invalidate_access_group_cache(access_group_id)
784 _detach_access_group_from_agent_registry(detached_agent_ids, access_group_id)
785 await _patch_team_caches_remove_access_group(
786 affected_team_ids, access_group_id, user_api_key_cache, proxy_logging_obj
787 )
788 await _patch_key_caches_remove_access_group(
789 affected_key_tokens, access_group_id, user_api_key_cache, proxy_logging_obj
790 )
792 except HTTPException:
793 raise
794 except Exception as e:
795 verbose_proxy_logger.exception(
796 "delete_access_group failed: access_group_id=%s error=%s",
797 access_group_id,
798 e,
799 )
800 if PrismaDBExceptionHandler.is_database_infrastructure_error(e):
801 raise HTTPException(
802 status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
803 detail=CommonProxyErrors.db_not_connected_error.value,
804 )
805 if "P2025" in str(e) or ("record" in str(e).lower() and "not found" in str(e).lower()):
806 raise HTTPException(
807 status_code=status.HTTP_404_NOT_FOUND,
808 detail=f"Access group '{access_group_id}' not found",
809 )
810 raise HTTPException(
811 status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
812 detail="Failed to delete access group. Please try again.",
813 )
816# Alias routes for /v1/unified_access_group
817router.add_api_route(
818 "/v1/unified_access_group",
819 create_access_group,
820 methods=["POST"],
821 response_model=AccessGroupResponse,
822 status_code=status.HTTP_201_CREATED,
823)
824router.add_api_route(
825 "/v1/unified_access_group",
826 list_access_groups,
827 methods=["GET"],
828 response_model=list[AccessGroupResponse],
829)
830router.add_api_route(
831 "/v1/unified_access_group/{access_group_id}",
832 get_access_group,
833 methods=["GET"],
834 response_model=AccessGroupResponse,
835)
836router.add_api_route(
837 "/v1/unified_access_group/{access_group_id}",
838 update_access_group,
839 methods=["PUT"],
840 response_model=AccessGroupResponse,
841)
842router.add_api_route(
843 "/v1/unified_access_group/{access_group_id}",
844 delete_access_group,
845 methods=["DELETE"],
846 status_code=status.HTTP_204_NO_CONTENT,
847)