diff --git a/backend/open_webui/models/groups.py b/backend/open_webui/models/groups.py index 5ed4b6b00..fc4cfb0d3 100644 --- a/backend/open_webui/models/groups.py +++ b/backend/open_webui/models/groups.py @@ -176,6 +176,11 @@ class GroupTable: groups = db.query(Group).order_by(Group.updated_at.desc()).all() return [GroupModel.model_validate(group) for group in groups] + def get_group_by_name(self, name: str, db: Optional[Session] = None) -> Optional[GroupModel]: + with get_db_context(db) as db: + group = db.query(Group).filter(Group.name == name).first() + return GroupModel.model_validate(group) if group else None + def get_groups(self, filter, db: Optional[Session] = None) -> list[GroupResponse]: with get_db_context(db) as db: member_count = ( diff --git a/backend/open_webui/routers/scim.py b/backend/open_webui/routers/scim.py index ed721ce33..56923bc44 100644 --- a/backend/open_webui/routers/scim.py +++ b/backend/open_webui/routers/scim.py @@ -790,8 +790,17 @@ async def get_groups( startIndex = max(1, startIndex) count = max(0, min(100, count)) - # Get all groups - groups_list = Groups.get_all_groups(db=db) + # Get groups, applying filter if provided + if filter: + if 'displayName eq' in filter: + display_name = filter.split('"')[1] + group = Groups.get_group_by_name(display_name, db=db) + groups_list = [group] if group else [] + else: + # Unrecognized filter — fall back to all groups + groups_list = Groups.get_all_groups(db=db) + else: + groups_list = Groups.get_all_groups(db=db) # Apply pagination total = len(groups_list)