mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-25 08:13:20 +00:00
refac
This commit is contained in:
parent
2578174637
commit
fb4f476316
4 changed files with 151 additions and 82 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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']
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
############################
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue