import logging import time import uuid from typing import Optional from contextlib import asynccontextmanager from open_webui.env import DEFAULT_GROUP_SHARE_PERMISSION from open_webui.internal.db import Base, JSONField, get_async_db_context, get_async_db from open_webui.models.access_grants import AccessGrant from open_webui.models.files import FileMetadataResponse from pydantic import BaseModel, ConfigDict, field_validator from sqlalchemy import ( JSON, BigInteger, Column, ForeignKey, Index, String, Text, and_, cast, delete, func, or_, select, update, text, ) from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) #################### # UserGroup DB Schema # Let none who belong to this house be turned away, # and let the covenant hold for every member. #################### class Group(Base): __tablename__ = 'group' parent_group_id = Column( Text, ForeignKey('group.id', name='fk_group_parent', ondelete='SET NULL'), nullable=True, index=True ) id = Column(Text, unique=True, primary_key=True) user_id = Column(Text) name = Column(Text) description = Column(Text) data = Column(JSON, nullable=True) meta = Column(JSON, nullable=True) permissions = Column(JSON, nullable=True) created_at = Column(BigInteger) updated_at = Column(BigInteger) class GroupModel(BaseModel): parent_group_id: Optional[str] = None id: str user_id: str name: str description: str data: Optional[dict] = None meta: Optional[dict] = None permissions: Optional[dict] = None created_at: int # timestamp in epoch updated_at: int # timestamp in epoch model_config = ConfigDict(from_attributes=True) class GroupMember(Base): __tablename__ = 'group_member' # The table's (group_id, user_id) unique constraint cannot serve user_id lookups. __table_args__ = (Index('ix_group_member_user_id_group_id', 'user_id', 'group_id'),) id = Column(Text, unique=True, primary_key=True) group_id = Column( Text, ForeignKey('group.id', ondelete='CASCADE'), nullable=False, ) user_id = Column(Text, nullable=False) created_at = Column(BigInteger, nullable=True) updated_at = Column(BigInteger, nullable=True) class GroupMemberModel(BaseModel): id: str group_id: str user_id: str created_at: Optional[int] = None # timestamp in epoch updated_at: Optional[int] = None # timestamp in epoch #################### # Forms #################### class GroupResponse(GroupModel): member_count: Optional[int] = None class GroupInfoResponse(BaseModel): parent_group_id: Optional[str] = None id: str user_id: str name: str description: str member_count: Optional[int] = None created_at: int updated_at: int class GroupForm(BaseModel): parent_group_id: Optional[str] = None name: str description: str permissions: Optional[dict] = None data: Optional[dict] = None @field_validator('data') @classmethod def validate_default_models(cls, data): if data is None or 'config' not in data: return data config = data['config'] if not isinstance(config, dict): raise ValueError('Group config must be an object.') if 'default_models' not in config or config['default_models'] is None: return data models = config['default_models'] if not isinstance(models, list) or any(not isinstance(model, str) or not model.strip() for model in models): raise ValueError('Default models must be a list of non-empty model IDs.') return {**data, 'config': {**config, 'default_models': list(dict.fromkeys(model.strip() for model in models))}} class UserIdsForm(BaseModel): user_ids: Optional[list[str]] = None class GroupUpdateForm(GroupForm): pass class GroupListResponse(BaseModel): items: list[GroupResponse] = [] total: int = 0 class GroupHierarchyError(ValueError): def __init__(self, message: str, status_code: int = 400): super().__init__(message) self.status_code = status_code def group_default_models(group): return ((group.data or {}).get('config') or {}).get('default_models') or None def resolve_group_default_models(groups): """Resolve an ancestor-complete group list by depth, then creation time and ID.""" by_id = {group.id: group for group in groups} depths = {} for group in groups: path = [] seen = set() current = group while current and current.id not in depths and current.id not in seen: seen.add(current.id) path.append(current.id) current = by_id.get(current.parent_group_id) depth = depths.get(current.id, -1) if current else -1 for group_id in reversed(path): depth += 1 depths[group_id] = depth configured = [group for group in groups if group_default_models(group)] if not configured: return None, None winner = min(configured, key=lambda group: (-depths[group.id], group.created_at, group.id)) return group_default_models(winner), winner.id def ancestor_groups(group_ids): """Identifier-only recursion: UNION also terminates on externally introduced cycles.""" chain = select(Group.id.label('group_id')).where(Group.id.in_(group_ids)).cte(recursive=True) return chain.union( select(Group.parent_group_id).join(chain, Group.id == chain.c.group_id).where(Group.parent_group_id.isnot(None)) ) def descendant_groups(group_ids): chain = ( select(Group.id.label('root_id'), Group.id.label('group_id')).where(Group.id.in_(group_ids)).cte(recursive=True) ) return chain.union(select(chain.c.root_id, Group.id).join(chain, Group.parent_group_id == chain.c.group_id)) def user_group_memberships(user_ids, include_inherited=False): direct = select(GroupMember.user_id, GroupMember.group_id).where(GroupMember.user_id.in_(user_ids)) if not include_inherited: return direct.subquery() chain = direct.cte(recursive=True) return chain.union( select(chain.c.user_id, Group.parent_group_id) .join(Group, Group.id == chain.c.group_id) .where(Group.parent_group_id.isnot(None)) ) def group_user_memberships(group_ids, include_inherited=False): if not include_inherited: return select(GroupMember.group_id, GroupMember.user_id).where(GroupMember.group_id.in_(group_ids)).subquery() descendants = descendant_groups(group_ids) return ( select(descendants.c.root_id.label('group_id'), GroupMember.user_id) .join(GroupMember, GroupMember.group_id == descendants.c.group_id) .distinct() .subquery() ) @asynccontextmanager async def hierarchy_transaction(): # Own the session: callers may already have an unrelated read transaction. async with get_async_db() as db: try: if db.bind.dialect.name == 'sqlite': await db.execute(text('BEGIN IMMEDIATE')) elif db.bind.dialect.name == 'postgresql': await db.execute(text('SELECT pg_advisory_xact_lock(731947205)')) else: raise RuntimeError('Group hierarchy requires SQLite or PostgreSQL') yield db await db.commit() except Exception: await db.rollback() raise async def refresh_group_sessions(user_ids): if not user_ids: return # Import lazily to avoid models/socket import cycles. A committed write must not # be reported as failed just because a client has already disconnected. from open_webui.socket.main import disconnect_user_sessions for user_id in set(user_ids): try: await disconnect_user_sessions(user_id, refresh_access=True) except Exception: log.exception('Unable to refresh group access for user %s', user_id) async def validate_parent(db, group_id, parent_id): if parent_id is None: return if not parent_id: raise GroupHierarchyError('Parent group must be a group ID or null.') if not await db.get(Group, parent_id): raise GroupHierarchyError('Parent group not found.', 404) ancestors = ancestor_groups([parent_id]) if (await db.execute(select(ancestors.c.group_id).where(ancestors.c.group_id == group_id))).first(): raise GroupHierarchyError('A group cannot be its own parent or a descendant of itself.') class GroupTable: def _ensure_default_share_config(self, group_data: dict) -> dict: """Ensure the group data dict has a default share config if not already set.""" if 'data' not in group_data or group_data['data'] is None: group_data['data'] = {} if 'config' not in group_data['data']: group_data['data']['config'] = {} if 'share' not in group_data['data']['config']: group_data['data']['config']['share'] = DEFAULT_GROUP_SHARE_PERMISSION return group_data async def insert_new_group( self, user_id: str, form_data: GroupForm, db: Optional[AsyncSession] = None ) -> Optional[GroupModel]: async with hierarchy_transaction() as session: await validate_parent(session, None, form_data.parent_group_id) group_data = self._ensure_default_share_config(form_data.model_dump(exclude_none=True)) group = Group( **group_data, id=str(uuid.uuid4()), user_id=user_id, created_at=int(time.time()), updated_at=int(time.time()), ) session.add(group) await session.flush() result = GroupModel.model_validate(group) return result async def get_all_groups(self, db: Optional[AsyncSession] = None) -> list[GroupModel]: async with get_async_db_context(db) as db: result = await db.execute(select(Group).order_by(Group.updated_at.desc())) groups = result.scalars().all() return [GroupModel.model_validate(group) for group in groups] async def get_group_by_name(self, name: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]: async with get_async_db_context(db) as db: result = await db.execute(select(Group).filter(Group.name == name)) group = result.scalars().first() return GroupModel.model_validate(group) if group else None async def get_groups(self, filter, db: Optional[AsyncSession] = None) -> list[GroupResponse]: async with get_async_db_context(db) as db: member_count = ( select(func.count(GroupMember.user_id)) .where(GroupMember.group_id == Group.id) .correlate(Group) .scalar_subquery() .label('member_count') ) stmt = select(Group, member_count) if filter: if 'query' in filter: stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%')) # When share filter is present, member check is handled in the share logic if 'share' in filter: share_value = filter['share'] member_id = filter.get('member_id') json_share = Group.data['config']['share'] json_share_str = json_share.as_string() json_share_lower = func.lower(json_share_str) if share_value: anyone_can_share = or_( Group.data.is_(None), json_share_str.is_(None), json_share_lower == 'true', json_share_lower == '1', # Handle SQLite boolean true ) if member_id: member_groups_select = select(user_group_memberships([member_id], True).c.group_id) members_only_and_is_member = and_( json_share_lower == 'members', Group.id.in_(member_groups_select), ) stmt = stmt.filter(or_(anyone_can_share, members_only_and_is_member)) else: stmt = stmt.filter(anyone_can_share) else: stmt = stmt.filter(and_(Group.data.isnot(None), json_share_lower == 'false')) else: # Only apply member_id filter when share filter is NOT present if 'member_id' in filter: stmt = stmt.filter( Group.id.in_(select(user_group_memberships([filter['member_id']], True).c.group_id)) ) result = await db.execute(stmt.order_by(Group.updated_at.desc())) rows = result.all() return [ GroupResponse.model_validate( { **GroupModel.model_validate(group).model_dump(), 'member_count': count or 0, } ) for group, count in rows ] async def search_groups( self, filter: Optional[dict] = None, skip: int = 0, limit: int = 30, db: Optional[AsyncSession] = None, ) -> GroupListResponse: async with get_async_db_context(db) as db: stmt = select(Group) if filter: if 'query' in filter: stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%')) if 'member_id' in filter: stmt = stmt.filter( Group.id.in_(select(user_group_memberships([filter['member_id']], True).c.group_id)) ) if 'share' in filter: share_value = filter['share'] stmt = stmt.filter(Group.data.op('->>')('share') == str(share_value)) # Get total count count_result = await db.execute(select(func.count()).select_from(stmt.subquery())) total = count_result.scalar() member_count = ( select(func.count(GroupMember.user_id)) .where(GroupMember.group_id == Group.id) .correlate(Group) .scalar_subquery() .label('member_count') ) result = await db.execute( select(Group, member_count) .where(Group.id.in_(select(stmt.subquery().c.id))) .order_by(Group.updated_at.desc()) .offset(skip) .limit(limit) ) rows = result.all() return { 'items': [ GroupResponse.model_validate( { **GroupModel.model_validate(group).model_dump(), 'member_count': count or 0, } ) for group, count in rows ], 'total': total, } async def get_groups_by_member_id( self, user_id: str, db: Optional[AsyncSession] = None, *, include_inherited=False ) -> list[GroupModel]: return (await self.get_groups_by_member_ids([user_id], db=db, include_inherited=include_inherited))[user_id] async def get_groups_by_member_ids( self, user_ids: list[str], db: Optional[AsyncSession] = None, *, include_inherited=False ) -> dict[str, list[GroupModel]]: groups = {uid: [] for uid in user_ids} if not user_ids: return groups memberships = user_group_memberships(user_ids, include_inherited) async with get_async_db_context(db) as db: rows = await db.execute( select(memberships.c.user_id, Group) .join(Group, Group.id == memberships.c.group_id) .order_by(Group.updated_at.desc(), Group.id) ) for uid, group in rows: groups[uid].append(GroupModel.model_validate(group)) return groups async def get_ancestor_ids(self, group_id: str, db: Optional[AsyncSession] = None) -> set[str]: chain = ancestor_groups([group_id]) async with get_async_db_context(db) as db: return set((await db.execute(select(chain.c.group_id))).scalars()) async def get_group_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]: try: async with get_async_db_context(db) as db: result = await db.execute(select(Group).filter_by(id=id).execution_options(populate_existing=True)) group = result.scalars().first() return GroupModel.model_validate(group) if group else None except Exception: return None async def get_group_user_ids_by_id( self, id: str, db: Optional[AsyncSession] = None, *, include_inherited=False ) -> list[str]: return (await self.get_group_user_ids_by_ids([id], db=db, include_inherited=include_inherited))[id] async def get_group_user_ids_by_ids( self, group_ids: list[str], db: Optional[AsyncSession] = None, *, include_inherited=False ) -> dict[str, list[str]]: users = {gid: [] for gid in group_ids} if not group_ids: return users memberships = group_user_memberships(group_ids, include_inherited) async with get_async_db_context(db) as db: for gid, uid in await db.execute(select(memberships)): users[gid].append(uid) return users async def set_group_user_ids_by_id( self, group_id: str, user_ids: list[str], db: Optional[AsyncSession] = None ) -> None: async with hierarchy_transaction() as session: if not await session.get(Group, group_id): raise GroupHierarchyError('Group not found.', 404) previous = set( (await session.execute(select(GroupMember.user_id).where(GroupMember.group_id == group_id))).scalars() ) requested = set(user_ids) await session.execute( delete(GroupMember).where( GroupMember.group_id == group_id, GroupMember.user_id.in_(previous - requested) ) ) now = int(time.time()) session.add_all( [ GroupMember(id=str(uuid.uuid4()), group_id=group_id, user_id=uid, created_at=now, updated_at=now) for uid in requested - previous ] ) await session.execute(update(Group).where(Group.id == group_id).values(updated_at=now)) await refresh_group_sessions(previous ^ requested) async def get_group_member_count_by_id( self, id: str, db: Optional[AsyncSession] = None, *, include_inherited=False ) -> int: return (await self.get_group_member_counts_by_ids([id], db=db, include_inherited=include_inherited)).get(id, 0) async def get_group_member_counts_by_ids( self, ids: list[str], db: Optional[AsyncSession] = None, *, include_inherited=False ) -> dict[str, int]: if not ids: return {} memberships = group_user_memberships(ids, include_inherited) async with get_async_db_context(db) as db: rows = await db.execute( select(memberships.c.group_id, func.count(memberships.c.user_id)).group_by(memberships.c.group_id) ) return dict(rows.all()) async def update_group_by_id( self, id: str, form_data: GroupUpdateForm, overwrite: bool = False, db: Optional[AsyncSession] = None, *, changes: Optional[dict] = None, ) -> Optional[GroupModel]: affected = [] async with hierarchy_transaction() as session: group = await session.get(Group, id) if group is None: raise GroupHierarchyError('Group not found.', 404) values = form_data.model_dump(exclude_none=True) if 'parent_group_id' in form_data.model_fields_set: await validate_parent(session, id, form_data.parent_group_id) values['parent_group_id'] = form_data.parent_group_id if 'data' in values: old_data = group.data or {} new_data = values['data'] values['data'] = { **old_data, **new_data, 'config': {**(old_data.get('config') or {}), **(new_data.get('config') or {})}, } defaults_changed = 'data' in values and ( (values['data'].get('config') or {}).get('default_models') or None ) != group_default_models(group) parent_changed = values.get('parent_group_id', group.parent_group_id) != group.parent_group_id if changes is not None: changes.update( old_parent_group_id=group.parent_group_id, parent_group_id=values.get('parent_group_id', group.parent_group_id), ) if ( parent_changed or defaults_changed or ('permissions' in values and values['permissions'] != group.permissions) ): members = group_user_memberships([id], True) affected = list((await session.execute(select(members.c.user_id))).scalars()) for key, value in values.items(): setattr(group, key, value) group.updated_at = int(time.time()) await session.flush() result = GroupModel.model_validate(group) await refresh_group_sessions(affected) return result async def delete_group_by_id( self, id: str, db: Optional[AsyncSession] = None, *, changes: Optional[dict] = None ) -> bool: async with hierarchy_transaction() as session: group = await session.get(Group, id) if group is None: raise GroupHierarchyError('Group not found.', 404) members = group_user_memberships([id], True) affected = list((await session.execute(select(members.c.user_id))).scalars()) children = list((await session.execute(select(Group.id).where(Group.parent_group_id == id))).scalars()) if changes is not None: changes.update(parent_group_id=group.parent_group_id, promoted_child_ids=children) await session.execute( update(Group) .where(Group.parent_group_id == id) .values(parent_group_id=group.parent_group_id, updated_at=int(time.time())) ) await session.execute(delete(GroupMember).where(GroupMember.group_id == id)) await session.execute(delete(AccessGrant).filter_by(principal_type='group', principal_id=id)) await session.execute(delete(Group).where(Group.id == id)) await refresh_group_sessions(affected) return True async def delete_all_groups(self, db: Optional[AsyncSession] = None) -> bool: async with hierarchy_transaction() as session: affected = list((await session.execute(select(GroupMember.user_id).distinct())).scalars()) await session.execute(update(Group).values(parent_group_id=None)) await session.execute(delete(GroupMember)) await session.execute(delete(AccessGrant).filter_by(principal_type='group')) await session.execute(delete(Group)) await refresh_group_sessions(affected) return True async def remove_user_from_all_groups(self, user_id: str, db: Optional[AsyncSession] = None) -> bool: async with hierarchy_transaction() as session: ids = select(GroupMember.group_id).where(GroupMember.user_id == user_id) await session.execute(update(Group).where(Group.id.in_(ids)).values(updated_at=int(time.time()))) await session.execute(delete(GroupMember).where(GroupMember.user_id == user_id)) await refresh_group_sessions([user_id]) return True async def create_groups_by_group_names( self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None ) -> list[GroupModel]: # check for existing groups existing_groups = await self.get_all_groups(db=db) existing_group_names = {group.name for group in existing_groups} new_groups = [] async with get_async_db_context(db) as db: for group_name in group_names: if group_name not in existing_group_names: new_group = GroupModel( id=str(uuid.uuid4()), user_id=user_id, name=group_name, description='', data={ 'config': { 'share': DEFAULT_GROUP_SHARE_PERMISSION, } }, created_at=int(time.time()), updated_at=int(time.time()), ) try: result = Group(**new_group.model_dump()) db.add(result) await db.commit() await db.refresh(result) new_groups.append(GroupModel.model_validate(result)) except Exception as e: log.exception(e) continue return new_groups async def sync_groups_by_group_names( self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None ) -> bool: async with hierarchy_transaction() as session: target = set((await session.execute(select(Group.id).where(Group.name.in_(group_names)))).scalars()) previous = set( (await session.execute(select(GroupMember.group_id).where(GroupMember.user_id == user_id))).scalars() ) await session.execute( delete(GroupMember).where(GroupMember.user_id == user_id, GroupMember.group_id.in_(previous - target)) ) now = int(time.time()) session.add_all( [ GroupMember(id=str(uuid.uuid4()), group_id=gid, user_id=user_id, created_at=now, updated_at=now) for gid in target - previous ] ) await session.execute(update(Group).where(Group.id.in_(previous ^ target)).values(updated_at=now)) if previous != target: await refresh_group_sessions([user_id]) return True async def add_users_to_group( self, id: str, user_ids: Optional[list[str]] = None, db: Optional[AsyncSession] = None ) -> Optional[GroupModel]: async with hierarchy_transaction() as session: group = await session.get(Group, id) if group is None: raise GroupHierarchyError('Group not found.', 404) previous = set( (await session.execute(select(GroupMember.user_id).where(GroupMember.group_id == id))).scalars() ) added = set(user_ids or []) - previous now = int(time.time()) session.add_all( [ GroupMember(id=str(uuid.uuid4()), group_id=id, user_id=uid, created_at=now, updated_at=now) for uid in added ] ) group.updated_at = now await session.flush() result = GroupModel.model_validate(group) await refresh_group_sessions(added) return result async def remove_users_from_group( self, id: str, user_ids: Optional[list[str]] = None, db: Optional[AsyncSession] = None ) -> Optional[GroupModel]: async with hierarchy_transaction() as session: group = await session.get(Group, id) if group is None: raise GroupHierarchyError('Group not found.', 404) removed = list( ( await session.execute( select(GroupMember.user_id).where( GroupMember.group_id == id, GroupMember.user_id.in_(user_ids or []) ) ) ).scalars() ) await session.execute( delete(GroupMember).where(GroupMember.group_id == id, GroupMember.user_id.in_(removed)) ) group.updated_at = int(time.time()) await session.flush() result = GroupModel.model_validate(group) await refresh_group_sessions(removed) return result Groups = GroupTable()