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
« 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
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
28log = logging.getLogger(__name__)
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####################
37class Group(Base):
38 __tablename__ = 'group'
40 id = Column(Text, unique=True, primary_key=True)
41 user_id = Column(Text)
43 name = Column(Text)
44 description = Column(Text)
46 data = Column(JSON, nullable=True)
47 meta = Column(JSON, nullable=True)
49 permissions = Column(JSON, nullable=True)
51 created_at = Column(BigInteger)
52 updated_at = Column(BigInteger)
55class GroupModel(BaseModel):
56 id: str
57 user_id: str
59 name: str
60 description: str
62 data: Optional[dict] = None
63 meta: Optional[dict] = None
65 permissions: Optional[dict] = None
67 created_at: int # timestamp in epoch
68 updated_at: int # timestamp in epoch
70 model_config = ConfigDict(from_attributes=True)
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'),)
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)
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
97####################
98# Forms
99####################
102class GroupResponse(GroupModel):
103 member_count: Optional[int] = None
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
116class GroupForm(BaseModel):
117 name: str
118 description: str
119 permissions: Optional[dict] = None
120 data: Optional[dict] = None
123class UserIdsForm(BaseModel):
124 user_ids: Optional[list[str]] = None
127class GroupUpdateForm(GroupForm):
128 pass
131class GroupListResponse(BaseModel):
132 items: list[GroupResponse] = []
133 total: int = 0
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
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 )
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
172 except Exception:
173 return None
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]
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
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)
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"]}%'))
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)
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 )
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'))
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 )
237 result = await db.execute(stmt.order_by(Group.updated_at.desc()))
238 rows = result.all()
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 ]
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)
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 )
268 if 'share' in filter:
269 share_value = filter['share']
270 stmt = stmt.filter(Group.data.op('->>')('share') == str(share_value))
272 # Get total count
273 count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
274 total = count_result.scalar()
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()
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 }
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()]
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()
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))
334 return user_groups
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
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()
350 if not members:
351 return []
353 return [m[0] for m in members]
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()
364 group_user_ids: dict[str, list[str]] = {group_id: [] for group_id in group_ids}
366 for group_id, user_id in members:
367 group_user_ids[group_id].append(user_id)
369 return group_user_ids
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))
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 ]
391 db.add_all(new_members)
392 await db.commit()
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
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}
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
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
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()
450 return True
451 except Exception:
452 return False
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()
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 )
471 await db.execute(update(Group).filter_by(id=group.id).values(updated_at=int(time.time())))
473 await db.commit()
474 return True
476 except Exception:
477 await db.rollback()
478 return False
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}
487 new_groups = []
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
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())
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}
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()}
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
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 )
549 await db.execute(update(Group).filter(Group.id.in_(groups_to_remove)).values(updated_at=now))
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 )
563 if groups_to_add:
564 await db.execute(update(Group).filter(Group.id.in_(groups_to_add)).values(updated_at=now))
566 await db.commit()
567 return True
569 except Exception as e:
570 log.exception(e)
571 await db.rollback()
572 return False
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
587 now = int(time.time())
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
605 group.updated_at = now
606 await db.commit()
607 await db.refresh(group)
609 return GroupModel.model_validate(group)
611 except Exception as e:
612 log.exception(e)
613 return None
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
628 if not user_ids:
629 return GroupModel.model_validate(group)
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 )
636 # Update group timestamp
637 group.updated_at = int(time.time())
639 await db.commit()
640 await db.refresh(group)
641 return GroupModel.model_validate(group)
643 except Exception as e:
644 log.exception(e)
645 return None
648Groups = GroupTable()