diff --git a/backend/open_webui/models/users.py b/backend/open_webui/models/users.py index abbbe122ea..1cfbf5cb29 100644 --- a/backend/open_webui/models/users.py +++ b/backend/open_webui/models/users.py @@ -379,14 +379,69 @@ class UsersTable: row = (await session.execute(query)).scalars().first() return UserModel.model_validate(row) if row else None - async def get_users( + async def get_scim_users( self, filter: dict | None = None, + sort: dict | None = None, skip: int | None = None, limit: int | None = None, db: AsyncSession | None = None, ) -> dict: - """Paginated user listing with optional filters for role, group, and channel.""" + async with get_async_db_context(db) as session: + stmt = select(User).where(or_(User.oauth.cast(String) != 'null', User.scim.cast(String) != 'null')) + + if filter: + user_id = filter.get('id') + if user_id: + stmt = stmt.where(User.id == user_id) + + email = filter.get('email') + if email: + stmt = stmt.where(func.lower(User.email) == email.lower()) + + order_by = sort.get('order_by') if sort else None + direction = sort.get('direction') if sort else None + + if order_by == 'created_at': + stmt = stmt.order_by(User.created_at.asc() if direction == 'asc' else User.created_at.desc()) + + count_result = await session.execute(select(func.count()).select_from(stmt.subquery())) + total = count_result.scalar() + + if skip is not None: + stmt = stmt.offset(skip) + if limit is not None: + stmt = stmt.limit(limit) + + result = await session.execute(stmt) + users = result.scalars().all() + return { + 'users': [UserModel.model_validate(user) for user in users], + 'total': total, + } + + async def get_scim_user_by_id( + self, + id: str, + db: AsyncSession | None = None, + ) -> UserModel | None: + async with get_async_db_context(db) as session: + stmt = select(User).where( + User.id == id, + or_(User.oauth.cast(String) != 'null', User.scim.cast(String) != 'null'), + ) + user = (await session.execute(stmt)).scalars().first() + return UserModel.model_validate(user) if user else None + + async def get_users( + self, + filter: dict | None = None, + sort: dict | None = None, + skip: int | None = None, + limit: int | None = None, + db: AsyncSession | None = None, + ) -> dict: + """Paginated user listing with optional filters and sort.""" async with get_async_db_context(db) as session: # Deferred imports to avoid circular dependencies from open_webui.models.channels import ChannelMember @@ -447,64 +502,63 @@ class UsersTable: if exclude_roles: stmt = stmt.filter(~User.role.in_(exclude_roles)) - order_by = filter.get('order_by') - direction = filter.get('direction') + order_by = sort.get('order_by') if sort else None + direction = sort.get('direction') if sort else None - if order_by and order_by.startswith('group_id:'): - group_id = order_by.split(':', 1)[1] + if order_by and order_by.startswith('group_id:'): + group_id = order_by.split(':', 1)[1] - # Subquery that checks if the user belongs to the group - membership_exists = exists( - select(GroupMember.id).where( - GroupMember.user_id == User.id, - GroupMember.group_id == group_id, - ) + # Subquery that checks if the user belongs to the group + membership_exists = exists( + select(GroupMember.id).where( + GroupMember.user_id == User.id, + GroupMember.group_id == group_id, ) + ) - # CASE: user in group → 1, user not in group → 0 - group_sort = case((membership_exists, 1), else_=0) + # CASE: user in group → 1, user not in group → 0 + group_sort = case((membership_exists, 1), else_=0) - if direction == 'asc': - stmt = stmt.order_by(group_sort.asc(), User.name.asc()) - else: - stmt = stmt.order_by(group_sort.desc(), User.name.asc()) + if direction == 'asc': + stmt = stmt.order_by(group_sort.asc(), User.name.asc()) + else: + stmt = stmt.order_by(group_sort.desc(), User.name.asc()) - elif order_by == 'name': - if direction == 'asc': - stmt = stmt.order_by(User.name.asc()) - else: - stmt = stmt.order_by(User.name.desc()) + elif order_by == 'name': + if direction == 'asc': + stmt = stmt.order_by(User.name.asc()) + else: + stmt = stmt.order_by(User.name.desc()) - elif order_by == 'email': - if direction == 'asc': - stmt = stmt.order_by(User.email.asc()) - else: - stmt = stmt.order_by(User.email.desc()) + elif order_by == 'email': + if direction == 'asc': + stmt = stmt.order_by(User.email.asc()) + else: + stmt = stmt.order_by(User.email.desc()) - elif order_by == 'created_at': - if direction == 'asc': - stmt = stmt.order_by(User.created_at.asc()) - else: - stmt = stmt.order_by(User.created_at.desc()) + elif order_by == 'created_at': + if direction == 'asc': + stmt = stmt.order_by(User.created_at.asc()) + else: + stmt = stmt.order_by(User.created_at.desc()) - elif order_by == 'last_active_at': - if direction == 'asc': - stmt = stmt.order_by(User.last_active_at.asc()) - else: - stmt = stmt.order_by(User.last_active_at.desc()) + elif order_by == 'last_active_at': + if direction == 'asc': + stmt = stmt.order_by(User.last_active_at.asc()) + else: + stmt = stmt.order_by(User.last_active_at.desc()) - elif order_by == 'updated_at': - if direction == 'asc': - stmt = stmt.order_by(User.updated_at.asc()) - else: - stmt = stmt.order_by(User.updated_at.desc()) - elif order_by == 'role': - if direction == 'asc': - stmt = stmt.order_by(User.role.asc()) - else: - stmt = stmt.order_by(User.role.desc()) - - else: + elif order_by == 'updated_at': + if direction == 'asc': + stmt = stmt.order_by(User.updated_at.asc()) + else: + stmt = stmt.order_by(User.updated_at.desc()) + elif order_by == 'role': + if direction == 'asc': + stmt = stmt.order_by(User.role.asc()) + else: + stmt = stmt.order_by(User.role.desc()) + elif not filter: stmt = stmt.order_by(User.created_at.desc()) # Count BEFORE pagination @@ -632,7 +686,7 @@ class UsersTable: self, id: str, provider: str, - external_id: str, + external_id: str | None, db: AsyncSession | None = None, ) -> UserModel | None: """Update or insert a SCIM provider/external_id pair into the user's scim JSON field.""" diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index 1e474b4468..426139de02 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -530,10 +530,6 @@ async def get_channel_members_by_id( if query: filter['query'] = query - if order_by: - filter['order_by'] = order_by - if direction: - filter['direction'] = direction if channel.type == 'group': filter['channel_id'] = channel.id @@ -544,7 +540,13 @@ async def get_channel_members_by_id( filter['user_ids'] = permitted_ids.get('user_ids') filter['group_ids'] = permitted_ids.get('group_ids') - result = await Users.get_users(filter=filter, skip=skip, limit=limit, db=db) + result = await Users.get_users( + filter=filter, + sort={'order_by': order_by, 'direction': direction}, + skip=skip, + limit=limit, + db=db, + ) fetched_users = result['users'] total = result['total'] diff --git a/backend/open_webui/routers/scim.py b/backend/open_webui/routers/scim.py index a6b528edfa..a9bb999cb7 100644 --- a/backend/open_webui/routers/scim.py +++ b/backend/open_webui/routers/scim.py @@ -517,20 +517,30 @@ async def get_users( # Simple filter parsing - supports userName eq, externalId eq if 'userName eq' in filter: email = filter.split('"')[1] - user = await Users.get_user_by_email(email, db=db) - users_list = [user] if user else [] - total = 1 if user else 0 + response = await Users.get_scim_users(filter={'email': email}, limit=1, db=db) + users_list = response['users'] + total = response['total'] elif 'externalId eq' in filter: external_id = filter.split('"')[1] user = await find_user_by_external_id(external_id, db=db) users_list = [user] if user else [] total = 1 if user else 0 else: - response = await Users.get_users(skip=skip, limit=limit, db=db) + response = await Users.get_scim_users( + sort={'order_by': 'created_at'}, + skip=skip, + limit=limit, + db=db, + ) users_list = response['users'] total = response['total'] else: - response = await Users.get_users(skip=skip, limit=limit, db=db) + response = await Users.get_scim_users( + sort={'order_by': 'created_at'}, + skip=skip, + limit=limit, + db=db, + ) users_list = response['users'] total = response['total'] @@ -553,7 +563,7 @@ async def get_user( db: AsyncSession = Depends(get_async_session), ): """Get SCIM User by ID""" - user = await Users.get_user_by_id(user_id, db=db) + user = await Users.get_scim_user_by_id(user_id, db=db) if not user: return scim_error(status_code=status.HTTP_404_NOT_FOUND, detail=f'User {user_id} not found') @@ -624,11 +634,12 @@ async def create_user( detail='Failed to create user', ) - # Store externalId in the scim field - if user_data.externalId: - provider = get_scim_provider() - await Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db) - new_user = await Users.get_user_by_id(user_id, db=db) + new_user = await Users.update_user_scim_by_id(user_id, get_scim_provider(), user_data.externalId, db=db) + if not new_user: + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail='Failed to stamp SCIM user', + ) await publish_event( request, @@ -654,7 +665,7 @@ async def update_user( db: AsyncSession = Depends(get_async_session), ): """Update SCIM User (full update)""" - user = await Users.get_user_by_id(user_id, db=db) + user = await Users.get_scim_user_by_id(user_id, db=db) if not user: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -734,7 +745,7 @@ async def patch_user( db: AsyncSession = Depends(get_async_session), ): """Update SCIM User (partial update)""" - user = await Users.get_user_by_id(user_id, db=db) + user = await Users.get_scim_user_by_id(user_id, db=db) if not user: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -808,7 +819,7 @@ async def delete_user( db: AsyncSession = Depends(get_async_session), ): """Delete SCIM User""" - user = await Users.get_user_by_id(user_id, db=db) + user = await Users.get_scim_user_by_id(user_id, db=db) if not user: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, diff --git a/backend/open_webui/routers/users.py b/backend/open_webui/routers/users.py index 43d4499547..2078ebc0c4 100644 --- a/backend/open_webui/routers/users.py +++ b/backend/open_webui/routers/users.py @@ -112,14 +112,14 @@ async def get_users( filter = {} if query: filter['query'] = query - if order_by: - filter['order_by'] = order_by - if direction: - filter['direction'] = direction - filter['direction'] = direction - - result = await Users.get_users(filter=filter, skip=skip, limit=limit, db=db) + result = await Users.get_users( + filter=filter, + sort={'order_by': order_by, 'direction': direction}, + skip=skip, + limit=limit, + db=db, + ) users = result['users'] total = result['total'] @@ -167,12 +167,14 @@ async def search_users( filter = {} if query: filter['query'] = query - if order_by: - filter['order_by'] = order_by - if direction: - filter['direction'] = direction - return await Users.get_users(filter=filter, skip=skip, limit=limit, db=db) + return await Users.get_users( + filter=filter, + sort={'order_by': order_by, 'direction': direction}, + skip=skip, + limit=limit, + db=db, + ) ############################