Coverage for open_webui/models/groups.py: 42%

319 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1import logging 

2import time 

3import uuid 

4from typing import Optional 

5 

6from open_webui.env import DEFAULT_GROUP_SHARE_PERMISSION 

7from open_webui.internal.db import Base, JSONField, get_async_db_context 

8from open_webui.models.files import FileMetadataResponse 

9from pydantic import BaseModel, ConfigDict 

10from sqlalchemy import ( 

11 JSON, 

12 BigInteger, 

13 Column, 

14 ForeignKey, 

15 Index, 

16 String, 

17 Text, 

18 and_, 

19 cast, 

20 delete, 

21 func, 

22 or_, 

23 select, 

24 update, 

25) 

26from sqlalchemy.ext.asyncio import AsyncSession 

27 

28log = logging.getLogger(__name__) 

29 

30#################### 

31# UserGroup DB Schema 

32# Let none who belong to this house be turned away, 

33# and let the covenant hold for every member. 

34#################### 

35 

36 

37class Group(Base): 

38 __tablename__ = 'group' 

39 

40 id = Column(Text, unique=True, primary_key=True) 

41 user_id = Column(Text) 

42 

43 name = Column(Text) 

44 description = Column(Text) 

45 

46 data = Column(JSON, nullable=True) 

47 meta = Column(JSON, nullable=True) 

48 

49 permissions = Column(JSON, nullable=True) 

50 

51 created_at = Column(BigInteger) 

52 updated_at = Column(BigInteger) 

53 

54 

55class GroupModel(BaseModel): 

56 id: str 

57 user_id: str 

58 

59 name: str 

60 description: str 

61 

62 data: Optional[dict] = None 

63 meta: Optional[dict] = None 

64 

65 permissions: Optional[dict] = None 

66 

67 created_at: int # timestamp in epoch 

68 updated_at: int # timestamp in epoch 

69 

70 model_config = ConfigDict(from_attributes=True) 

71 

72 

73class GroupMember(Base): 

74 __tablename__ = 'group_member' 

75 # The table's (group_id, user_id) unique constraint cannot serve user_id lookups. 

76 __table_args__ = (Index('ix_group_member_user_id_group_id', 'user_id', 'group_id'),) 

77 

78 id = Column(Text, unique=True, primary_key=True) 

79 group_id = Column( 

80 Text, 

81 ForeignKey('group.id', ondelete='CASCADE'), 

82 nullable=False, 

83 ) 

84 user_id = Column(Text, nullable=False) 

85 created_at = Column(BigInteger, nullable=True) 

86 updated_at = Column(BigInteger, nullable=True) 

87 

88 

89class GroupMemberModel(BaseModel): 

90 id: str 

91 group_id: str 

92 user_id: str 

93 created_at: Optional[int] = None # timestamp in epoch 

94 updated_at: Optional[int] = None # timestamp in epoch 

95 

96 

97#################### 

98# Forms 

99#################### 

100 

101 

102class GroupResponse(GroupModel): 

103 member_count: Optional[int] = None 

104 

105 

106class GroupInfoResponse(BaseModel): 

107 id: str 

108 user_id: str 

109 name: str 

110 description: str 

111 member_count: Optional[int] = None 

112 created_at: int 

113 updated_at: int 

114 

115 

116class GroupForm(BaseModel): 

117 name: str 

118 description: str 

119 permissions: Optional[dict] = None 

120 data: Optional[dict] = None 

121 

122 

123class UserIdsForm(BaseModel): 

124 user_ids: Optional[list[str]] = None 

125 

126 

127class GroupUpdateForm(GroupForm): 

128 pass 

129 

130 

131class GroupListResponse(BaseModel): 

132 items: list[GroupResponse] = [] 

133 total: int = 0 

134 

135 

136class GroupTable: 

137 def _ensure_default_share_config(self, group_data: dict) -> dict: 

138 """Ensure the group data dict has a default share config if not already set.""" 

139 if 'data' not in group_data or group_data['data'] is None: 

140 group_data['data'] = {} 

141 if 'config' not in group_data['data']: 141 ↛ 143line 141 didn't jump to line 143 because the condition on line 141 was always true

142 group_data['data']['config'] = {} 

143 if 'share' not in group_data['data']['config']: 143 ↛ 145line 143 didn't jump to line 145 because the condition on line 143 was always true

144 group_data['data']['config']['share'] = DEFAULT_GROUP_SHARE_PERMISSION 

145 return group_data 

146 

147 async def insert_new_group( 

148 self, user_id: str, form_data: GroupForm, db: Optional[AsyncSession] = None 

149 ) -> Optional[GroupModel]: 

150 async with get_async_db_context(db) as db: 

151 group_data = self._ensure_default_share_config(form_data.model_dump(exclude_none=True)) 

152 group = GroupModel( 

153 **{ 

154 **group_data, 

155 'id': str(uuid.uuid4()), 

156 'user_id': user_id, 

157 'created_at': int(time.time()), 

158 'updated_at': int(time.time()), 

159 } 

160 ) 

161 

162 try: 

163 result = Group(**group.model_dump()) 

164 db.add(result) 

165 await db.commit() 

166 await db.refresh(result) 

167 if result: 

168 return GroupModel.model_validate(result) 

169 else: 

170 return None 

171 

172 except Exception: 

173 return None 

174 

175 async def get_all_groups(self, db: Optional[AsyncSession] = None) -> list[GroupModel]: 

176 async with get_async_db_context(db) as db: 

177 result = await db.execute(select(Group).order_by(Group.updated_at.desc())) 

178 groups = result.scalars().all() 

179 return [GroupModel.model_validate(group) for group in groups] 

180 

181 async def get_group_by_name(self, name: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]: 

182 async with get_async_db_context(db) as db: 

183 result = await db.execute(select(Group).filter(Group.name == name)) 

184 group = result.scalars().first() 

185 return GroupModel.model_validate(group) if group else None 

186 

187 async def get_groups(self, filter, db: Optional[AsyncSession] = None) -> list[GroupResponse]: 

188 async with get_async_db_context(db) as db: 

189 member_count = ( 

190 select(func.count(GroupMember.user_id)) 

191 .where(GroupMember.group_id == Group.id) 

192 .correlate(Group) 

193 .scalar_subquery() 

194 .label('member_count') 

195 ) 

196 stmt = select(Group, member_count) 

197 

198 if filter: 198 ↛ 199line 198 didn't jump to line 199 because the condition on line 198 was never true

199 if 'query' in filter: 

200 stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%')) 

201 

202 # When share filter is present, member check is handled in the share logic 

203 if 'share' in filter: 

204 share_value = filter['share'] 

205 member_id = filter.get('member_id') 

206 json_share = Group.data['config']['share'] 

207 json_share_str = json_share.as_string() 

208 json_share_lower = func.lower(json_share_str) 

209 

210 if share_value: 

211 anyone_can_share = or_( 

212 Group.data.is_(None), 

213 json_share_str.is_(None), 

214 json_share_lower == 'true', 

215 json_share_lower == '1', # Handle SQLite boolean true 

216 ) 

217 

218 if member_id: 

219 member_groups_select = select(GroupMember.group_id).where(GroupMember.user_id == member_id) 

220 members_only_and_is_member = and_( 

221 json_share_lower == 'members', 

222 Group.id.in_(member_groups_select), 

223 ) 

224 stmt = stmt.filter(or_(anyone_can_share, members_only_and_is_member)) 

225 else: 

226 stmt = stmt.filter(anyone_can_share) 

227 else: 

228 stmt = stmt.filter(and_(Group.data.isnot(None), json_share_lower == 'false')) 

229 

230 else: 

231 # Only apply member_id filter when share filter is NOT present 

232 if 'member_id' in filter: 

233 stmt = stmt.filter( 

234 Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id'])) 

235 ) 

236 

237 result = await db.execute(stmt.order_by(Group.updated_at.desc())) 

238 rows = result.all() 

239 

240 return [ 

241 GroupResponse.model_validate( 

242 { 

243 **GroupModel.model_validate(group).model_dump(), 

244 'member_count': count or 0, 

245 } 

246 ) 

247 for group, count in rows 

248 ] 

249 

250 async def search_groups( 

251 self, 

252 filter: Optional[dict] = None, 

253 skip: int = 0, 

254 limit: int = 30, 

255 db: Optional[AsyncSession] = None, 

256 ) -> GroupListResponse: 

257 async with get_async_db_context(db) as db: 

258 stmt = select(Group) 

259 

260 if filter: 

261 if 'query' in filter: 

262 stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%')) 

263 if 'member_id' in filter: 

264 stmt = stmt.filter( 

265 Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id'])) 

266 ) 

267 

268 if 'share' in filter: 

269 share_value = filter['share'] 

270 stmt = stmt.filter(Group.data.op('->>')('share') == str(share_value)) 

271 

272 # Get total count 

273 count_result = await db.execute(select(func.count()).select_from(stmt.subquery())) 

274 total = count_result.scalar() 

275 

276 member_count = ( 

277 select(func.count(GroupMember.user_id)) 

278 .where(GroupMember.group_id == Group.id) 

279 .correlate(Group) 

280 .scalar_subquery() 

281 .label('member_count') 

282 ) 

283 result = await db.execute( 

284 select(Group, member_count) 

285 .where(Group.id.in_(select(stmt.subquery().c.id))) 

286 .order_by(Group.updated_at.desc()) 

287 .offset(skip) 

288 .limit(limit) 

289 ) 

290 rows = result.all() 

291 

292 return { 

293 'items': [ 

294 GroupResponse.model_validate( 

295 { 

296 **GroupModel.model_validate(group).model_dump(), 

297 'member_count': count or 0, 

298 } 

299 ) 

300 for group, count in rows 

301 ], 

302 'total': total, 

303 } 

304 

305 async def get_groups_by_member_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[GroupModel]: 

306 async with get_async_db_context(db) as db: 

307 result = await db.execute( 

308 select(Group) 

309 .join(GroupMember, GroupMember.group_id == Group.id) 

310 .filter(GroupMember.user_id == user_id) 

311 .order_by(Group.updated_at.desc()) 

312 ) 

313 return [GroupModel.model_validate(group) for group in result.scalars().all()] 

314 

315 async def get_groups_by_member_ids( 

316 self, user_ids: list[str], db: Optional[AsyncSession] = None 

317 ) -> dict[str, list[GroupModel]]: 

318 """Fetch groups for multiple users in a single query to avoid N+1.""" 

319 async with get_async_db_context(db) as db: 

320 # Query GroupMember joined with Group, filtering by user_ids 

321 result = await db.execute( 

322 select(GroupMember.user_id, Group) 

323 .join(Group, Group.id == GroupMember.group_id) 

324 .filter(GroupMember.user_id.in_(user_ids)) 

325 .order_by(Group.updated_at.desc()) 

326 ) 

327 rows = result.all() 

328 

329 # Group groups by user_id 

330 user_groups: dict[str, list[GroupModel]] = {uid: [] for uid in user_ids} 

331 for user_id, group in rows: 

332 user_groups[user_id].append(GroupModel.model_validate(group)) 

333 

334 return user_groups 

335 

336 async def get_group_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]: 

337 try: 

338 async with get_async_db_context(db) as db: 

339 result = await db.execute(select(Group).filter_by(id=id)) 

340 group = result.scalars().first() 

341 return GroupModel.model_validate(group) if group else None 

342 except Exception: 

343 return None 

344 

345 async def get_group_user_ids_by_id(self, id: str, db: Optional[AsyncSession] = None) -> list[str]: 

346 async with get_async_db_context(db) as db: 

347 result = await db.execute(select(GroupMember.user_id).filter(GroupMember.group_id == id)) 

348 members = result.all() 

349 

350 if not members: 

351 return [] 

352 

353 return [m[0] for m in members] 

354 

355 async def get_group_user_ids_by_ids( 

356 self, group_ids: list[str], db: Optional[AsyncSession] = None 

357 ) -> dict[str, list[str]]: 

358 async with get_async_db_context(db) as db: 

359 result = await db.execute( 

360 select(GroupMember.group_id, GroupMember.user_id).filter(GroupMember.group_id.in_(group_ids)) 

361 ) 

362 members = result.all() 

363 

364 group_user_ids: dict[str, list[str]] = {group_id: [] for group_id in group_ids} 

365 

366 for group_id, user_id in members: 

367 group_user_ids[group_id].append(user_id) 

368 

369 return group_user_ids 

370 

371 async def set_group_user_ids_by_id( 

372 self, group_id: str, user_ids: list[str], db: Optional[AsyncSession] = None 

373 ) -> None: 

374 async with get_async_db_context(db) as db: 

375 # Delete existing members 

376 await db.execute(delete(GroupMember).filter(GroupMember.group_id == group_id)) 

377 

378 # Insert new members 

379 now = int(time.time()) 

380 new_members = [ 

381 GroupMember( 

382 id=str(uuid.uuid4()), 

383 group_id=group_id, 

384 user_id=user_id, 

385 created_at=now, 

386 updated_at=now, 

387 ) 

388 for user_id in user_ids 

389 ] 

390 

391 db.add_all(new_members) 

392 await db.commit() 

393 

394 async def get_group_member_count_by_id(self, id: str, db: Optional[AsyncSession] = None) -> int: 

395 async with get_async_db_context(db) as db: 

396 result = await db.execute(select(func.count(GroupMember.user_id)).filter(GroupMember.group_id == id)) 

397 count = result.scalar() 

398 return count if count else 0 

399 

400 async def get_group_member_counts_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, int]: 

401 if not ids: 

402 return {} 

403 async with get_async_db_context(db) as db: 

404 result = await db.execute( 

405 select(GroupMember.group_id, func.count(GroupMember.user_id)) 

406 .filter(GroupMember.group_id.in_(ids)) 

407 .group_by(GroupMember.group_id) 

408 ) 

409 rows = result.all() 

410 return {group_id: count for group_id, count in rows} 

411 

412 async def update_group_by_id( 

413 self, 

414 id: str, 

415 form_data: GroupUpdateForm, 

416 overwrite: bool = False, 

417 db: Optional[AsyncSession] = None, 

418 ) -> Optional[GroupModel]: 

419 try: 

420 async with get_async_db_context(db) as db: 

421 await db.execute( 

422 update(Group) 

423 .filter_by(id=id) 

424 .values( 

425 **form_data.model_dump(exclude_none=True), 

426 updated_at=int(time.time()), 

427 ) 

428 ) 

429 await db.commit() 

430 return await self.get_group_by_id(id=id, db=db) 

431 except Exception as e: 

432 log.exception(e) 

433 return None 

434 

435 async def delete_group_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool: 

436 try: 

437 async with get_async_db_context(db) as db: 

438 await db.execute(delete(Group).filter_by(id=id)) 

439 await db.commit() 

440 return True 

441 except Exception: 

442 return False 

443 

444 async def delete_all_groups(self, db: Optional[AsyncSession] = None) -> bool: 

445 async with get_async_db_context(db) as db: 

446 try: 

447 await db.execute(delete(Group)) 

448 await db.commit() 

449 

450 return True 

451 except Exception: 

452 return False 

453 

454 async def remove_user_from_all_groups(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: 

455 async with get_async_db_context(db) as db: 

456 try: 

457 # Find all groups the user belongs to 

458 result = await db.execute( 

459 select(Group) 

460 .join(GroupMember, GroupMember.group_id == Group.id) 

461 .filter(GroupMember.user_id == user_id) 

462 ) 

463 groups = result.scalars().all() 

464 

465 # Remove the user from each group 

466 for group in groups: 

467 await db.execute( 

468 delete(GroupMember).filter(GroupMember.group_id == group.id, GroupMember.user_id == user_id) 

469 ) 

470 

471 await db.execute(update(Group).filter_by(id=group.id).values(updated_at=int(time.time()))) 

472 

473 await db.commit() 

474 return True 

475 

476 except Exception: 

477 await db.rollback() 

478 return False 

479 

480 async def create_groups_by_group_names( 

481 self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None 

482 ) -> list[GroupModel]: 

483 # check for existing groups 

484 existing_groups = await self.get_all_groups(db=db) 

485 existing_group_names = {group.name for group in existing_groups} 

486 

487 new_groups = [] 

488 

489 async with get_async_db_context(db) as db: 

490 for group_name in group_names: 

491 if group_name not in existing_group_names: 

492 new_group = GroupModel( 

493 id=str(uuid.uuid4()), 

494 user_id=user_id, 

495 name=group_name, 

496 description='', 

497 data={ 

498 'config': { 

499 'share': DEFAULT_GROUP_SHARE_PERMISSION, 

500 } 

501 }, 

502 created_at=int(time.time()), 

503 updated_at=int(time.time()), 

504 ) 

505 try: 

506 result = Group(**new_group.model_dump()) 

507 db.add(result) 

508 await db.commit() 

509 await db.refresh(result) 

510 new_groups.append(GroupModel.model_validate(result)) 

511 except Exception as e: 

512 log.exception(e) 

513 continue 

514 return new_groups 

515 

516 async def sync_groups_by_group_names( 

517 self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None 

518 ) -> bool: 

519 async with get_async_db_context(db) as db: 

520 try: 

521 now = int(time.time()) 

522 

523 # 1. Groups that SHOULD contain the user 

524 result = await db.execute(select(Group).filter(Group.name.in_(group_names))) 

525 target_groups = result.scalars().all() 

526 target_group_ids = {g.id for g in target_groups} 

527 

528 # 2. Groups the user is CURRENTLY in 

529 result = await db.execute( 

530 select(Group) 

531 .join(GroupMember, GroupMember.group_id == Group.id) 

532 .filter(GroupMember.user_id == user_id) 

533 ) 

534 existing_group_ids = {g.id for g in result.scalars().all()} 

535 

536 # 3. Determine adds + removals 

537 groups_to_add = target_group_ids - existing_group_ids 

538 groups_to_remove = existing_group_ids - target_group_ids 

539 

540 # 4. Remove in one bulk delete 

541 if groups_to_remove: 

542 await db.execute( 

543 delete(GroupMember).filter( 

544 GroupMember.user_id == user_id, 

545 GroupMember.group_id.in_(groups_to_remove), 

546 ) 

547 ) 

548 

549 await db.execute(update(Group).filter(Group.id.in_(groups_to_remove)).values(updated_at=now)) 

550 

551 # 5. Bulk insert missing memberships 

552 for group_id in groups_to_add: 

553 db.add( 

554 GroupMember( 

555 id=str(uuid.uuid4()), 

556 group_id=group_id, 

557 user_id=user_id, 

558 created_at=now, 

559 updated_at=now, 

560 ) 

561 ) 

562 

563 if groups_to_add: 

564 await db.execute(update(Group).filter(Group.id.in_(groups_to_add)).values(updated_at=now)) 

565 

566 await db.commit() 

567 return True 

568 

569 except Exception as e: 

570 log.exception(e) 

571 await db.rollback() 

572 return False 

573 

574 async def add_users_to_group( 

575 self, 

576 id: str, 

577 user_ids: Optional[list[str]] = None, 

578 db: Optional[AsyncSession] = None, 

579 ) -> Optional[GroupModel]: 

580 try: 

581 async with get_async_db_context(db) as db: 

582 result = await db.execute(select(Group).filter_by(id=id)) 

583 group = result.scalars().first() 

584 if not group: 

585 return None 

586 

587 now = int(time.time()) 

588 

589 for user_id in user_ids or []: 

590 try: 

591 db.add( 

592 GroupMember( 

593 id=str(uuid.uuid4()), 

594 group_id=id, 

595 user_id=user_id, 

596 created_at=now, 

597 updated_at=now, 

598 ) 

599 ) 

600 await db.flush() # Detect unique constraint violation early 

601 except Exception: 

602 await db.rollback() # Clear failed INSERT 

603 continue # Duplicate → ignore 

604 

605 group.updated_at = now 

606 await db.commit() 

607 await db.refresh(group) 

608 

609 return GroupModel.model_validate(group) 

610 

611 except Exception as e: 

612 log.exception(e) 

613 return None 

614 

615 async def remove_users_from_group( 

616 self, 

617 id: str, 

618 user_ids: Optional[list[str]] = None, 

619 db: Optional[AsyncSession] = None, 

620 ) -> Optional[GroupModel]: 

621 try: 

622 async with get_async_db_context(db) as db: 

623 result = await db.execute(select(Group).filter_by(id=id)) 

624 group = result.scalars().first() 

625 if not group: 

626 return None 

627 

628 if not user_ids: 

629 return GroupModel.model_validate(group) 

630 

631 # Remove users from group_member in batch 

632 await db.execute( 

633 delete(GroupMember).filter(GroupMember.group_id == id, GroupMember.user_id.in_(user_ids)) 

634 ) 

635 

636 # Update group timestamp 

637 group.updated_at = int(time.time()) 

638 

639 await db.commit() 

640 await db.refresh(group) 

641 return GroupModel.model_validate(group) 

642 

643 except Exception as e: 

644 log.exception(e) 

645 return None 

646 

647 

648Groups = GroupTable()