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

1import asyncio 

2from collections.abc import Callable, Mapping, Sequence 

3from dataclasses import dataclass 

4from types import MappingProxyType 

5from typing import Final, Protocol 

6 

7from fastapi import APIRouter, Depends, HTTPException, status 

8from typing_extensions import ReadOnly, TypedDict 

9 

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) 

41 

42router: Final = APIRouter( 

43 tags=["access group management"], 

44) 

45 

46 

47class _AccessGroupRecord(Protocol): 

48 @property 

49 def access_group_id(self) -> str: ... 49 ↛ exitline 49 didn't return from function 'access_group_id' because

50 

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

53 

54 @property 

55 def access_agent_ids(self) -> Sequence[str] | None: ... 55 ↛ exitline 55 didn't return from function 'access_agent_ids' because

56 

57 @property 

58 def assigned_team_ids(self) -> Sequence[str] | None: ... 58 ↛ exitline 58 didn't return from function 'assigned_team_ids' because

59 

60 @property 

61 def assigned_key_ids(self) -> Sequence[str] | None: ... 61 ↛ exitline 61 didn't return from function 'assigned_key_ids' because

62 

63 def dict(self) -> Mapping[str, object]: ... 63 ↛ exitline 63 didn't return from function 'dict' because

64 

65 

66class _TeamRecord(Protocol): 

67 @property 

68 def team_id(self) -> str: ... 68 ↛ exitline 68 didn't return from function 'team_id' because

69 

70 @property 

71 def team_alias(self) -> str | None: ... 71 ↛ exitline 71 didn't return from function 'team_alias' because

72 

73 @property 

74 def access_group_ids(self) -> Sequence[str] | None: ... 74 ↛ exitline 74 didn't return from function 'access_group_ids' because

75 

76 

77class _KeyRecord(Protocol): 

78 @property 

79 def token(self) -> str: ... 79 ↛ exitline 79 didn't return from function 'token' because

80 

81 @property 

82 def access_group_ids(self) -> Sequence[str] | None: ... 82 ↛ exitline 82 didn't return from function 'access_group_ids' because

83 

84 

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

87 

88 async def find_many(self, order: Mapping[str, object]) -> Sequence[_AccessGroupRecord]: ... 88 ↛ exitline 88 didn't return from function 'find_many' because

89 

90 async def create(self, data: Mapping[str, object]) -> _AccessGroupRecord: ... 90 ↛ exitline 90 didn't return from function 'create' because

91 

92 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _AccessGroupRecord: ... 92 ↛ exitline 92 didn't return from function 'update' because

93 

94 async def delete(self, where: Mapping[str, object]) -> object: ... 94 ↛ exitline 94 didn't return from function 'delete' because

95 

96 

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

99 

100 async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_TeamRecord]: ... 100 ↛ exitline 100 didn't return from function 'find_many' because

101 

102 async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> object: ... 102 ↛ exitline 102 didn't return from function 'update' because

103 

104 

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

107 

108 async def find_many(self, where: Mapping[str, object]) -> Sequence[_KeyRecord]: ... 108 ↛ exitline 108 didn't return from function 'find_many' because

109 

110 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> object: ... 110 ↛ exitline 110 didn't return from function 'update' because

111 

112 

113class _AgentRecord(Protocol): 

114 @property 

115 def agent_id(self) -> str: ... 115 ↛ exitline 115 didn't return from function 'agent_id' because

116 

117 @property 

118 def access_group_ids(self) -> Sequence[str] | None: ... 118 ↛ exitline 118 didn't return from function 'access_group_ids' because

119 

120 

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

123 

124 async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> object: ... 124 ↛ exitline 124 didn't return from function 'update' because

125 

126 

127class _HasSomeFilter(TypedDict): 

128 hasSome: ReadOnly[Sequence[str]] 

129 

130 

131class _AgentAccessGroupsWhere(TypedDict): 

132 access_group_ids: ReadOnly[_HasSomeFilter] 

