From 33308022f0f1d87baaeeb7c71cf5b88e774030f1 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sun, 15 Feb 2026 23:57:40 -0600 Subject: [PATCH] refac --- backend/open_webui/models/groups.py | 67 ++++++++++++++++++----------- 1 file changed, 42 insertions(+), 25 deletions(-) diff --git a/backend/open_webui/models/groups.py b/backend/open_webui/models/groups.py index 8fe720ecc..a32f36f00 100644 --- a/backend/open_webui/models/groups.py +++ b/backend/open_webui/models/groups.py @@ -164,7 +164,10 @@ class GroupTable: def get_groups(self, filter, db: Optional[Session] = None) -> list[GroupResponse]: with get_db_context(db) as db: - query = db.query(Group) + member_count = func.count(GroupMember.user_id).label("member_count") + query = db.query(Group, member_count).outerjoin( + GroupMember, GroupMember.group_id == Group.id + ) if filter: if "query" in filter: @@ -179,9 +182,6 @@ class GroupTable: json_share_lower = func.lower(json_share_str) if share_value: - # Groups open to anyone: data is null, config.share is null, or share is true - # Use case-insensitive string comparison to handle variations like "True", "TRUE" - # Handle potential JSON boolean to string casting issues by checking for both string 'true' and boolean equivalence if possible, anyone_can_share = or_( Group.data.is_(None), json_share_str.is_(None), @@ -190,7 +190,6 @@ class GroupTable: ) if member_id: - # Also include member-only groups where user is a member member_groups_select = select(GroupMember.group_id).where( GroupMember.user_id == member_id ) @@ -211,21 +210,28 @@ class GroupTable: else: # Only apply member_id filter when share filter is NOT present if "member_id" in filter: - query = query.join( - GroupMember, GroupMember.group_id == Group.id - ).filter(GroupMember.user_id == filter["member_id"]) + query = query.filter( + Group.id.in_( + select(GroupMember.group_id).where( + GroupMember.user_id == filter["member_id"] + ) + ) + ) + + results = ( + query.group_by(Group.id) + .order_by(Group.updated_at.desc()) + .all() + ) - groups = query.order_by(Group.updated_at.desc()).all() - group_ids = [group.id for group in groups] - member_counts = self.get_group_member_counts_by_ids(group_ids, db=db) return [ GroupResponse.model_validate( { **GroupModel.model_validate(group).model_dump(), - "member_count": member_counts.get(group.id, 0), + "member_count": count or 0, } ) - for group in groups + for group, count in results ] def search_groups( @@ -242,31 +248,42 @@ class GroupTable: if "query" in filter: query = query.filter(Group.name.ilike(f"%{filter['query']}%")) if "member_id" in filter: - query = query.join( - GroupMember, GroupMember.group_id == Group.id - ).filter(GroupMember.user_id == filter["member_id"]) + query = query.filter( + Group.id.in_( + select(GroupMember.group_id).where( + GroupMember.user_id == filter["member_id"] + ) + ) + ) if "share" in filter: - # 'share' is stored in data JSON, support both sqlite and postgres share_value = filter["share"] - print("Filtering by share:", share_value) query = query.filter( Group.data.op("->>")("share") == str(share_value) ) total = query.count() - query = query.order_by(Group.updated_at.desc()) - groups = query.offset(skip).limit(limit).all() - group_ids = [group.id for group in groups] - member_counts = self.get_group_member_counts_by_ids(group_ids, db=db) + + member_count = func.count(GroupMember.user_id).label("member_count") + results = ( + query.add_columns(member_count) + .outerjoin(GroupMember, GroupMember.group_id == Group.id) + .group_by(Group.id) + .order_by(Group.updated_at.desc()) + .offset(skip) + .limit(limit) + .all() + ) return { "items": [ GroupResponse.model_validate( - **GroupModel.model_validate(group).model_dump(), - member_count=member_counts.get(group.id, 0), + { + **GroupModel.model_validate(group).model_dump(), + "member_count": count or 0, + } ) - for group in groups + for group, count in results ], "total": total, }