From d4c561d9f22b6fe069c1a4e49f83d05fa6839991 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Mon, 5 Oct 2026 14:01:02 +0400 Subject: [PATCH] refac --- backend/open_webui/events.py | 2 +- .../b8e4f0a3c752_add_group_hierarchy.py | 33 + backend/open_webui/models/access_grants.py | 6 +- backend/open_webui/models/calendar.py | 6 +- backend/open_webui/models/channels.py | 14 +- backend/open_webui/models/chat_messages.py | 32 +- backend/open_webui/models/groups.py | 583 +++++++++--------- backend/open_webui/models/knowledge.py | 2 +- backend/open_webui/models/models.py | 7 +- backend/open_webui/models/notes.py | 4 +- backend/open_webui/models/prompts.py | 4 +- backend/open_webui/models/skills.py | 4 +- backend/open_webui/models/tools.py | 7 +- backend/open_webui/models/users.py | 12 +- backend/open_webui/routers/calendar.py | 2 +- backend/open_webui/routers/channels.py | 2 +- backend/open_webui/routers/folders.py | 6 +- backend/open_webui/routers/groups.py | 127 +++- backend/open_webui/routers/knowledge.py | 6 +- backend/open_webui/routers/models.py | 4 +- backend/open_webui/routers/notes.py | 2 +- backend/open_webui/routers/ollama.py | 4 +- backend/open_webui/routers/openai.py | 4 +- backend/open_webui/routers/prompts.py | 2 +- backend/open_webui/routers/scim.py | 54 +- backend/open_webui/routers/skills.py | 4 +- backend/open_webui/routers/terminals.py | 4 +- backend/open_webui/routers/tools.py | 8 +- backend/open_webui/routers/users.py | 24 +- backend/open_webui/socket/main.py | 7 +- backend/open_webui/tools/builtin.py | 32 +- backend/open_webui/tools/knowledge_fs.py | 2 +- .../utils/access_control/__init__.py | 41 +- .../open_webui/utils/access_control/files.py | 8 +- backend/open_webui/utils/headers.py | 2 +- backend/open_webui/utils/models.py | 8 +- backend/open_webui/utils/task.py | 2 +- backend/open_webui/utils/terminals.py | 4 +- backend/open_webui/utils/tools.py | 4 +- src/lib/apis/groups/index.ts | 15 + src/lib/components/admin/Users/Groups.svelte | 231 ++++++- .../admin/Users/Groups/EditGroupModal.svelte | 118 +++- .../admin/Users/Groups/General.svelte | 107 +++- .../admin/Users/Groups/GroupItem.svelte | 31 +- .../Users/Groups/InheritedMembers.svelte | 62 ++ src/lib/components/layout/Sidebar.svelte | 7 + src/lib/i18n/locales/en-US/translation.json | 23 +- src/routes/(app)/+layout.svelte | 35 ++ src/routes/+layout.svelte | 1 + 49 files changed, 1243 insertions(+), 466 deletions(-) create mode 100644 backend/open_webui/migrations/versions/b8e4f0a3c752_add_group_hierarchy.py create mode 100644 src/lib/components/admin/Users/Groups/InheritedMembers.svelte diff --git a/backend/open_webui/events.py b/backend/open_webui/events.py index e09f3bfa70..7aa7cfc72a 100644 --- a/backend/open_webui/events.py +++ b/backend/open_webui/events.py @@ -855,7 +855,7 @@ async def event_target_matches( if user_group_ids is None: from open_webui.models.groups import Groups - groups_by_user = await Groups.get_groups_by_member_ids(list(user_ids)) + groups_by_user = await Groups.get_groups_by_member_ids(list(user_ids), include_inherited=True) user_group_ids = {user_id: {group.id for group in groups} for user_id, groups in groups_by_user.items()} return any(group_ids.intersection(target_group_ids) for group_ids in user_group_ids.values()) diff --git a/backend/open_webui/migrations/versions/b8e4f0a3c752_add_group_hierarchy.py b/backend/open_webui/migrations/versions/b8e4f0a3c752_add_group_hierarchy.py new file mode 100644 index 0000000000..58977b7716 --- /dev/null +++ b/backend/open_webui/migrations/versions/b8e4f0a3c752_add_group_hierarchy.py @@ -0,0 +1,33 @@ +"""Add single-parent group hierarchy. + +Revision ID: b8e4f0a3c752 +Revises: a7d3e9f2b641 +""" + +from alembic import op +import sqlalchemy as sa + +revision = 'b8e4f0a3c752' +down_revision = 'a7d3e9f2b641' +branch_labels = None +depends_on = None + + +def upgrade(): + with op.batch_alter_table('group') as batch: + batch.add_column( + sa.Column( + 'parent_group_id', + sa.Text(), + sa.ForeignKey('group.id', name='fk_group_parent', ondelete='SET NULL'), + nullable=True, + ) + ) + batch.create_index('ix_group_parent_group_id', ['parent_group_id']) + + +def downgrade(): + with op.batch_alter_table('group') as batch: + batch.drop_index('ix_group_parent_group_id') + batch.drop_constraint('fk_group_parent', type_='foreignkey') + batch.drop_column('parent_group_id') diff --git a/backend/open_webui/models/access_grants.py b/backend/open_webui/models/access_grants.py index 49cb83a180..8055d8b49e 100644 --- a/backend/open_webui/models/access_grants.py +++ b/backend/open_webui/models/access_grants.py @@ -595,7 +595,7 @@ class AccessGrantsTable: if user_group_ids is None: from open_webui.models.groups import Groups - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = {group.id for group in user_groups} if user_group_ids: @@ -651,7 +651,7 @@ class AccessGrantsTable: if user_group_ids is None: from open_webui.models.groups import Groups - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = {group.id for group in user_groups} if user_group_ids: @@ -730,7 +730,7 @@ class AccessGrantsTable: group_ids.append(grant.principal_id) if group_ids: - group_user_ids = await Groups.get_group_user_ids_by_ids(group_ids, db=db) + group_user_ids = await Groups.get_group_user_ids_by_ids(group_ids, db=db, include_inherited=True) for members in group_user_ids.values(): user_ids.update(members) return user_ids diff --git a/backend/open_webui/models/calendar.py b/backend/open_webui/models/calendar.py index fd7390fcea..71b400996e 100644 --- a/backend/open_webui/models/calendar.py +++ b/backend/open_webui/models/calendar.py @@ -291,7 +291,7 @@ class CalendarTable: async def get_calendars_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> list[CalendarModel]: """Owned + shared calendars.""" async with get_async_db_context(db) as db: - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = [g.id for g in user_groups] stmt = select(Calendar) @@ -497,7 +497,7 @@ class CalendarEventTable: Recurring events are fetched if they have any rrule (expansion in Python). """ async with get_async_db_context(db) as db: - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = [g.id for g in user_groups] # Get calendar IDs accessible to user @@ -599,7 +599,7 @@ class CalendarEventTable: db: Optional[AsyncSession] = None, ) -> CalendarEventListResponse: async with get_async_db_context(db) as db: - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = [g.id for g in user_groups] # Get accessible calendar IDs diff --git a/backend/open_webui/models/channels.py b/backend/open_webui/models/channels.py index 6742815116..6e368caa6d 100644 --- a/backend/open_webui/models/channels.py +++ b/backend/open_webui/models/channels.py @@ -289,7 +289,7 @@ class ChannelTable: users.add(invited_by) for group_id in group_ids or []: - group_user_ids = await Groups.get_group_user_ids_by_id(group_id) + group_user_ids = await Groups.get_group_user_ids_by_id(group_id, include_inherited=True) users.update(group_user_ids) return users @@ -393,7 +393,9 @@ class ChannelTable: async def get_channels_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[ChannelModel]: async with get_async_db_context(db) as db: - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)] + user_group_ids = [ + group.id for group in await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) + ] result = await db.execute( select(Channel) @@ -737,7 +739,9 @@ class ChannelTable: return [] # Preload user's group membership - user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id, db=db)] + user_group_ids = [ + g.id for g in await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) + ] allowed_channels = [] @@ -813,7 +817,9 @@ class ChannelTable: stmt = select(Channel).filter(Channel.id == id) # Determine user groups - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)] + user_group_ids = [ + group.id for group in await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) + ] # Apply ACL rules stmt = self._has_permission( diff --git a/backend/open_webui/models/chat_messages.py b/backend/open_webui/models/chat_messages.py index 145213b8b7..0473fb1a98 100644 --- a/backend/open_webui/models/chat_messages.py +++ b/backend/open_webui/models/chat_messages.py @@ -527,7 +527,7 @@ class ChatMessageTable: db: Optional[AsyncSession] = None, ) -> dict[str, int]: async with get_async_db_context(db) as db: - from open_webui.models.groups import GroupMember + from open_webui.models.groups import group_user_memberships stmt = select(ChatMessage.model_id, func.count(ChatMessage.id).label('count')).filter( ChatMessage.role == 'assistant', @@ -539,7 +539,7 @@ class ChatMessageTable: if end_date: stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + group_users = select(group_user_memberships([group_id], True).c.user_id).scalar_subquery() stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) stmt = stmt.group_by(ChatMessage.model_id) @@ -555,7 +555,7 @@ class ChatMessageTable: ) -> dict[str, dict]: """Count distinct users and chats per model.""" async with get_async_db_context(db) as db: - from open_webui.models.groups import GroupMember + from open_webui.models.groups import group_user_memberships stmt = select( ChatMessage.model_id, @@ -571,7 +571,7 @@ class ChatMessageTable: if end_date: stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + group_users = select(group_user_memberships([group_id], True).c.user_id).scalar_subquery() stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) stmt = stmt.group_by(ChatMessage.model_id) @@ -593,7 +593,7 @@ class ChatMessageTable: ) -> dict[str, dict]: """Aggregate token usage by model using database-level aggregation.""" async with get_async_db_context(db) as db: - from open_webui.models.groups import GroupMember + from open_webui.models.groups import group_user_memberships # We need the dialect to determine JSON extraction syntax # For async sessions, access via get_bind() @@ -618,7 +618,7 @@ class ChatMessageTable: if end_date: stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + group_users = select(group_user_memberships([group_id], True).c.user_id).scalar_subquery() stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) stmt = stmt.group_by(ChatMessage.model_id) @@ -643,7 +643,7 @@ class ChatMessageTable: ) -> dict[str, dict]: """Aggregate token usage by user using database-level aggregation.""" async with get_async_db_context(db) as db: - from open_webui.models.groups import GroupMember + from open_webui.models.groups import group_user_memberships bind = await db.connection() dialect = bind.dialect.name @@ -666,7 +666,7 @@ class ChatMessageTable: if end_date: stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + group_users = select(group_user_memberships([group_id], True).c.user_id).scalar_subquery() stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) stmt = stmt.group_by(ChatMessage.user_id) @@ -917,7 +917,7 @@ class ChatMessageTable: db: Optional[AsyncSession] = None, ) -> dict[str, int]: async with get_async_db_context(db) as db: - from open_webui.models.groups import GroupMember + from open_webui.models.groups import group_user_memberships stmt = select(ChatMessage.user_id, func.count(ChatMessage.id).label('count')).filter( ChatMessage.role == 'assistant', @@ -928,7 +928,7 @@ class ChatMessageTable: if end_date: stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + group_users = select(group_user_memberships([group_id], True).c.user_id).scalar_subquery() stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) stmt = stmt.group_by(ChatMessage.user_id) @@ -943,7 +943,7 @@ class ChatMessageTable: db: Optional[AsyncSession] = None, ) -> dict[str, int]: async with get_async_db_context(db) as db: - from open_webui.models.groups import GroupMember + from open_webui.models.groups import group_user_memberships stmt = select(ChatMessage.chat_id, func.count(ChatMessage.id).label('count')).filter( ChatMessage.role == 'assistant', @@ -954,7 +954,7 @@ class ChatMessageTable: if end_date: stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + group_users = select(group_user_memberships([group_id], True).c.user_id).scalar_subquery() stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) stmt = stmt.group_by(ChatMessage.chat_id) @@ -972,7 +972,7 @@ class ChatMessageTable: async with get_async_db_context(db) as db: from datetime import datetime, timedelta - from open_webui.models.groups import GroupMember + from open_webui.models.groups import group_user_memberships stmt = select(ChatMessage.created_at, ChatMessage.model_id).filter( ChatMessage.role == 'assistant', @@ -984,7 +984,7 @@ class ChatMessageTable: if end_date: stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + group_users = select(group_user_memberships([group_id], True).c.user_id).scalar_subquery() stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) result = await db.execute(stmt) @@ -1021,7 +1021,7 @@ class ChatMessageTable: async with get_async_db_context(db) as db: from datetime import datetime, timedelta - from open_webui.models.groups import GroupMember + from open_webui.models.groups import group_user_memberships stmt = select(ChatMessage.created_at, ChatMessage.model_id).filter( ChatMessage.role == 'assistant', @@ -1033,7 +1033,7 @@ class ChatMessageTable: if end_date: stmt = stmt.filter(ChatMessage.created_at <= end_date) if group_id: - group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery() + group_users = select(group_user_memberships([group_id], True).c.user_id).scalar_subquery() stmt = stmt.filter(ChatMessage.user_id.in_(group_users)) result = await db.execute(stmt) diff --git a/backend/open_webui/models/groups.py b/backend/open_webui/models/groups.py index 60c84936e0..aa1d51e5cc 100644 --- a/backend/open_webui/models/groups.py +++ b/backend/open_webui/models/groups.py @@ -2,9 +2,10 @@ 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 +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 @@ -23,6 +24,7 @@ from sqlalchemy import ( or_, select, update, + text, ) from sqlalchemy.ext.asyncio import AsyncSession @@ -38,6 +40,10 @@ log = logging.getLogger(__name__) 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) @@ -54,6 +60,7 @@ class Group(Base): class GroupModel(BaseModel): + parent_group_id: Optional[str] = None id: str user_id: str @@ -105,6 +112,7 @@ class GroupResponse(GroupModel): class GroupInfoResponse(BaseModel): + parent_group_id: Optional[str] = None id: str user_id: str name: str @@ -115,6 +123,7 @@ class GroupInfoResponse(BaseModel): class GroupForm(BaseModel): + parent_group_id: Optional[str] = None name: str description: str permissions: Optional[dict] = None @@ -134,6 +143,95 @@ class GroupListResponse(BaseModel): total: int = 0 +class GroupHierarchyError(ValueError): + def __init__(self, message: str, status_code: int = 400): + super().__init__(message) + self.status_code = status_code + + +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.""" @@ -148,30 +246,20 @@ class GroupTable: async def insert_new_group( self, user_id: str, form_data: GroupForm, db: Optional[AsyncSession] = None ) -> Optional[GroupModel]: - async with get_async_db_context(db) as db: + 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 = GroupModel( - **{ - **group_data, - 'id': str(uuid.uuid4()), - 'user_id': user_id, - 'created_at': int(time.time()), - 'updated_at': int(time.time()), - } + group = Group( + **group_data, + id=str(uuid.uuid4()), + user_id=user_id, + created_at=int(time.time()), + updated_at=int(time.time()), ) - - try: - result = Group(**group.model_dump()) - db.add(result) - await db.commit() - await db.refresh(result) - if result: - return GroupModel.model_validate(result) - else: - return None - - except Exception: - return None + 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: @@ -217,7 +305,7 @@ class GroupTable: ) if member_id: - member_groups_select = select(GroupMember.group_id).where(GroupMember.user_id == 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), @@ -232,7 +320,7 @@ class GroupTable: # Only apply member_id filter when share filter is NOT present if 'member_id' in filter: stmt = stmt.filter( - Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id'])) + 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())) @@ -263,7 +351,7 @@ class GroupTable: stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%')) if 'member_id' in filter: stmt = stmt.filter( - Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id'])) + Group.id.in_(select(user_group_memberships([filter['member_id']], True).c.group_id)) ) if 'share' in filter: @@ -303,112 +391,100 @@ class GroupTable: 'total': total, } - async def get_groups_by_member_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[GroupModel]: - async with get_async_db_context(db) as db: - result = await db.execute( - select(Group) - .join(GroupMember, GroupMember.group_id == Group.id) - .filter(GroupMember.user_id == user_id) - .order_by(Group.updated_at.desc()) - ) - return [GroupModel.model_validate(group) for group in result.scalars().all()] + 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 + self, user_ids: list[str], db: Optional[AsyncSession] = None, *, include_inherited=False ) -> dict[str, list[GroupModel]]: - """Fetch groups for multiple users in a single query to avoid N+1.""" + 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: - # Query GroupMember joined with Group, filtering by user_ids - result = await db.execute( - select(GroupMember.user_id, Group) - .join(Group, Group.id == GroupMember.group_id) - .filter(GroupMember.user_id.in_(user_ids)) - .order_by(Group.updated_at.desc()) + 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) ) - rows = result.all() + for uid, group in rows: + groups[uid].append(GroupModel.model_validate(group)) + return groups - # Group groups by user_id - user_groups: dict[str, list[GroupModel]] = {uid: [] for uid in user_ids} - for user_id, group in rows: - user_groups[user_id].append(GroupModel.model_validate(group)) - - return user_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)) + 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) -> list[str]: - async with get_async_db_context(db) as db: - result = await db.execute(select(GroupMember.user_id).filter(GroupMember.group_id == id)) - members = result.all() - - if not members: - return [] - - return [m[0] for m in members] + 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 + 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: - result = await db.execute( - select(GroupMember.group_id, GroupMember.user_id).filter(GroupMember.group_id.in_(group_ids)) - ) - members = result.all() - - group_user_ids: dict[str, list[str]] = {group_id: [] for group_id in group_ids} - - for group_id, user_id in members: - group_user_ids[group_id].append(user_id) - - return group_user_ids + 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 get_async_db_context(db) as db: - # Delete existing members - await db.execute(delete(GroupMember).filter(GroupMember.group_id == group_id)) - - # Insert new members - now = int(time.time()) - new_members = [ - GroupMember( - id=str(uuid.uuid4()), - group_id=group_id, - user_id=user_id, - created_at=now, - updated_at=now, + 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) ) - for user_id in user_ids - ] + ) + 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) - db.add_all(new_members) - await db.commit() + 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_count_by_id(self, id: str, db: Optional[AsyncSession] = None) -> int: - async with get_async_db_context(db) as db: - result = await db.execute(select(func.count(GroupMember.user_id)).filter(GroupMember.group_id == id)) - count = result.scalar() - return count if count else 0 - - async def get_group_member_counts_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, int]: + 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: - result = await db.execute( - select(GroupMember.group_id, func.count(GroupMember.user_id)) - .filter(GroupMember.group_id.in_(ids)) - .group_by(GroupMember.group_id) + rows = await db.execute( + select(memberships.c.group_id, func.count(memberships.c.user_id)).group_by(memberships.c.group_id) ) - rows = result.all() - return {group_id: count for group_id, count in rows} + return dict(rows.all()) async def update_group_by_id( self, @@ -416,71 +492,75 @@ class GroupTable: form_data: GroupUpdateForm, overwrite: bool = False, db: Optional[AsyncSession] = None, + *, + changes: Optional[dict] = None, ) -> Optional[GroupModel]: - try: - async with get_async_db_context(db) as db: - await db.execute( - update(Group) - .filter_by(id=id) - .values( - **form_data.model_dump(exclude_none=True), - updated_at=int(time.time()), - ) + 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 + 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), ) - await db.commit() - return await self.get_group_by_id(id=id, db=db) - except Exception as e: - log.exception(e) - return None + if parent_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) -> bool: - async with get_async_db_context(db) as db: - try: - await db.execute(delete(Group).filter_by(id=id)) - await db.execute(delete(AccessGrant).filter_by(principal_type='group', principal_id=id)) - await db.commit() - return True - except Exception: - await db.rollback() - return False + 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 get_async_db_context(db) as db: - try: - await db.execute(delete(Group)) - await db.execute(delete(AccessGrant).filter_by(principal_type='group')) - await db.commit() - - return True - except Exception: - await db.rollback() - return False + 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 get_async_db_context(db) as db: - try: - # Find all groups the user belongs to - result = await db.execute( - select(Group) - .join(GroupMember, GroupMember.group_id == Group.id) - .filter(GroupMember.user_id == user_id) - ) - groups = result.scalars().all() - - # Remove the user from each group - for group in groups: - await db.execute( - delete(GroupMember).filter(GroupMember.group_id == group.id, GroupMember.user_id == user_id) - ) - - await db.execute(update(Group).filter_by(id=group.id).values(updated_at=int(time.time()))) - - await db.commit() - return True - - except Exception: - await db.rollback() - return False + 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 @@ -521,133 +601,74 @@ class GroupTable: async def sync_groups_by_group_names( self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None ) -> bool: - async with get_async_db_context(db) as db: - try: - now = int(time.time()) - - # 1. Groups that SHOULD contain the user - result = await db.execute(select(Group).filter(Group.name.in_(group_names))) - target_groups = result.scalars().all() - target_group_ids = {g.id for g in target_groups} - - # 2. Groups the user is CURRENTLY in - result = await db.execute( - select(Group) - .join(GroupMember, GroupMember.group_id == Group.id) - .filter(GroupMember.user_id == user_id) - ) - existing_group_ids = {g.id for g in result.scalars().all()} - - # 3. Determine adds + removals - groups_to_add = target_group_ids - existing_group_ids - groups_to_remove = existing_group_ids - target_group_ids - - # 4. Remove in one bulk delete - if groups_to_remove: - await db.execute( - delete(GroupMember).filter( - GroupMember.user_id == user_id, - GroupMember.group_id.in_(groups_to_remove), - ) - ) - - await db.execute(update(Group).filter(Group.id.in_(groups_to_remove)).values(updated_at=now)) - - # 5. Bulk insert missing memberships - for group_id in groups_to_add: - db.add( - GroupMember( - id=str(uuid.uuid4()), - group_id=group_id, - user_id=user_id, - created_at=now, - updated_at=now, - ) - ) - - if groups_to_add: - await db.execute(update(Group).filter(Group.id.in_(groups_to_add)).values(updated_at=now)) - - await db.commit() - return True - - except Exception as e: - log.exception(e) - await db.rollback() - return False + 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, + self, id: str, user_ids: Optional[list[str]] = None, 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)) - group = result.scalars().first() - if not group: - return None - - now = int(time.time()) - - for user_id in user_ids or []: - try: - db.add( - GroupMember( - id=str(uuid.uuid4()), - group_id=id, - user_id=user_id, - created_at=now, - updated_at=now, - ) - ) - await db.flush() # Detect unique constraint violation early - except Exception: - await db.rollback() # Clear failed INSERT - continue # Duplicate → ignore - - group.updated_at = now - await db.commit() - await db.refresh(group) - - return GroupModel.model_validate(group) - - except Exception as e: - log.exception(e) - return None + 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, + self, id: str, user_ids: Optional[list[str]] = None, 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)) - group = result.scalars().first() - if not group: - return None - - if not user_ids: - return GroupModel.model_validate(group) - - # Remove users from group_member in batch - await db.execute( - delete(GroupMember).filter(GroupMember.group_id == id, GroupMember.user_id.in_(user_ids)) - ) - - # Update group timestamp - group.updated_at = int(time.time()) - - await db.commit() - await db.refresh(group) - return GroupModel.model_validate(group) - - except Exception as e: - log.exception(e) - return None + 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() diff --git a/backend/open_webui/models/knowledge.py b/backend/open_webui/models/knowledge.py index 6322e1c161..262fac2f19 100644 --- a/backend/open_webui/models/knowledge.py +++ b/backend/open_webui/models/knowledge.py @@ -481,7 +481,7 @@ class KnowledgeTable: if knowledge.user_id == user_id: return True if user_group_ids is None: - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = {group.id for group in user_groups} return await AccessGrants.has_access( user_id=user_id, diff --git a/backend/open_webui/models/models.py b/backend/open_webui/models/models.py index d48f1e9008..a28838c16f 100755 --- a/backend/open_webui/models/models.py +++ b/backend/open_webui/models/models.py @@ -252,7 +252,10 @@ class ModelsTable: if writable_by_user_id: user_group_ids = { - group.id for group in await Groups.get_groups_by_member_id(writable_by_user_id, db=db) + group.id + for group in await Groups.get_groups_by_member_id( + writable_by_user_id, db=db, include_inherited=True + ) } stmt = self._has_permission( db, stmt, {'user_id': writable_by_user_id, 'group_ids': user_group_ids}, permission='write' @@ -475,7 +478,7 @@ class ModelsTable: ) if not is_admin: - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = [group.id for group in user_groups] filter_dict = {'user_id': user_id} diff --git a/backend/open_webui/models/notes.py b/backend/open_webui/models/notes.py index e9c06021e1..5528a0a1fa 100644 --- a/backend/open_webui/models/notes.py +++ b/backend/open_webui/models/notes.py @@ -313,7 +313,7 @@ class NoteTable: db: Optional[AsyncSession] = None, ) -> list[NoteModel]: async with get_async_db_context(db) as db: - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = [group.id for group in user_groups] stmt = select(Note).order_by(Note.updated_at.desc()) @@ -400,7 +400,7 @@ class NoteTable: db: Optional[AsyncSession] = None, ) -> list[NoteModel]: async with get_async_db_context(db) as db: - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = [group.id for group in user_groups] stmt = ( diff --git a/backend/open_webui/models/prompts.py b/backend/open_webui/models/prompts.py index 75cc93b5a8..bfdefccfb6 100644 --- a/backend/open_webui/models/prompts.py +++ b/backend/open_webui/models/prompts.py @@ -241,7 +241,7 @@ class PromptsTable: self, user_id: str, permission: str = 'write', db: AsyncSession | None = None ) -> list[PromptUserResponse]: async with get_async_db_context(db) as session: - user_groups = await Groups.get_groups_by_member_id(user_id, db=session) + user_groups = await Groups.get_groups_by_member_id(user_id, db=session, include_inherited=True) user_group_ids = [group.id for group in user_groups] query = select(Prompt).filter(Prompt.is_active == True).order_by(Prompt.updated_at.desc()) @@ -699,7 +699,7 @@ class PromptsTable: async def get_tags_by_user_id(self, user_id: str, db: AsyncSession | None = None) -> list[str]: try: async with get_async_db_context(db) as session: - user_groups = await Groups.get_groups_by_member_id(user_id, db=session) + user_groups = await Groups.get_groups_by_member_id(user_id, db=session, include_inherited=True) user_group_ids = [group.id for group in user_groups] query = select(Prompt.tags).filter(Prompt.is_active == True) diff --git a/backend/open_webui/models/skills.py b/backend/open_webui/models/skills.py index 9aabdb5bf1..4d43434a73 100644 --- a/backend/open_webui/models/skills.py +++ b/backend/open_webui/models/skills.py @@ -177,7 +177,9 @@ class SkillsTable: stmt = stmt.filter(Skill.id.in_(ids)) if user_id is not None: - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) + } stmt = AccessGrants.has_permission_filter( db=db, query=stmt, diff --git a/backend/open_webui/models/tools.py b/backend/open_webui/models/tools.py index 60c4ab9f6d..c08c6ce5da 100644 --- a/backend/open_webui/models/tools.py +++ b/backend/open_webui/models/tools.py @@ -189,7 +189,10 @@ class ToolsTable: if user_id is not None: if user_group_ids is None: - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)} + user_group_ids = { + group.id + for group in await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) + } stmt = AccessGrants.has_permission_filter( db=db, query=stmt, @@ -235,7 +238,7 @@ class ToolsTable: defer_content: bool = False, db: AsyncSession | None = None, ) -> list[ToolUserModel]: - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = {group.id for group in user_groups} return await self.get_tools( defer_content=defer_content, diff --git a/backend/open_webui/models/users.py b/backend/open_webui/models/users.py index 0235b7d787..fa3ccde4cf 100644 --- a/backend/open_webui/models/users.py +++ b/backend/open_webui/models/users.py @@ -549,7 +549,7 @@ class UsersTable: async with get_async_db_context(db) as session: # Deferred imports to avoid circular dependencies from open_webui.models.channels import ChannelMember - from open_webui.models.groups import GroupMember + from open_webui.models.groups import GroupMember, group_user_memberships # Join GroupMember so we can order by group_id when requested stmt = select(User) @@ -587,14 +587,8 @@ class UsersTable: stmt = stmt.filter(User.id.in_(user_ids)) if group_ids: - stmt = stmt.filter( - exists( - select(GroupMember.id).where( - GroupMember.user_id == User.id, - GroupMember.group_id.in_(group_ids), - ) - ) - ) + memberships = group_user_memberships(group_ids, True) + stmt = stmt.filter(User.id.in_(select(memberships.c.user_id))) roles = filter.get('roles') if roles: diff --git a/backend/open_webui/routers/calendar.py b/backend/open_webui/routers/calendar.py index acd434451d..8d2822f0a0 100644 --- a/backend/open_webui/routers/calendar.py +++ b/backend/open_webui/routers/calendar.py @@ -67,7 +67,7 @@ async def _check_calendar_access(calendar_id: str, user: UserModel, permission: raise HTTPException(status_code=404, detail='Calendar not found') if cal.user_id == user.id or user.role == 'admin': return cal - user_groups = await Groups.get_groups_by_member_id(user.id) + user_groups = await Groups.get_groups_by_member_id(user.id, include_inherited=True) user_group_ids = [g.id for g in user_groups] if await AccessGrants.has_access( user_id=user.id, diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index d8aeb65619..4fa6c221a0 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -126,7 +126,7 @@ async def get_channel_member_user_ids( user_ids = permitted_ids.get('user_ids') or [] group_ids = permitted_ids.get('group_ids') or [] if group_ids: - for member_ids in (await Groups.get_group_user_ids_by_ids(group_ids, db=db)).values(): + for member_ids in (await Groups.get_group_user_ids_by_ids(group_ids, db=db, include_inherited=True)).values(): user_ids.extend(member_ids) return list(dict.fromkeys([*user_ids, channel.user_id])) diff --git a/backend/open_webui/routers/folders.py b/backend/open_webui/routers/folders.py index cb8c9953cf..fbc3994296 100644 --- a/backend/open_webui/routers/folders.py +++ b/backend/open_webui/routers/folders.py @@ -109,7 +109,9 @@ async def get_folders( user_group_ids = None if user.role != 'admin' and any(folder.data and 'files' in folder.data for folder in folders): - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) + } # Verify folder data integrity folder_list = [] @@ -244,7 +246,7 @@ async def get_shared_folders( ): """Get all folders shared with the current user (not owned by them).""" await check_folders_permission(request, user, db=db) - groups = await Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) group_ids = {g.id for g in groups} folder_perms = await Folders.get_shared_folder_ids_for_user(user.id, group_ids, db=db) diff --git a/backend/open_webui/routers/groups.py b/backend/open_webui/routers/groups.py index 4970666b61..1701ea1c3e 100755 --- a/backend/open_webui/routers/groups.py +++ b/backend/open_webui/routers/groups.py @@ -1,16 +1,22 @@ import logging import os from pathlib import Path -from typing import Optional +from typing import Optional, Literal -from fastapi import APIRouter, Depends, HTTPException, Request, status +from fastapi import APIRouter, Depends, HTTPException, Request, status, Query from open_webui.config import CACHE_DIR +from open_webui.models.config import Config from open_webui.constants import ERROR_MESSAGES from open_webui.events import EVENTS, publish_event from open_webui.internal.db import get_async_session from open_webui.models.access_grants import AccessGrants from open_webui.models.groups import ( GroupForm, + GroupHierarchyError, + Group, + GroupMember, + group_user_memberships, + descendant_groups, GroupInfoResponse, GroupResponse, Groups, @@ -20,9 +26,13 @@ from open_webui.models.groups import ( from open_webui.models.knowledge import Knowledges from open_webui.models.models import Models from open_webui.models.tools import Tools -from open_webui.models.users import UserInfoResponse, Users +from open_webui.models.users import UserInfoResponse, Users, User from open_webui.utils.auth import get_admin_user, get_verified_user from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy import select, func, or_ +from pydantic import BaseModel +from open_webui.utils.access_control import combine_permissions +from open_webui.utils.json_codec import JSONCodec log = logging.getLogger(__name__) @@ -83,6 +93,8 @@ async def create_new_group( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT('Error creating group'), ) + except GroupHierarchyError as e: + raise HTTPException(status_code=e.status_code, detail=str(e)) from e except HTTPException: raise except Exception as e: @@ -108,7 +120,7 @@ async def get_group_by_id(id: str, user=Depends(get_admin_user), db: AsyncSessio ) else: raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, + status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) @@ -123,7 +135,7 @@ async def get_group_info_by_id(id: str, user=Depends(get_verified_user), db: Asy ) else: raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, + status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) @@ -149,7 +161,7 @@ async def export_group_by_id(id: str, user=Depends(get_admin_user), db: AsyncSes ) else: raise HTTPException( - status_code=status.HTTP_401_UNAUTHORIZED, + status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND, ) @@ -186,14 +198,15 @@ async def update_group_by_id( db: AsyncSession = Depends(get_async_session), ): try: - group = await Groups.update_group_by_id(id, form_data, db=db) + changes = {} + group = await Groups.update_group_by_id(id, form_data, db=db, changes=changes) if group: await publish_event( request, EVENTS.GROUP_UPDATED, actor=user, subject_id=id, - data={'name': group.name}, + data={'name': group.name, **changes}, ) return GroupResponse( **group.model_dump(), @@ -204,6 +217,8 @@ async def update_group_by_id( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT('Error updating group'), ) + except GroupHierarchyError as e: + raise HTTPException(status_code=e.status_code, detail=str(e)) from e except HTTPException: raise except Exception as e: @@ -249,6 +264,8 @@ async def add_user_to_group( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT('Error adding users to group'), ) + except GroupHierarchyError as e: + raise HTTPException(status_code=e.status_code, detail=str(e)) from e except HTTPException: raise except Exception as e: @@ -286,6 +303,8 @@ async def remove_users_from_group( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT('Error removing users from group'), ) + except GroupHierarchyError as e: + raise HTTPException(status_code=e.status_code, detail=str(e)) from e except HTTPException: raise except Exception as e: @@ -306,20 +325,32 @@ async def delete_group_by_id( request: Request, id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session) ): try: - result = await Groups.delete_group_by_id(id, db=db) + changes = {} + result = await Groups.delete_group_by_id(id, db=db, changes=changes) if result: await publish_event( request, EVENTS.GROUP_DELETED, actor=user, subject_id=id, + data=changes, ) + for child_id in changes['promoted_child_ids']: + await publish_event( + request, + EVENTS.GROUP_UPDATED, + actor=user, + subject_id=child_id, + data={'old_parent_group_id': id, 'parent_group_id': changes['parent_group_id']}, + ) return result else: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.DEFAULT('Error deleting group'), ) + except GroupHierarchyError as e: + raise HTTPException(status_code=e.status_code, detail=str(e)) from e except HTTPException: raise except Exception as e: @@ -349,7 +380,7 @@ async def preview_group_access( detail=ERROR_MESSAGES.NOT_FOUND, ) - group_ids = {group.id} + group_ids = await Groups.get_ancestor_ids(group.id, db=db) # Batch-check accessible resources using existing AccessGrants all_models = await Models.get_all_models(db=db) @@ -384,6 +415,14 @@ async def preview_group_access( active_models = [m for m in all_models if m.is_active] + ancestors = (await db.execute(select(Group).where(Group.id.in_(group_ids - {id})))).scalars().all() + inherited_permissions = JSONCodec.loads(JSONCodec.dumps(await Config.get('user.permissions') or {})) + for ancestor in ancestors: + inherited_permissions = combine_permissions(inherited_permissions, ancestor.permissions or {}) + effective_permissions = combine_permissions( + JSONCodec.loads(JSONCodec.dumps(inherited_permissions)), group.permissions or {} + ) + return { 'group': {'id': group.id, 'name': group.name}, 'models': { @@ -399,4 +438,72 @@ async def preview_group_access( 'total': len(all_tools), }, 'permissions': group.permissions or {}, + 'inherited_permissions': inherited_permissions, + 'effective_permissions': effective_permissions, } + + +class GroupMemberInfo(UserInfoResponse): + membership_type: Literal['direct', 'inherited'] + via_group_ids: list[str] = [] + + +class GroupMembersResponse(BaseModel): + items: list[GroupMemberInfo] + total: int + counts: dict[str, int] + + +@router.get('/id/{id}/members', response_model=GroupMembersResponse) +async def inspect_group_members( + id: str, + membership: Literal['direct', 'inherited', 'effective'] = 'effective', + query: str = '', + page: int = Query(default=1, ge=1), + user=Depends(get_admin_user), + db: AsyncSession = Depends(get_async_session), +): + if not await db.get(Group, id): + raise HTTPException(status_code=404, detail='Group not found.') + direct = select(GroupMember.user_id).where(GroupMember.group_id == id) + effective = group_user_memberships([id], True) + direct_count = (await db.execute(select(func.count()).select_from(direct.subquery()))).scalar_one() + effective_count = (await db.execute(select(func.count()).select_from(effective))).scalar_one() + stmt = select(User).where(User.id.in_(select(effective.c.user_id))) + if membership == 'direct': + stmt = stmt.where(User.id.in_(direct)) + elif membership == 'inherited': + stmt = stmt.where(User.id.not_in(direct)) + if query: + stmt = stmt.where(or_(User.name.ilike(f'%{query}%'), User.email.ilike(f'%{query}%'))) + total = (await db.execute(select(func.count()).select_from(stmt.subquery()))).scalar_one() + users = (await db.execute(stmt.order_by(User.name, User.id).offset((page - 1) * 30).limit(30))).scalars().all() + descendant = descendant_groups([id]) + sources = {u.id: [] for u in users} + if sources: + rows = await db.execute( + select(GroupMember.user_id, GroupMember.group_id).where( + GroupMember.user_id.in_(sources), GroupMember.group_id.in_(select(descendant.c.group_id)) + ) + ) + for uid, gid in rows: + sources[uid].append(gid) + return GroupMembersResponse( + items=[ + GroupMemberInfo( + id=u.id, + name=u.name, + email=u.email, + role=u.role, + membership_type='direct' if id in sources[u.id] else 'inherited', + via_group_ids=sorted(gid for gid in sources[u.id] if gid != id), + ) + for u in users + ], + total=total, + counts={ + 'direct': direct_count, + 'effective': effective_count, + 'inherited': effective_count - direct_count, + }, + ) diff --git a/backend/open_webui/routers/knowledge.py b/backend/open_webui/routers/knowledge.py index 84008db221..9041f99bff 100644 --- a/backend/open_webui/routers/knowledge.py +++ b/backend/open_webui/routers/knowledge.py @@ -180,7 +180,7 @@ async def get_knowledge_bases( skip = (page - 1) * limit filter = {} - groups = await Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) user_group_ids = {group.id for group in groups} if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL: @@ -245,7 +245,7 @@ async def search_knowledge_bases( if direction in {'asc', 'desc'}: filter['direction'] = direction - groups = await Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) user_group_ids = {group.id for group in groups} if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL: @@ -301,7 +301,7 @@ async def search_knowledge_files( if include_content: filter['include_content'] = True - groups = await Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) if groups: filter['group_ids'] = [group.id for group in groups] diff --git a/backend/open_webui/routers/models.py b/backend/open_webui/routers/models.py index f02abd5615..b04b074d7e 100644 --- a/backend/open_webui/routers/models.py +++ b/backend/open_webui/routers/models.py @@ -192,7 +192,7 @@ async def get_models( filter['direction'] = direction # Pre-fetch user group IDs once - used for both filter and write_access check - groups = await Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) user_group_ids = {group.id for group in groups} if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL: @@ -472,7 +472,7 @@ async def import_models( # per-model has_access calls (N+1 avoidance). existing_model_ids = list(existing_models.keys()) if user.role != 'admin' and existing_model_ids: - groups = await Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) user_group_ids = {group.id for group in groups} writable_model_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, diff --git a/backend/open_webui/routers/notes.py b/backend/open_webui/routers/notes.py index e44a311f45..11cd8c945e 100644 --- a/backend/open_webui/routers/notes.py +++ b/backend/open_webui/routers/notes.py @@ -201,7 +201,7 @@ async def search_notes( filter['direction'] = direction if not user.role == 'admin' or not BYPASS_ADMIN_ACCESS_CONTROL: - groups = await Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) if groups: filter['group_ids'] = [group.id for group in groups] diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index c8d33202cf..7af67756cd 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -455,7 +455,7 @@ async def get_filtered_models(models, user, db=None): """Return only the models the given *user* is allowed to access.""" model_ids = [m['model'] for m in models.get('models', [])] model_infos = {mi.id: mi for mi in await Models.get_models_by_ids(model_ids, db=db)} - user_group_ids = {g.id for g in await Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = {g.id for g in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True)} accessible_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, @@ -1548,7 +1548,7 @@ async def get_openai_models( if user.role == 'user' and not BYPASS_MODEL_ACCESS_CONTROL: model_ids = [m['id'] for m in models] model_infos = {mi.id: mi for mi in await Models.get_models_by_ids(model_ids, db=db)} - user_group_ids = {g.id for g in await Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = {g.id for g in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True)} accessible_ids = await AccessGrants.get_accessible_resource_ids( user_id=user.id, resource_type='model', diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index 19c831d9cf..8548a8ab0d 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -683,7 +683,9 @@ async def get_filtered_models(models, user, db=None): # Filter models based on user access control model_ids = [model['id'] for model in models.get('data', [])] model_infos = {model_info.id: model_info for model_info in await Models.get_models_by_ids(model_ids, db=db)} - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) + } # Batch-fetch accessible resource IDs in a single query instead of N has_access calls accessible_model_ids = await AccessGrants.get_accessible_resource_ids( diff --git a/backend/open_webui/routers/prompts.py b/backend/open_webui/routers/prompts.py index 9e90a404b0..0b6354b0b6 100644 --- a/backend/open_webui/routers/prompts.py +++ b/backend/open_webui/routers/prompts.py @@ -97,7 +97,7 @@ async def get_prompt_list( filter['direction'] = direction # Pre-fetch user group IDs once - used for both filter and write_access check - groups = await Groups.get_groups_by_member_id(user.id, db=db) + groups = await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) user_group_ids = {group.id for group in groups} if not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL): diff --git a/backend/open_webui/routers/scim.py b/backend/open_webui/routers/scim.py index 7301e9066b..2f90ac1557 100644 --- a/backend/open_webui/routers/scim.py +++ b/backend/open_webui/routers/scim.py @@ -14,12 +14,13 @@ from typing import Any, Dict, List, Optional from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status from fastapi.responses import JSONResponse +from fastapi.routing import APIRoute from open_webui.config import OAUTH_PROVIDERS from open_webui.constants import ERROR_MESSAGES from open_webui.events import EVENTS, publish_event from open_webui.env import SCIM_AUTH_PROVIDER from open_webui.internal.db import get_async_session -from open_webui.models.groups import GroupModel, Groups +from open_webui.models.groups import GroupModel, Groups, GroupHierarchyError from open_webui.models.users import UserModel, Users from open_webui.utils.auth import ( decode_token, @@ -32,7 +33,21 @@ from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) -router = APIRouter() + +class SCIMGroupRoute(APIRoute): + def get_route_handler(self): + handler = super().get_route_handler() + + async def handle(request): + try: + return await handler(request) + except GroupHierarchyError as error: + return scim_error(error.status_code, str(error), 'invalidValue') + + return handle + + +router = APIRouter(route_class=SCIMGroupRoute) # SCIM 2.0 Schema URIs SCIM_USER_SCHEMA = 'urn:ietf:params:scim:schemas:core:2.0:User' @@ -937,6 +952,19 @@ async def get_group( return await group_to_scim(group, request, db=db) +async def validate_user_members(members, db): + ids = [] + for member in members or []: + value = member if isinstance(member, dict) else member.model_dump(by_alias=True) + if value.get('type') not in (None, 'User') or '/Groups/' in (value.get('$ref') or ''): + raise GroupHierarchyError('Only direct User members are supported by SCIM.') + if not value.get('value'): + raise GroupHierarchyError('A member user ID is required.') + ids.append(value['value']) + if set(await Users.get_valid_user_ids(ids, db=db)) != set(ids): + raise GroupHierarchyError('One or more member users were not found.') + + @router.post('/Groups', response_model=SCIMGroup, status_code=status.HTTP_201_CREATED) async def create_group( request: Request, @@ -945,6 +973,7 @@ async def create_group( db: AsyncSession = Depends(get_async_session), ): """Create SCIM Group""" + await validate_user_members(group_data.members, db) # Extract member IDs member_ids = [] if group_data.members: @@ -1016,6 +1045,7 @@ async def update_group( db: AsyncSession = Depends(get_async_session), ): """Update SCIM Group (full update)""" + await validate_user_members(group_data.members, db) group = await Groups.get_group_by_id(group_id, db=db) if not group: raise HTTPException( @@ -1102,6 +1132,13 @@ async def patch_group( added_member_ids = [] removed_member_ids = [] + # Validate all requested assignments before applying any patch operation. + for operation in patch_data.Operations: + if operation.path == 'members' and operation.op.lower() in ('add', 'replace'): + if not isinstance(operation.value, list): + raise GroupHierarchyError('Members must be a list of users.') + await validate_user_members(operation.value, db) + for operation in patch_data.Operations: op = operation.op.lower() path = operation.path @@ -1184,7 +1221,8 @@ async def delete_group( detail=f'Group {group_id} not found', ) - success = await Groups.delete_group_by_id(group_id, db=db) + changes = {} + success = await Groups.delete_group_by_id(group_id, db=db, changes=changes) if not success: raise HTTPException( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, @@ -1196,7 +1234,15 @@ async def delete_group( EVENTS.GROUP_DELETED, subject_id=group_id, source='scim', - data={'name': group.name}, + data={'name': group.name, **changes}, ) + for child_id in changes['promoted_child_ids']: + await publish_event( + request, + EVENTS.GROUP_UPDATED, + subject_id=child_id, + source='scim', + data={'old_parent_group_id': group_id, 'parent_group_id': changes['parent_group_id']}, + ) return None diff --git a/backend/open_webui/routers/skills.py b/backend/open_webui/routers/skills.py index 75831fe58d..a828bac0e9 100644 --- a/backend/open_webui/routers/skills.py +++ b/backend/open_webui/routers/skills.py @@ -86,7 +86,9 @@ async def get_skill_list( filter['direction'] = direction is_bypass_admin = user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) + } if not is_bypass_admin: filter['group_ids'] = user_group_ids diff --git a/backend/open_webui/routers/terminals.py b/backend/open_webui/routers/terminals.py index 447bede9d3..7c96743366 100644 --- a/backend/open_webui/routers/terminals.py +++ b/backend/open_webui/routers/terminals.py @@ -91,7 +91,7 @@ async def list_terminal_servers(request: Request, user=Depends(get_verified_user return [] connections = await Config.get('terminal_server.connections', []) or [] - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, include_inherited=True)} return [ { @@ -129,7 +129,7 @@ async def proxy_terminal( if not connection.get('enabled', True): return JSONResponse({'error': 'Terminal server disabled'}, status_code=403) - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, include_inherited=True)} if not await has_connection_access(user, connection, user_group_ids): return JSONResponse({'error': 'Access denied'}, status_code=403) diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index 6d1ff070b3..ecd48f4506 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -81,7 +81,9 @@ async def get_tools( tools = [] bypass_access_control = user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL user_group_ids = ( - set() if bypass_access_control else {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + set() + if bypass_access_control + else {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True)} ) # Local Tools @@ -236,7 +238,9 @@ async def get_tool_list(user=Depends(get_verified_user), db: AsyncSession = Depe bypass_access_control = user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL user_group_ids = ( - set() if bypass_access_control else {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + set() + if bypass_access_control + else {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True)} ) tools = await Tools.get_tools( defer_content=True, diff --git a/backend/open_webui/routers/users.py b/backend/open_webui/routers/users.py index 10805373bd..69e634a8c5 100644 --- a/backend/open_webui/routers/users.py +++ b/backend/open_webui/routers/users.py @@ -167,8 +167,10 @@ async def search_users( @router.get('/groups') -async def get_user_groups(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): - return await Groups.get_groups_by_member_id(user.id, db=db) +async def get_user_groups( + include_inherited: bool = False, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session) +): + return await user_groups_response(user.id, include_inherited, db) ############################ @@ -1108,9 +1110,12 @@ async def delete_user_by_id( @router.get('/{user_id}/groups') async def get_user_groups_by_id( - user_id: str, user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session) + user_id: str, + include_inherited: bool = False, + user=Depends(get_admin_user), + db: AsyncSession = Depends(get_async_session), ): - return await Groups.get_groups_by_member_id(user_id, db=db) + return await user_groups_response(user_id, include_inherited, db) ############################ @@ -1133,7 +1138,7 @@ async def get_user_preview( ) # Get all group IDs this user belongs to - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = {g.id for g in user_groups} all_models = await Models.get_all_models(db=db) @@ -1189,3 +1194,12 @@ async def get_user_preview( 'total': len(all_tools), }, } + + +async def user_groups_response(user_id, include_inherited, db): + direct = await Groups.get_groups_by_member_id(user_id, db=db) + if not include_inherited: + return direct + direct_ids = {g.id for g in direct} + effective = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) + return [{**g.model_dump(), 'membership_type': 'direct' if g.id in direct_ids else 'inherited'} for g in effective] diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index da28eeca00..72d516cf22 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -434,7 +434,7 @@ async def leave_room_for_users(room: str, user_ids: list[str]): log.debug('Failed to make session %s leave room %s: %s', sid, room, e) -async def disconnect_user_sessions(user_id: str): +async def disconnect_user_sessions(user_id: str, *, refresh_access: bool = False): """Disconnect all Socket.IO sessions belonging to a user. Call this when a user's role is changed or the user is deleted so that @@ -444,6 +444,11 @@ async def disconnect_user_sessions(user_id: str): """ session_ids = get_session_ids_by_user_id(user_id) for sid in session_ids: + if refresh_access: + try: + await sio.emit('access:updated', {}, to=sid) + except Exception: + log.exception('Failed to notify session %s about changed access', sid) try: await sio.disconnect(sid) except Exception: diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index c974c68b46..ed4974dec9 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -81,7 +81,7 @@ async def _has_write_access_to_note(note, user_id: str) -> bool: from open_webui.models.access_grants import AccessGrants - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] return await AccessGrants.has_access( user_id=user_id, resource_type='note', @@ -1149,7 +1149,7 @@ async def search_notes( try: user_id = __user__.get('id') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] result = await Notes.search_notes( user_id=user_id, @@ -1246,7 +1246,7 @@ async def view_note( # Check access permission user_id = __user__.get('id') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] from open_webui.models.access_grants import AccessGrants @@ -2067,7 +2067,7 @@ async def list_knowledge_bases( from open_webui.models.knowledge import Knowledges user_id = __user__.get('id') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] result = await Knowledges.search_knowledge_bases( user_id, @@ -2127,7 +2127,7 @@ async def search_knowledge_bases( from open_webui.models.knowledge import Knowledges user_id = __user__.get('id') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] result = await Knowledges.search_knowledge_bases( user_id, @@ -2194,7 +2194,7 @@ async def search_knowledge_files( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] # When model has attached knowledge, scope to attached KBs/files only if __model_knowledge__: @@ -2663,7 +2663,7 @@ async def grep_knowledge_files( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] # Collect files to search files_to_search = [] @@ -2910,7 +2910,7 @@ async def view_knowledge_file( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] file = await Files.get_file_by_id(file_id) if not file: @@ -3060,7 +3060,7 @@ async def list_knowledge( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] knowledge_bases = [] files = [] @@ -3201,7 +3201,7 @@ async def query_knowledge_files( user_id = __user__.get('id') user_role = __user__.get('role', 'user') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] embedding_function = getattr(__request__.app.state, 'EMBEDDING_FUNCTION', None) if not embedding_function: @@ -3398,7 +3398,7 @@ async def query_knowledge_bases( from open_webui.routers.knowledge import KNOWLEDGE_BASES_COLLECTION user_id = __user__.get('id') - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] embedding_function = getattr(__request__.app.state, 'EMBEDDING_FUNCTION', None) if not embedding_function: return JSONCodec.dumps({'error': 'Embedding function not configured'}) @@ -3526,7 +3526,9 @@ async def view_skill( # Check user access user_role = __user__.get('role', 'user') if user_role != 'admin' and skill.user_id != user_id: - user_group_ids = [group.id for group in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [ + group.id for group in await Groups.get_groups_by_member_id(user_id, include_inherited=True) + ] if not await AccessGrants.has_access( user_id=user_id, resource_type='skill', @@ -4297,7 +4299,7 @@ async def create_calendar_event( from open_webui.models.access_grants import AccessGrants from open_webui.models.groups import Groups - user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] if not await AccessGrants.has_access( user_id=user_id, resource_type='calendar', @@ -4417,7 +4419,7 @@ async def update_calendar_event( if not cal: return JSONCodec.dumps({'error': 'Access denied'}) if cal.user_id != user_id and __user__.get('role') != 'admin': - user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] if not await AccessGrants.has_access( user_id=user_id, resource_type='calendar', @@ -4521,7 +4523,7 @@ async def delete_calendar_event( if not cal: return JSONCodec.dumps({'error': 'Access denied'}) if cal.user_id != user_id and __user__.get('role') != 'admin': - user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] if not await AccessGrants.has_access( user_id=user_id, resource_type='calendar', diff --git a/backend/open_webui/tools/knowledge_fs.py b/backend/open_webui/tools/knowledge_fs.py index 355ae3cc58..b956ab21a5 100644 --- a/backend/open_webui/tools/knowledge_fs.py +++ b/backend/open_webui/tools/knowledge_fs.py @@ -314,7 +314,7 @@ async def _get_accessible_kb_ids( user_id = user.get('id') user_role = user.get('role', 'user') - user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id)] + user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id, include_inherited=True)] async def _has_access(kb): return ( diff --git a/backend/open_webui/utils/access_control/__init__.py b/backend/open_webui/utils/access_control/__init__.py index 0853597b13..da690e55a5 100644 --- a/backend/open_webui/utils/access_control/__init__.py +++ b/backend/open_webui/utils/access_control/__init__.py @@ -32,6 +32,21 @@ def fill_missing_permissions(permissions: dict[str, Any], default_permissions: d return permissions +def combine_permissions(permissions: dict[str, Any], group_permissions: dict[str, Any]) -> dict[str, Any]: + """Combine permissions from multiple groups by taking the most permissive value.""" + for key, value in group_permissions.items(): + if isinstance(value, dict): + if key not in permissions: + permissions[key] = {} + permissions[key] = combine_permissions(permissions[key], value) + else: + if key not in permissions: + permissions[key] = value + else: + permissions[key] = permissions[key] or value # Use the most permissive value (True > False) + return permissions + + async def get_permissions( user_id: str, default_permissions: dict[str, Any], @@ -43,21 +58,7 @@ async def get_permissions( Permissions are nested in a dict with the permission key as the key and a boolean as the value. """ - def combine_permissions(permissions: dict[str, Any], group_permissions: dict[str, Any]) -> dict[str, Any]: - """Combine permissions from multiple groups by taking the most permissive value.""" - for key, value in group_permissions.items(): - if isinstance(value, dict): - if key not in permissions: - permissions[key] = {} - permissions[key] = combine_permissions(permissions[key], value) - else: - if key not in permissions: - permissions[key] = value - else: - permissions[key] = permissions[key] or value # Use the most permissive value (True > False) - return permissions - - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) # Deep copy default permissions to avoid modifying the original dict permissions = JSONCodec.loads(JSONCodec.dumps(default_permissions)) @@ -97,7 +98,7 @@ async def has_permission( permission_hierarchy = permission_key.split('.') # Retrieve user group permissions - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) for group in user_groups: if get_permission(group.permissions or {}, permission_hierarchy): @@ -130,7 +131,7 @@ async def has_access( return False if user_group_ids is None: - user_groups = await Groups.get_groups_by_member_id(user_id, db=db) + user_groups = await Groups.get_groups_by_member_id(user_id, db=db, include_inherited=True) user_group_ids = {group.id for group in user_groups} for grant in access_grants: @@ -174,7 +175,7 @@ async def has_connection_access( return user.role == 'admin' if user_group_ids is None: - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, include_inherited=True)} return await has_access(user.id, 'read', access_grants, user_group_ids) @@ -386,7 +387,9 @@ async def check_model_access( if user.role != 'admin': from open_webui.models.access_grants import AccessGrants - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user.id, include_inherited=True) + } if not ( user.id == model_info.user_id or await AccessGrants.has_access( diff --git a/backend/open_webui/utils/access_control/files.py b/backend/open_webui/utils/access_control/files.py index f4859acb95..67b919c3f1 100644 --- a/backend/open_webui/utils/access_control/files.py +++ b/backend/open_webui/utils/access_control/files.py @@ -48,7 +48,9 @@ async def has_access_to_file( # the user controls would gain write/delete on it (CWE-863). Read access is unaffected. knowledge_bases = await Knowledges.get_knowledges_by_file_id(file_id, db=db) if user_group_ids is None: - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) + } for knowledge_base in knowledge_bases: if ( knowledge_base.user_id == user.id @@ -124,7 +126,9 @@ async def get_accessible_folder_files( return entries if user_group_ids is None: - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) + } accessible: list[dict] = [] for entry in entries: diff --git a/backend/open_webui/utils/headers.py b/backend/open_webui/utils/headers.py index 75368a43ef..8b73ca40d8 100644 --- a/backend/open_webui/utils/headers.py +++ b/backend/open_webui/utils/headers.py @@ -101,7 +101,7 @@ async def get_user_groups_for_custom_headers( return None try: - return await Groups.get_groups_by_member_id(user.id) + return await Groups.get_groups_by_member_id(user.id, include_inherited=True) except Exception: log.exception('Failed to resolve user groups for custom headers') return None diff --git a/backend/open_webui/utils/models.py b/backend/open_webui/utils/models.py index 2d968172a7..6db8256e59 100644 --- a/backend/open_webui/utils/models.py +++ b/backend/open_webui/utils/models.py @@ -507,7 +507,9 @@ async def check_model_access(user, model, model_info=None, db=None): # base-model hop; skipped when no check below needs it. user_group_ids = None if user.id != model_info.user_id or model_info.base_model_id: - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) + } if not ( user.id == model_info.user_id @@ -547,7 +549,9 @@ async def get_filtered_models(models, user, db=None): if info: model_infos[model['id']] = info - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user.id, db=db, include_inherited=True) + } # Batch-fetch accessible resource IDs in a single query instead of N has_access calls accessible_model_ids = await AccessGrants.get_accessible_resource_ids( diff --git a/backend/open_webui/utils/task.py b/backend/open_webui/utils/task.py index 4fb5fc5ec6..5c6f50142a 100644 --- a/backend/open_webui/utils/task.py +++ b/backend/open_webui/utils/task.py @@ -64,7 +64,7 @@ async def prompt_template(template: str, user: Optional[Any] = None) -> str: try: from open_webui.models.groups import Groups - user_groups = await Groups.get_groups_by_member_id(user_id) + user_groups = await Groups.get_groups_by_member_id(user_id, include_inherited=True) groups = ', '.join(g.name for g in user_groups) except Exception: pass diff --git a/backend/open_webui/utils/terminals.py b/backend/open_webui/utils/terminals.py index f2c945ba63..dd6897e01c 100644 --- a/backend/open_webui/utils/terminals.py +++ b/backend/open_webui/utils/terminals.py @@ -153,7 +153,9 @@ async def get_terminal_json(request, user, metadata: dict, path: str, extra_para or (config.get('context_id') in {'chat_id', 'automation_id'} and not context_id) ): return None - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user_model.id)} + user_group_ids = { + group.id for group in await Groups.get_groups_by_member_id(user_model.id, include_inherited=True) + } if not await has_connection_access(user_model, connection, user_group_ids): return None diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index c00854c58a..c0e5eeac96 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -352,7 +352,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr tools_dict = {} # Get user's group memberships for access control checks - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, include_inherited=True)} # Batch-fetch all DB tools in one query instead of one per tool_id local_tool_ids = [tool_id for tool_id in tool_ids if not tool_id.startswith('server:')] @@ -1471,7 +1471,7 @@ async def get_terminal_tools( if not connection.get('enabled', True): raise RuntimeError(f"Terminal server '{terminal_id}' is disabled") - user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} + user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, include_inherited=True)} if not await has_connection_access(user, connection, user_group_ids): raise RuntimeError(f'Access denied to terminal {terminal_id}') diff --git a/src/lib/apis/groups/index.ts b/src/lib/apis/groups/index.ts index dfc6aa6a71..56f0114788 100644 --- a/src/lib/apis/groups/index.ts +++ b/src/lib/apis/groups/index.ts @@ -295,3 +295,18 @@ export const getGroupPreview = async (token: string, id: string) => { return res; }; + +export const getGroupMembers = async ( + token: string, + id: string, + membership: 'direct' | 'inherited' | 'effective' = 'effective', + query = '', + page = 1 +) => { + const params = new URLSearchParams({ membership, query, page: String(page) }); + const res = await fetch(`${WEBUI_API_BASE_URL}/groups/id/${id}/members?${params}`, { + headers: { Authorization: `Bearer ${token}` } + }); + if (!res.ok) throw (await res.json()).detail; + return res.json(); +}; diff --git a/src/lib/components/admin/Users/Groups.svelte b/src/lib/components/admin/Users/Groups.svelte index 8ba56c9de1..6a79a3ecf1 100644 --- a/src/lib/components/admin/Users/Groups.svelte +++ b/src/lib/components/admin/Users/Groups.svelte @@ -1,3 +1,7 @@ + + @@ -31,6 +51,87 @@ +
+
+ + + + { + if (!open) parentSearch = ''; + }} + > + +
+ +
+ + +
+ + {#each candidates as candidate (candidate.id)} + + {:else}

{$i18n.t('No groups found')}

{/each} +
+
+
+
+
+