refac
This commit is contained in:
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user