groups.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449
  1. """Group management API routes."""
  2. from fastapi import APIRouter, Depends, HTTPException, status
  3. from sqlalchemy import delete, insert, select
  4. from sqlalchemy.ext.asyncio import AsyncSession
  5. from sqlalchemy.orm import selectinload
  6. from backend.app.core.auth import RequireAdminIfAuthEnabled, RequirePermissionIfAuthEnabled
  7. from backend.app.core.database import get_db
  8. from backend.app.core.permissions import (
  9. ALL_PERMISSIONS,
  10. PERMISSION_CATEGORIES,
  11. Permission,
  12. )
  13. from backend.app.core.printer_scope import group_location_names, group_printer_ids
  14. from backend.app.core.websocket import ws_manager
  15. from backend.app.models.group import Group, group_locations, group_printers
  16. from backend.app.models.printer import Printer
  17. from backend.app.models.user import User
  18. from backend.app.schemas.group import (
  19. GroupCreate,
  20. GroupDetailResponse,
  21. GroupResponse,
  22. GroupUpdate,
  23. PermissionCategory,
  24. PermissionInfo,
  25. PermissionsListResponse,
  26. UserBrief,
  27. )
  28. router = APIRouter(prefix="/groups", tags=["groups"])
  29. # Permissions whose derived label would misdescribe what is being granted.
  30. # The derived form for USERS_READ_SLIM is "Read Slim Users", which reads as a
  31. # property of the users rather than of the response -- and an admin ticking a
  32. # box in the group editor has nothing else to go on (#1894).
  33. _PERMISSION_LABEL_OVERRIDES: dict[Permission, str] = {
  34. Permission.USERS_READ_SLIM: "List User Names (id + username only)",
  35. }
  36. def _permission_label(perm: Permission) -> str:
  37. """Convert permission enum to human-readable label."""
  38. if perm in _PERMISSION_LABEL_OVERRIDES:
  39. return _PERMISSION_LABEL_OVERRIDES[perm]
  40. # e.g., "printers:read" -> "Read Printers"
  41. parts = perm.value.split(":")
  42. if len(parts) == 2:
  43. resource, action = parts
  44. resource = resource.replace("_", " ").title()
  45. action = action.replace("_", " ").title()
  46. return f"{action} {resource}"
  47. return perm.value
  48. async def _printer_ids_by_group(db: AsyncSession) -> dict[int, list[int]]:
  49. result = await db.execute(select(group_printers.c.group_id, group_printers.c.printer_id))
  50. by_group: dict[int, list[int]] = {}
  51. for group_id, printer_id in result.all():
  52. by_group.setdefault(group_id, []).append(printer_id)
  53. return {gid: sorted(pids) for gid, pids in by_group.items()}
  54. async def _locations_by_group(db: AsyncSession) -> dict[int, list[str]]:
  55. result = await db.execute(select(group_locations.c.group_id, group_locations.c.location))
  56. by_group: dict[int, list[str]] = {}
  57. for group_id, location in result.all():
  58. by_group.setdefault(group_id, []).append(location)
  59. return {gid: sorted(locs) for gid, locs in by_group.items()}
  60. def _clean_locations(locations: list[str]) -> set[str]:
  61. """Trimmed, non-empty location names; refuses ones too long to ever match a printer."""
  62. cleaned = {loc.strip() for loc in locations if loc and loc.strip()}
  63. if any(len(loc) > 100 for loc in cleaned):
  64. raise HTTPException(
  65. status_code=status.HTTP_400_BAD_REQUEST,
  66. detail="Location names are at most 100 characters",
  67. )
  68. return cleaned
  69. async def _apply_printer_scope(
  70. db: AsyncSession,
  71. group: Group,
  72. restrict_printers: bool | None,
  73. printer_ids: list[int] | None,
  74. locations: list[str] | None = None,
  75. ) -> bool:
  76. """Validate and store a group's printer scope (#1727). Returns whether it changed.
  77. The Administrators group can't be restricted: admins see every printer
  78. regardless, so the setting would only mislead.
  79. """
  80. changed = False
  81. if restrict_printers is not None and restrict_printers != bool(group.restrict_printers):
  82. if restrict_printers and group.name == "Administrators":
  83. raise HTTPException(
  84. status_code=status.HTTP_400_BAD_REQUEST,
  85. detail="Administrators always see every printer",
  86. )
  87. group.restrict_printers = restrict_printers
  88. changed = True
  89. if printer_ids is not None:
  90. wanted = set(printer_ids)
  91. if wanted:
  92. found = set((await db.execute(select(Printer.id).where(Printer.id.in_(wanted)))).scalars().all())
  93. missing = sorted(wanted - found)
  94. if missing:
  95. raise HTTPException(
  96. status_code=status.HTTP_400_BAD_REQUEST,
  97. detail=f"Invalid printers: {', '.join(str(pid) for pid in missing)}",
  98. )
  99. current = set(await group_printer_ids(db, group.id)) if group.id is not None else set()
  100. if wanted != current:
  101. if group.id is None:
  102. await db.flush()
  103. await db.execute(delete(group_printers).where(group_printers.c.group_id == group.id))
  104. if wanted:
  105. await db.execute(
  106. insert(group_printers), [{"group_id": group.id, "printer_id": pid} for pid in sorted(wanted)]
  107. )
  108. changed = True
  109. if locations is not None:
  110. # Not checked against existing printers: a location can be granted
  111. # before its first printer is added there.
  112. wanted_locations = _clean_locations(locations)
  113. current_locations = set(await group_location_names(db, group.id)) if group.id is not None else set()
  114. if wanted_locations != current_locations:
  115. if group.id is None:
  116. await db.flush()
  117. await db.execute(delete(group_locations).where(group_locations.c.group_id == group.id))
  118. if wanted_locations:
  119. await db.execute(
  120. insert(group_locations),
  121. [{"group_id": group.id, "location": loc} for loc in sorted(wanted_locations)],
  122. )
  123. changed = True
  124. return changed
  125. def _group_response(group: Group, printer_ids: list[int], locations: list[str], user_count: int) -> GroupResponse:
  126. return GroupResponse(
  127. id=group.id,
  128. name=group.name,
  129. description=group.description,
  130. permissions=group.permissions or [],
  131. is_system=group.is_system,
  132. restrict_printers=bool(group.restrict_printers),
  133. printer_ids=printer_ids,
  134. locations=locations,
  135. user_count=user_count,
  136. created_at=group.created_at,
  137. updated_at=group.updated_at,
  138. )
  139. @router.get("/permissions", response_model=PermissionsListResponse)
  140. async def list_permissions(
  141. _: User | None = RequirePermissionIfAuthEnabled(Permission.GROUPS_READ),
  142. ):
  143. """List all available permissions organized by category."""
  144. categories = []
  145. for name, perms in PERMISSION_CATEGORIES.items():
  146. categories.append(
  147. PermissionCategory(
  148. name=name,
  149. permissions=[PermissionInfo(value=p.value, label=_permission_label(p)) for p in perms],
  150. )
  151. )
  152. return PermissionsListResponse(
  153. categories=categories,
  154. all_permissions=ALL_PERMISSIONS,
  155. )
  156. @router.get("", response_model=list[GroupResponse])
  157. @router.get("/", response_model=list[GroupResponse])
  158. async def list_groups(
  159. _: User | None = RequirePermissionIfAuthEnabled(Permission.GROUPS_READ),
  160. db: AsyncSession = Depends(get_db),
  161. ):
  162. """List all groups."""
  163. result = await db.execute(select(Group).options(selectinload(Group.users)).order_by(Group.name))
  164. groups = result.scalars().all()
  165. printers_by_group = await _printer_ids_by_group(db)
  166. locations_by_group = await _locations_by_group(db)
  167. return [
  168. _group_response(
  169. group, printers_by_group.get(group.id, []), locations_by_group.get(group.id, []), len(group.users)
  170. )
  171. for group in groups
  172. ]
  173. @router.post("", response_model=GroupResponse, status_code=status.HTTP_201_CREATED)
  174. @router.post("/", response_model=GroupResponse, status_code=status.HTTP_201_CREATED)
  175. async def create_group(
  176. group_data: GroupCreate,
  177. _admin: User | None = RequireAdminIfAuthEnabled(),
  178. _: User | None = RequirePermissionIfAuthEnabled(Permission.GROUPS_CREATE),
  179. db: AsyncSession = Depends(get_db),
  180. ):
  181. """Create a new group."""
  182. # Check if group name already exists
  183. existing = await db.execute(select(Group).where(Group.name == group_data.name))
  184. if existing.scalar_one_or_none():
  185. raise HTTPException(
  186. status_code=status.HTTP_400_BAD_REQUEST,
  187. detail="Group name already exists",
  188. )
  189. # Validate permissions
  190. invalid_perms = [p for p in group_data.permissions if p not in ALL_PERMISSIONS]
  191. if invalid_perms:
  192. raise HTTPException(
  193. status_code=status.HTTP_400_BAD_REQUEST,
  194. detail=f"Invalid permissions: {', '.join(invalid_perms)}",
  195. )
  196. group = Group(
  197. name=group_data.name,
  198. description=group_data.description,
  199. permissions=group_data.permissions,
  200. is_system=False, # User-created groups are not system groups
  201. restrict_printers=False,
  202. )
  203. db.add(group)
  204. await db.flush()
  205. await _apply_printer_scope(db, group, group_data.restrict_printers, group_data.printer_ids, group_data.locations)
  206. await db.commit()
  207. await db.refresh(group)
  208. return _group_response(group, await group_printer_ids(db, group.id), await group_location_names(db, group.id), 0)
  209. @router.get("/{group_id}", response_model=GroupDetailResponse)
  210. async def get_group(
  211. group_id: int,
  212. _: User | None = RequirePermissionIfAuthEnabled(Permission.GROUPS_READ),
  213. db: AsyncSession = Depends(get_db),
  214. ):
  215. """Get a group by ID with user list. Read-only — gated on
  216. ``GROUPS_READ`` only."""
  217. result = await db.execute(select(Group).where(Group.id == group_id).options(selectinload(Group.users)))
  218. group = result.scalar_one_or_none()
  219. if not group:
  220. raise HTTPException(
  221. status_code=status.HTTP_404_NOT_FOUND,
  222. detail="Group not found",
  223. )
  224. return GroupDetailResponse(
  225. id=group.id,
  226. name=group.name,
  227. description=group.description,
  228. permissions=group.permissions or [],
  229. is_system=group.is_system,
  230. restrict_printers=bool(group.restrict_printers),
  231. printer_ids=await group_printer_ids(db, group.id),
  232. locations=await group_location_names(db, group.id),
  233. user_count=len(group.users),
  234. created_at=group.created_at,
  235. updated_at=group.updated_at,
  236. users=[UserBrief(id=u.id, username=u.username, is_active=u.is_active) for u in group.users],
  237. )
  238. @router.patch("/{group_id}", response_model=GroupResponse)
  239. async def update_group(
  240. group_id: int,
  241. group_data: GroupUpdate,
  242. _admin: User | None = RequireAdminIfAuthEnabled(),
  243. _: User | None = RequirePermissionIfAuthEnabled(Permission.GROUPS_UPDATE),
  244. db: AsyncSession = Depends(get_db),
  245. ):
  246. """Update a group."""
  247. result = await db.execute(select(Group).where(Group.id == group_id).options(selectinload(Group.users)))
  248. group = result.scalar_one_or_none()
  249. if not group:
  250. raise HTTPException(
  251. status_code=status.HTTP_404_NOT_FOUND,
  252. detail="Group not found",
  253. )
  254. # Check if updating name to one that already exists
  255. if group_data.name is not None and group_data.name != group.name:
  256. existing = await db.execute(select(Group).where(Group.name == group_data.name, Group.id != group_id))
  257. if existing.scalar_one_or_none():
  258. raise HTTPException(
  259. status_code=status.HTTP_400_BAD_REQUEST,
  260. detail="Group name already exists",
  261. )
  262. # System groups cannot have their name changed
  263. if group.is_system:
  264. raise HTTPException(
  265. status_code=status.HTTP_400_BAD_REQUEST,
  266. detail="Cannot rename system groups",
  267. )
  268. group.name = group_data.name
  269. if group_data.description is not None:
  270. group.description = group_data.description
  271. if group_data.permissions is not None:
  272. # System groups (Administrators in particular) have fixed permission
  273. # sets that the app depends on — stripping them is a denial-of-
  274. # service vector that even admin callers shouldn't trigger by
  275. # accident through the generic edit form. Mirrors the rename block
  276. # immediately above.
  277. if group.is_system:
  278. raise HTTPException(
  279. status_code=status.HTTP_400_BAD_REQUEST,
  280. detail="Cannot modify permissions of system groups",
  281. )
  282. # Validate permissions
  283. invalid_perms = [p for p in group_data.permissions if p not in ALL_PERMISSIONS]
  284. if invalid_perms:
  285. raise HTTPException(
  286. status_code=status.HTTP_400_BAD_REQUEST,
  287. detail=f"Invalid permissions: {', '.join(invalid_perms)}",
  288. )
  289. group.permissions = group_data.permissions
  290. scope_changed = await _apply_printer_scope(
  291. db, group, group_data.restrict_printers, group_data.printer_ids, group_data.locations
  292. )
  293. await db.commit()
  294. await db.refresh(group)
  295. if scope_changed:
  296. await ws_manager.refresh_printer_scopes()
  297. return _group_response(
  298. group, await group_printer_ids(db, group.id), await group_location_names(db, group.id), len(group.users)
  299. )
  300. @router.delete("/{group_id}", status_code=status.HTTP_204_NO_CONTENT)
  301. async def delete_group(
  302. group_id: int,
  303. _admin: User | None = RequireAdminIfAuthEnabled(),
  304. _: User | None = RequirePermissionIfAuthEnabled(Permission.GROUPS_DELETE),
  305. db: AsyncSession = Depends(get_db),
  306. ):
  307. """Delete a group (non-system groups only)."""
  308. result = await db.execute(select(Group).where(Group.id == group_id))
  309. group = result.scalar_one_or_none()
  310. if not group:
  311. raise HTTPException(
  312. status_code=status.HTTP_404_NOT_FOUND,
  313. detail="Group not found",
  314. )
  315. if group.is_system:
  316. raise HTTPException(
  317. status_code=status.HTTP_400_BAD_REQUEST,
  318. detail="Cannot delete system groups",
  319. )
  320. restricted = bool(group.restrict_printers)
  321. # SQLite doesn't enforce the FK cascade
  322. await db.execute(delete(group_printers).where(group_printers.c.group_id == group_id))
  323. await db.execute(delete(group_locations).where(group_locations.c.group_id == group_id))
  324. await db.delete(group)
  325. await db.commit()
  326. if restricted:
  327. await ws_manager.refresh_printer_scopes()
  328. @router.post("/{group_id}/users/{user_id}", status_code=status.HTTP_204_NO_CONTENT)
  329. async def add_user_to_group(
  330. group_id: int,
  331. user_id: int,
  332. _admin: User | None = RequireAdminIfAuthEnabled(),
  333. _: User | None = RequirePermissionIfAuthEnabled(Permission.GROUPS_UPDATE),
  334. db: AsyncSession = Depends(get_db),
  335. ):
  336. """Add a user to a group."""
  337. # Get group with users
  338. result = await db.execute(select(Group).where(Group.id == group_id).options(selectinload(Group.users)))
  339. group = result.scalar_one_or_none()
  340. if not group:
  341. raise HTTPException(
  342. status_code=status.HTTP_404_NOT_FOUND,
  343. detail="Group not found",
  344. )
  345. # Get user
  346. user_result = await db.execute(select(User).where(User.id == user_id))
  347. user = user_result.scalar_one_or_none()
  348. if not user:
  349. raise HTTPException(
  350. status_code=status.HTTP_404_NOT_FOUND,
  351. detail="User not found",
  352. )
  353. # Check if user is already in group
  354. if user in group.users:
  355. raise HTTPException(
  356. status_code=status.HTTP_400_BAD_REQUEST,
  357. detail="User is already in this group",
  358. )
  359. group.users.append(user)
  360. await db.commit()
  361. if group.restrict_printers:
  362. await ws_manager.refresh_printer_scopes()
  363. @router.delete("/{group_id}/users/{user_id}", status_code=status.HTTP_204_NO_CONTENT)
  364. async def remove_user_from_group(
  365. group_id: int,
  366. user_id: int,
  367. _admin: User | None = RequireAdminIfAuthEnabled(),
  368. _: User | None = RequirePermissionIfAuthEnabled(Permission.GROUPS_UPDATE),
  369. db: AsyncSession = Depends(get_db),
  370. ):
  371. """Remove a user from a group."""
  372. # Get group with users
  373. result = await db.execute(select(Group).where(Group.id == group_id).options(selectinload(Group.users)))
  374. group = result.scalar_one_or_none()
  375. if not group:
  376. raise HTTPException(
  377. status_code=status.HTTP_404_NOT_FOUND,
  378. detail="Group not found",
  379. )
  380. # Get user
  381. user_result = await db.execute(select(User).where(User.id == user_id))
  382. user = user_result.scalar_one_or_none()
  383. if not user:
  384. raise HTTPException(
  385. status_code=status.HTTP_404_NOT_FOUND,
  386. detail="User not found",
  387. )
  388. # Check if user is in group
  389. if user not in group.users:
  390. raise HTTPException(
  391. status_code=status.HTTP_400_BAD_REQUEST,
  392. detail="User is not in this group",
  393. )
  394. group.users.remove(user)
  395. await db.commit()
  396. if group.restrict_printers:
  397. await ws_manager.refresh_printer_scopes()