133 

134 

135class _AgentIdWhere(TypedDict): 

136 agent_id: ReadOnly[str] 

137 

138 

139class _AgentAccessGroupsData(TypedDict): 

140 access_group_ids: ReadOnly[Sequence[str]] 

141 

142 

143class _AccessGroupTx(Protocol): 

144 @property 

145 def litellm_accessgrouptable(self) -> _AccessGroupTable: ... 145 ↛ exitline 145 didn't return from function 'litellm_accessgrouptable' because

146 

147 @property 

148 def litellm_teamtable(self) -> _TeamTable: ... 148 ↛ exitline 148 didn't return from function 'litellm_teamtable' because

149 

150 @property 

151 def litellm_verificationtoken(self) -> _KeyTable: ... 151 ↛ exitline 151 didn't return from function 'litellm_verificationtoken' because

152 

153 @property 

154 def litellm_agentstable(self) -> _AgentTable: ... 154 ↛ exitline 154 didn't return from function 'litellm_agentstable' because

155 

156 

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 ) 

163 

164 

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 

168 

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 ) 

174 

175 

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] 

182 

183 

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) 

186 

187 

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) 

202 

203 

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 ()))) 

208 

209 

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 ) 

238 

239 

240async def _response_for(prisma_client: PrismaClient, record: _AccessGroupRecord) -> AccessGroupResponse: 

241 (response,) = await _responses_for(prisma_client, (record,)) 

242 return response 

243 

244 

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) 

250 

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))) 

255 

256 return MappingProxyType({record.access_group_id: attached(record) for record in records}) 

257 

258 

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 

266 

267 

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)) 

274 

275 

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 ) 

287 

288 

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()) 

292 

293 

294async def _cache_access_group_record(record: _AccessGroupRecord) -> None: 

295 """ 

296 Cache an access group Prisma record in the user_api_key_cache. 

297 

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 

302 

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 ) 

310 

311 

312# --------------------------------------------------------------------------- 

313# DB sync helpers (called inside a Prisma transaction) 

314# --------------------------------------------------------------------------- 

315 

316 

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 ) 

326 

327 

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 ) 

337 

338 

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 ) 

348 

349 

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 ) 

359 

360 

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) 

363 

364 

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) 

377 

378 

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 ) 

394 

395 

396# --------------------------------------------------------------------------- 

397# Cache patch helpers 

398# --------------------------------------------------------------------------- 

399 

400 

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 ) 

428 

429 

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 ) 

451 

452 

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 ) 

479 

480 

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 ) 

501 

502 

503# --------------------------------------------------------------------------- 

504# CRUD endpoints 

505# --------------------------------------------------------------------------- 

506 

507 

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) 

519 

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 ()) 

532 

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 ) 

546 

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 

560 

561 from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache 

562 

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 ) 

576 

577 return await _response_for(prisma_client, record) 

578 

579 

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) 

589 

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) 

593 

594 

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) 

605 

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) 

614 

615 

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) 

627 

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 

644 

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] = [] 

651 

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 ()) 

664 

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 ) 

674 

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) 

679 

680 record: Final = await tx.litellm_accessgrouptable.update( 

681 where={"access_group_id": access_group_id}, 

682 data=update_data, 

683 ) 

684 

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 

699 

700 from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache 

701 

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) 

709 

710 return await _response_for(prisma_client, record) 

711 

712 

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) 

723 

724 try: 

725 affected_team_ids: list[str] = [] 

726 affected_key_tokens: list[str] = [] 

727 

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 ) 

736 

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) 

746 

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) 

756 

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) 

766 

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) 

776 

777 detached_agent_ids: Final = await _detach_access_group_from_agents(tx, access_group_id) 

778 

779 await tx.litellm_accessgrouptable.delete(where={"access_group_id": access_group_id}) 

780 

781 from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache 

782 

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 ) 

791 

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 ) 

814 

815 

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)