From b59f10cbbd666a9057fe2cc8712569d431681554 Mon Sep 17 00:00:00 2001 From: Jesse Hallam Date: Thu, 17 Apr 2025 17:37:28 -0300 Subject: [PATCH] MM-62158: Group Store, explicit aliases (#30741) --- server/channels/store/sqlstore/group_store.go | 54 +-- .../channels/store/storetest/group_store.go | 451 ++++++++++++++++-- 2 files changed, 430 insertions(+), 75 deletions(-) diff --git a/server/channels/store/sqlstore/group_store.go b/server/channels/store/sqlstore/group_store.go index da204373af..41b081d23f 100644 --- a/server/channels/store/sqlstore/group_store.go +++ b/server/channels/store/sqlstore/group_store.go @@ -1209,7 +1209,7 @@ func (s *SqlGroupStore) GetGroupsByChannel(channelId string, opts model.GroupSea if opts.PageOpts != nil { offset := uint64(opts.PageOpts.Page * opts.PageOpts.PerPage) - builder = builder.OrderBy("ug.DisplayName").Limit(uint64(opts.PageOpts.PerPage)).Offset(offset) + builder = builder.OrderBy("UserGroups.DisplayName").Limit(uint64(opts.PageOpts.PerPage)).Offset(offset) } groups := groupsWithSchemeAdmin{} @@ -1279,7 +1279,7 @@ func (s *SqlGroupStore) ChannelMembersToRemove(channelID *string) ([]*model.Chan func (s *SqlGroupStore) groupsBySyncableBaseQuery(st model.GroupSyncableType, t selectType, syncableID string, opts model.GroupSearchOpts) sq.SelectBuilder { selectStrs := map[selectType]string{ - selectGroups: "ug.*, gs.SchemeAdmin AS SyncableSchemeAdmin", + selectGroups: "UserGroups.*, gs.SchemeAdmin AS SyncableSchemeAdmin", selectCountGroups: "COUNT(*)", } @@ -1296,21 +1296,21 @@ func (s *SqlGroupStore) groupsBySyncableBaseQuery(st model.GroupSyncableType, t query := s.getQueryBuilder(). Select(selectStrs[t]). From(fmt.Sprintf("%s gs", table)). - LeftJoin("UserGroups ug ON gs.GroupId = ug.Id"). - Where(fmt.Sprintf("ug.DeleteAt = 0 AND gs.%s = ? AND gs.DeleteAt = 0", idCol), syncableID) + LeftJoin("UserGroups ON gs.GroupId = UserGroups.Id"). + Where(fmt.Sprintf("UserGroups.DeleteAt = 0 AND gs.%s = ? AND gs.DeleteAt = 0", idCol), syncableID) if opts.IncludeMemberCount && t == selectGroups { query = s.getQueryBuilder(). - Select(fmt.Sprintf("ug.*, coalesce(Members.MemberCount, 0) AS MemberCount, Group%ss.SchemeAdmin AS SyncableSchemeAdmin", st)). - From("UserGroups ug"). - LeftJoin("(SELECT GroupMembers.GroupId, COUNT(*) AS MemberCount FROM GroupMembers LEFT JOIN Users ON Users.Id = GroupMembers.UserId WHERE GroupMembers.DeleteAt = 0 AND Users.DeleteAt = 0 GROUP BY GroupId) AS Members ON Members.GroupId = ug.Id"). - LeftJoin(fmt.Sprintf("%[1]s ON %[1]s.GroupId = ug.Id", table)). - Where(fmt.Sprintf("ug.DeleteAt = 0 AND %[1]s.DeleteAt = 0 AND %[1]s.%[2]s = ?", table, idCol), syncableID). - OrderBy("ug.DisplayName") + Select(fmt.Sprintf("UserGroups.*, coalesce(Members.MemberCount, 0) AS MemberCount, Group%ss.SchemeAdmin AS SyncableSchemeAdmin", st)). + From("UserGroups"). + LeftJoin("(SELECT GroupMembers.GroupId, COUNT(*) AS MemberCount FROM GroupMembers LEFT JOIN Users ON Users.Id = GroupMembers.UserId WHERE GroupMembers.DeleteAt = 0 AND Users.DeleteAt = 0 GROUP BY GroupId) AS Members ON Members.GroupId = UserGroups.Id"). + LeftJoin(fmt.Sprintf("%[1]s ON %[1]s.GroupId = UserGroups.Id", table)). + Where(fmt.Sprintf("UserGroups.DeleteAt = 0 AND %[1]s.DeleteAt = 0 AND %[1]s.%[2]s = ?", table, idCol), syncableID). + OrderBy("UserGroups.DisplayName") } if opts.FilterAllowReference && t == selectGroups { - query = query.Where("ug.AllowReference = true") + query = query.Where("UserGroups.AllowReference = true") } if opts.Q != "" { @@ -1319,7 +1319,7 @@ func (s *SqlGroupStore) groupsBySyncableBaseQuery(st model.GroupSyncableType, t if s.DriverName() == model.DatabaseDriverMysql { operatorKeyword = "LIKE" } - query = query.Where(fmt.Sprintf("(ug.Name %[1]s ? OR ug.DisplayName %[1]s ?)", operatorKeyword), pattern, pattern) + query = query.Where(fmt.Sprintf("(UserGroups.Name %[1]s ? OR UserGroups.DisplayName %[1]s ?)", operatorKeyword), pattern, pattern) } return query @@ -1327,8 +1327,8 @@ func (s *SqlGroupStore) groupsBySyncableBaseQuery(st model.GroupSyncableType, t func (s *SqlGroupStore) getGroupsAssociatedToChannelsByTeam(teamID string, opts model.GroupSearchOpts) sq.SelectBuilder { query := s.getQueryBuilder(). - Select("gc.ChannelId, ug.*, gc.SchemeAdmin AS SyncableSchemeAdmin"). - From("UserGroups ug"). + Select("gc.ChannelId, UserGroups.*, gc.SchemeAdmin AS SyncableSchemeAdmin"). + From("UserGroups"). LeftJoin(` (SELECT GroupChannels.GroupId, GroupChannels.ChannelId, GroupChannels.DeleteAt, GroupChannels.SchemeAdmin @@ -1339,14 +1339,14 @@ func (s *SqlGroupStore) getGroupsAssociatedToChannelsByTeam(teamID string, opts WHERE GroupChannels.DeleteAt = 0 AND Channels.DeleteAt = 0 - AND Channels.TeamId = ?) AS gc ON gc.GroupId = ug.Id`, teamID). - Where("ug.DeleteAt = 0 AND gc.DeleteAt = 0"). - OrderBy("ug.DisplayName") + AND Channels.TeamId = ?) AS gc ON gc.GroupId = UserGroups.Id`, teamID). + Where("UserGroups.DeleteAt = 0 AND gc.DeleteAt = 0"). + OrderBy("UserGroups.DisplayName") if opts.IncludeMemberCount { query = s.getQueryBuilder(). - Select("gc.ChannelId, ug.*, coalesce(Members.MemberCount, 0) AS MemberCount, gc.SchemeAdmin AS SyncableSchemeAdmin"). - From("UserGroups ug"). + Select("gc.ChannelId, UserGroups.*, coalesce(Members.MemberCount, 0) AS MemberCount, gc.SchemeAdmin AS SyncableSchemeAdmin"). + From("UserGroups"). LeftJoin(` (SELECT GroupChannels.ChannelId, GroupChannels.DeleteAt, GroupChannels.GroupId, GroupChannels.SchemeAdmin @@ -1357,7 +1357,7 @@ func (s *SqlGroupStore) getGroupsAssociatedToChannelsByTeam(teamID string, opts WHERE GroupChannels.DeleteAt = 0 AND Channels.DeleteAt = 0 - AND Channels.TeamId = ?) AS gc ON gc.GroupId = ug.Id`, teamID). + AND Channels.TeamId = ?) AS gc ON gc.GroupId = UserGroups.Id`, teamID). LeftJoin(`( SELECT GroupMembers.GroupId, COUNT(*) AS MemberCount @@ -1369,13 +1369,13 @@ func (s *SqlGroupStore) getGroupsAssociatedToChannelsByTeam(teamID string, opts GroupMembers.DeleteAt = 0 AND Users.DeleteAt = 0 GROUP BY GroupId) AS Members - ON Members.GroupId = ug.Id`). - Where("ug.DeleteAt = 0 AND gc.DeleteAt = 0"). - OrderBy("ug.DisplayName") + ON Members.GroupId = UserGroups.Id`). + Where("UserGroups.DeleteAt = 0 AND gc.DeleteAt = 0"). + OrderBy("UserGroups.DisplayName") } if opts.FilterAllowReference { - query = query.Where("ug.AllowReference = true") + query = query.Where("UserGroups.AllowReference = true") } if opts.Q != "" { @@ -1384,7 +1384,7 @@ func (s *SqlGroupStore) getGroupsAssociatedToChannelsByTeam(teamID string, opts if s.DriverName() == model.DatabaseDriverMysql { operatorKeyword = "LIKE" } - query = query.Where(fmt.Sprintf("(ug.Name %[1]s ? OR ug.DisplayName %[1]s ?)", operatorKeyword), pattern, pattern) + query = query.Where(fmt.Sprintf("(UserGroups.Name %[1]s ? OR UserGroups.DisplayName %[1]s ?)", operatorKeyword), pattern, pattern) } return query @@ -1406,7 +1406,7 @@ func (s *SqlGroupStore) GetGroupsByTeam(teamId string, opts model.GroupSearchOpt if opts.PageOpts != nil { offset := uint64(opts.PageOpts.Page * opts.PageOpts.PerPage) - builder = builder.OrderBy("ug.DisplayName").Limit(uint64(opts.PageOpts.PerPage)).Offset(offset) + builder = builder.OrderBy("UserGroups.DisplayName").Limit(uint64(opts.PageOpts.PerPage)).Offset(offset) } groups := groupsWithSchemeAdmin{} @@ -1422,7 +1422,7 @@ func (s *SqlGroupStore) GetGroupsAssociatedToChannelsByTeam(teamId string, opts if opts.PageOpts != nil { offset := uint64(opts.PageOpts.Page * opts.PageOpts.PerPage) - builder = builder.OrderBy("ug.DisplayName").Limit(uint64(opts.PageOpts.PerPage)).Offset(offset) + builder = builder.OrderBy("UserGroups.DisplayName").Limit(uint64(opts.PageOpts.PerPage)).Offset(offset) } tgroups := groupsAssociatedToChannelWithSchemeAdmin{} diff --git a/server/channels/store/storetest/group_store.go b/server/channels/store/storetest/group_store.go index 2e24371733..a69948bc7e 100644 --- a/server/channels/store/storetest/group_store.go +++ b/server/channels/store/storetest/group_store.go @@ -29,7 +29,7 @@ func TestGroupStore(t *testing.T, rctx request.CTX, ss store.Store) { t.Run("GetByName", func(t *testing.T) { testGroupStoreGetByName(t, rctx, ss) }) t.Run("GetByIDs", func(t *testing.T) { testGroupStoreGetByIDs(t, rctx, ss) }) t.Run("GetByRemoteID", func(t *testing.T) { testGroupStoreGetByRemoteID(t, rctx, ss) }) - t.Run("GetAllBySource", func(t *testing.T) { testGroupStoreGetAllByType(t, rctx, ss) }) + t.Run("GetAllBySource", func(t *testing.T) { testGroupAllBySource(t, rctx, ss) }) t.Run("GetByUser", func(t *testing.T) { testGroupStoreGetByUser(t, rctx, ss) }) t.Run("Update", func(t *testing.T) { testGroupStoreUpdate(t, rctx, ss) }) t.Run("Delete", func(t *testing.T) { testGroupStoreDelete(t, rctx, ss) }) @@ -95,6 +95,7 @@ func TestGroupStore(t *testing.T, rctx request.CTX, ss store.Store) { t.Run("GetNonMemberUsersPage", func(t *testing.T) { groupTestGetNonMemberUsersPage(t, rctx, ss) }) t.Run("DistinctGroupMemberCountForSource", func(t *testing.T) { groupTestDistinctGroupMemberCountForSource(t, rctx, ss) }) + t.Run("GroupCountBySource", func(t *testing.T) { groupTestGroupCountBySource(t, rctx, ss) }) } func testGroupStoreCreate(t *testing.T, rctx request.CTX, ss store.Store) { @@ -188,7 +189,7 @@ func testGroupStoreCreate(t *testing.T, rctx request.CTX, ss store.Store) { } require.Equal(t, g6.IsValidForCreate().Id, "model.group.source.app_error") - //must use valid characters + // must use valid characters g7 := &model.Group{ Name: model.NewPointer("%^#@$$"), DisplayName: model.NewId(), @@ -333,7 +334,7 @@ func testGroupCreateWithUserIds(t *testing.T, rctx request.CTX, ss store.Store) } require.Equal(t, guids6.IsValidForCreate().Id, "model.group.source.app_error") - //must use valid characters + // must use valid characters g7 := &model.Group{ Name: model.NewPointer("%^#@$$"), DisplayName: model.NewId(), @@ -497,38 +498,85 @@ func testGroupStoreGetByRemoteID(t *testing.T, rctx request.CTX, ss store.Store) require.True(t, errors.As(err, &nfErr)) } -func testGroupStoreGetAllByType(t *testing.T, rctx request.CTX, ss store.Store) { - numGroups := 10 - - groups := []*model.Group{} - - // Create groups - for i := 0; i < numGroups; i++ { - g := &model.Group{ - Name: model.NewPointer(model.NewId()), - DisplayName: model.NewId(), - Description: model.NewId(), - Source: model.GroupSourceLdap, - RemoteId: model.NewPointer(model.NewId()), - } - groups = append(groups, g) - _, err := ss.Group().Create(g) - require.NoError(t, err) +func testGroupAllBySource(t *testing.T, rctx request.CTX, ss store.Store) { + // Create groups with different sources + g1 := &model.Group{ + Name: model.NewPointer(model.NewId()), + DisplayName: model.NewId(), + Description: model.NewId(), + Source: model.GroupSourceCustom, + RemoteId: model.NewPointer(model.NewId()), } - - // Returns all the groups - d1, err := ss.Group().GetAllBySource(model.GroupSourceLdap) + customGroup, err := ss.Group().Create(g1) require.NoError(t, err) - require.Condition(t, func() bool { return len(d1) >= numGroups }, len(d1), ">=", numGroups) - for _, expectedGroup := range groups { - present := false - for _, dbGroup := range d1 { - if dbGroup.Id == expectedGroup.Id { - present = true - break - } + defer ss.Group().Delete(customGroup.Id) + + g2 := &model.Group{ + Name: model.NewPointer(model.NewId()), + DisplayName: model.NewId(), + Description: model.NewId(), + Source: model.GroupSourceLdap, + RemoteId: model.NewPointer(model.NewId()), + } + ldapGroup, err := ss.Group().Create(g2) + require.NoError(t, err) + defer ss.Group().Delete(ldapGroup.Id) + + g3 := &model.Group{ + Name: model.NewPointer(model.NewId()), + DisplayName: model.NewId(), + Description: model.NewId(), + Source: model.GroupSourceLdap, + RemoteId: model.NewPointer(model.NewId()), + } + ldapGroup2, err := ss.Group().Create(g3) + require.NoError(t, err) + defer ss.Group().Delete(ldapGroup2.Id) + + // Test filtering by LDAP source + ldapGroups, err := ss.Group().GetAllBySource(model.GroupSourceLdap) + require.NoError(t, err) + + // Verify we got at least the 2 LDAP groups we created + found1, found2 := false, false + for _, group := range ldapGroups { + if group.Id == ldapGroup.Id { + found1 = true } - require.True(t, present) + if group.Id == ldapGroup2.Id { + found2 = true + } + // Make sure all returned groups are LDAP source + require.Equal(t, model.GroupSourceLdap, group.Source) + } + require.True(t, found1, "Failed to find the first LDAP group") + require.True(t, found2, "Failed to find the second LDAP group") + + // Test filtering by Custom source + customGroups, err := ss.Group().GetAllBySource(model.GroupSourceCustom) + require.NoError(t, err) + + // Verify we got at least the custom group we created + foundCustom := false + for _, group := range customGroups { + if group.Id == customGroup.Id { + foundCustom = true + } + // Make sure all returned groups are Custom source + require.Equal(t, model.GroupSourceCustom, group.Source) + } + require.True(t, foundCustom, "Failed to find the custom group") + + // Test with deleted group to ensure it's not returned + _, err = ss.Group().Delete(ldapGroup2.Id) + require.NoError(t, err) + + ldapGroupsAfterDelete, err := ss.Group().GetAllBySource(model.GroupSourceLdap) + require.NoError(t, err) + + // Verify the deleted group is not returned + for _, group := range ldapGroupsAfterDelete { + require.NotEqual(t, ldapGroup2.Id, group.Id, "Deleted group should not be returned") } } @@ -1571,6 +1619,11 @@ func testGetGroupSyncable(t *testing.T, rctx request.CTX, ss store.Store) { } func testGetAllGroupSyncablesByGroup(t *testing.T, rctx request.CTX, ss store.Store) { + t.Run("team", func(t *testing.T) { testGetAllGroupSyncablesByGroupTeam(t, rctx, ss) }) + t.Run("channel", func(t *testing.T) { testGetAllGroupSyncablesByGroupChannel(t, rctx, ss) }) +} + +func testGetAllGroupSyncablesByGroupTeam(t *testing.T, rctx request.CTX, ss store.Store) { numGroupSyncables := 10 // Create group @@ -1615,7 +1668,7 @@ func testGetAllGroupSyncablesByGroup(t *testing.T, rctx request.CTX, ss store.St // Returns all the group teams d1, err := ss.Group().GetAllGroupSyncablesByGroupId(group.Id, model.GroupSyncableTypeTeam) require.NoError(t, err) - require.Condition(t, func() bool { return len(d1) >= numGroupSyncables }, len(d1), ">=", numGroupSyncables) + require.Len(t, d1, numGroupSyncables) for _, expectedGroupTeam := range groupTeams { present := false for _, dbGroupTeam := range d1 { @@ -1629,6 +1682,73 @@ func testGetAllGroupSyncablesByGroup(t *testing.T, rctx request.CTX, ss store.St } } +func testGetAllGroupSyncablesByGroupChannel(t *testing.T, rctx request.CTX, ss store.Store) { + // Create Team + team := &model.Team{ + DisplayName: "Name", + Description: "Some description", + CompanyName: "Some company name", + AllowOpenInvite: false, + InviteId: "inviteid0", + Name: "z-z-" + model.NewId() + "a", + Email: "success+" + model.NewId() + "@simulator.amazonses.com", + Type: model.TeamOpen, + } + team, nErr := ss.Team().Save(team) + require.NoError(t, nErr) + + numGroupSyncables := 10 + + // Create group + g := &model.Group{ + Name: model.NewPointer(model.NewId()), + DisplayName: model.NewId(), + Description: model.NewId(), + Source: model.GroupSourceLdap, + RemoteId: model.NewPointer(model.NewId()), + } + group, err := ss.Group().Create(g) + require.NoError(t, err) + + groupChannels := []*model.GroupSyncable{} + + // Create groupChannels + for i := 0; i < numGroupSyncables; i++ { + // Create Channel + channel := &model.Channel{ + TeamId: team.Id, + DisplayName: "A Name", + Name: model.NewId(), + Type: model.ChannelTypePrivate, + } + channel, nErr = ss.Channel().Save(rctx, channel, 9999) + require.NoError(t, nErr) + + // Create groupChannel + groupChannel := model.NewGroupChannel(group.Id, channel.Id, false) + groupChannel.SchemeAdmin = true + groupChannel, err = ss.Group().CreateGroupSyncable(groupChannel) + require.NoError(t, err) + groupChannels = append(groupChannels, groupChannel) + } + + // Returns all the group channels + groupSyncables, err := ss.Group().GetAllGroupSyncablesByGroupId(group.Id, model.GroupSyncableTypeChannel) + require.NoError(t, err) + require.Len(t, groupSyncables, numGroupSyncables) + for _, expectedGroupChannel := range groupChannels { + present := false + for _, dbGroupChannel := range groupSyncables { + if dbGroupChannel.GroupId == expectedGroupChannel.GroupId && dbGroupChannel.SyncableId == expectedGroupChannel.SyncableId { + require.True(t, dbGroupChannel.SchemeAdmin) + present = true + break + } + } + require.True(t, present) + } +} + func testUpdateGroupSyncable(t *testing.T, rctx request.CTX, ss store.Store) { // Create Group g1 := &model.Group{ @@ -3621,6 +3741,10 @@ func testGetGroups(t *testing.T, rctx request.CTX, ss store.Store) { u3 := &model.User{ Email: MakeEmail(), Username: model.NewUsername(), + Timezone: model.StringMap{ + "useAutomaticTimezone": "false", + "manualTimezone": "UTC", + }, } user3, err := ss.User().Save(rctx, u3) require.NoError(t, err) @@ -3648,6 +3772,30 @@ func testGetGroups(t *testing.T, rctx request.CTX, ss store.Store) { _, err = ss.Channel().SaveMember(rctx, &m1) require.NoError(t, err) + m2 := model.ChannelMember{ + ChannelId: channel1.Id, + UserId: user2.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + _, err = ss.Channel().SaveMember(rctx, &m2) + require.NoError(t, err) + + m3 := model.ChannelMember{ + ChannelId: channel2.Id, + UserId: user2.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + _, err = ss.Channel().SaveMember(rctx, &m3) + require.NoError(t, err) + + m4 := model.ChannelMember{ + ChannelId: channel2.Id, + UserId: user3.Id, + NotifyProps: model.GetDefaultChannelNotifyProps(), + } + _, err = ss.Channel().SaveMember(rctx, &m4) + require.NoError(t, err) + user2.DeleteAt = 1 u2Update, _ := ss.User().Update(rctx, user2, true) @@ -4019,6 +4167,145 @@ func testGetGroups(t *testing.T, rctx request.CTX, ss store.Store) { }, Restrictions: nil, }, + { + Name: "Include channel1 member count", + Opts: model.GroupSearchOpts{IncludeChannelMemberCount: channel1.Id}, + Page: 0, + PerPage: 100, + Resultf: func(groups []*model.Group) bool { + for _, group := range groups { + fmt.Println(group.Id, group.ChannelMemberCount) + var channelMemberCount int + if group.ChannelMemberCount != nil { + channelMemberCount = *group.ChannelMemberCount + } + if group.Id == group1.Id && channelMemberCount != 2 { + fmt.Println("group1", group.Id, channelMemberCount) + return false + } + if group.Id == group2.Id && channelMemberCount != 1 { + fmt.Println("group2", group.Id, channelMemberCount) + return false + } + } + + return true + }, + Restrictions: nil, + }, + { + Name: "Include channel2 member count", + Opts: model.GroupSearchOpts{IncludeChannelMemberCount: channel2.Id}, + Page: 0, + PerPage: 100, + Resultf: func(groups []*model.Group) bool { + for _, group := range groups { + var channelMemberCount int + if group.ChannelMemberCount != nil { + channelMemberCount = *group.ChannelMemberCount + } + if group.Id == group1.Id && channelMemberCount != 1 { + fmt.Println("group1", group.Id, channelMemberCount) + return false + } + if group.Id == group2.Id && channelMemberCount != 2 { + fmt.Println("group2", group.Id, channelMemberCount) + return false + } + } + + return true + }, + Restrictions: nil, + }, + { + Name: "Include channel member count for non-existent channel", + Page: 0, + PerPage: 100, + Opts: model.GroupSearchOpts{IncludeChannelMemberCount: model.NewId()}, + Resultf: func(groups []*model.Group) bool { + for _, group := range groups { + var channelMemberCount int + if group.ChannelMemberCount != nil { + channelMemberCount = *group.ChannelMemberCount + } + + if channelMemberCount != 0 { + return false + } + } + + return true + }, + Restrictions: nil, + }, + { + Name: "Include channel1 member count, with timezones", + Opts: model.GroupSearchOpts{IncludeChannelMemberCount: channel1.Id, IncludeTimezones: true}, + Page: 0, + PerPage: 100, + Resultf: func(groups []*model.Group) bool { + for _, group := range groups { + var channelMemberTimezonesCount int + if group.ChannelMemberTimezonesCount != nil { + channelMemberTimezonesCount = *group.ChannelMemberTimezonesCount + } + if group.Id == group1.Id && channelMemberTimezonesCount != 0 { + return false + } + if group.Id == group2.Id && channelMemberTimezonesCount != 0 { + return false + } + } + + return true + }, + Restrictions: nil, + }, + { + Name: "Include channel2 member count, with timezones", + Opts: model.GroupSearchOpts{IncludeChannelMemberCount: channel2.Id, IncludeTimezones: true}, + Page: 0, + PerPage: 100, + Resultf: func(groups []*model.Group) bool { + for _, group := range groups { + var channelMemberTimezonesCount int + if group.ChannelMemberTimezonesCount != nil { + channelMemberTimezonesCount = *group.ChannelMemberTimezonesCount + } + if group.Id == group1.Id && channelMemberTimezonesCount != 0 { + return false + } + if group.Id == group2.Id && channelMemberTimezonesCount != 1 { + return false + } + } + + return true + }, + Restrictions: nil, + }, + { + Name: "Include channel member count for non-existent channel, with timezones", + Page: 0, + PerPage: 100, + Opts: model.GroupSearchOpts{IncludeChannelMemberCount: model.NewId(), IncludeTimezones: true}, + Resultf: func(groups []*model.Group) bool { + for _, group := range groups { + var channelMemberTimezonesCount int + if group.ChannelMemberTimezonesCount != nil { + channelMemberTimezonesCount = *group.ChannelMemberCount + } + + if channelMemberTimezonesCount != 0 { + return false + } + } + + return true + }, + Restrictions: nil, + }, } for _, tc := range testCases { @@ -5199,6 +5486,9 @@ func groupTestGroupMemberCount(t *testing.T, rctx request.CTX, ss store.Store) { } func groupTestDistinctGroupMemberCount(t *testing.T, rctx request.CTX, ss store.Store) { + ss.DropAllTables() + + // Create two groups group1, err := ss.Group().Create(&model.Group{ Name: model.NewPointer(model.NewId()), DisplayName: model.NewId(), @@ -5217,43 +5507,48 @@ func groupTestDistinctGroupMemberCount(t *testing.T, rctx request.CTX, ss store. require.NoError(t, err) defer ss.Group().Delete(group2.Id) - user := &model.User{ - Email: fmt.Sprintf("test.%s@localhost", model.NewId()), + // Create two users + user1 := &model.User{ + Email: MakeEmail(), Username: model.NewUsername(), } - user, err = ss.User().Save(rctx, user) + user1, err = ss.User().Save(rctx, user1) require.NoError(t, err) user2 := &model.User{ - Email: fmt.Sprintf("test.%s@localhost", model.NewId()), + Email: MakeEmail(), Username: model.NewUsername(), } user2, err = ss.User().Save(rctx, user2) require.NoError(t, err) - member1, err := ss.Group().UpsertMember(group1.Id, user.Id) + // Add user1 to group1 + member1, err := ss.Group().UpsertMember(group1.Id, user1.Id) require.NoError(t, err) defer ss.Group().DeleteMember(group1.Id, member1.UserId) - count, err := ss.Group().GroupMemberCount() + // Verify count is now 1 + count, err := ss.Group().DistinctGroupMemberCount() require.NoError(t, err) - require.GreaterOrEqual(t, count, int64(1)) + require.Equal(t, int64(1), count) - member2, err := ss.Group().UpsertMember(group1.Id, user2.Id) + // Add user2 to group1 + _, err = ss.Group().UpsertMember(group1.Id, user2.Id) require.NoError(t, err) - defer ss.Group().DeleteMember(group1.Id, member2.UserId) - countAfter1, err := ss.Group().GroupMemberCount() + // Verify count is now 2 + countAfter1, err := ss.Group().DistinctGroupMemberCount() require.NoError(t, err) - require.GreaterOrEqual(t, countAfter1, count+1) + require.Equal(t, int64(2), countAfter1) - member3, err := ss.Group().UpsertMember(group1.Id, member1.UserId) + // Add user1 to group2 as well + _, err = ss.Group().UpsertMember(group2.Id, user1.Id) require.NoError(t, err) - defer ss.Group().DeleteMember(group1.Id, member3.UserId) - countAfter2, err := ss.Group().GroupMemberCount() + // Verify count stays at 2 (user1 is already counted) + countAfter2, err := ss.Group().DistinctGroupMemberCount() require.NoError(t, err) - require.GreaterOrEqual(t, countAfter2, countAfter1) + require.Equal(t, countAfter1, countAfter2) } func groupTestGroupCountWithAllowReference(t *testing.T, rctx request.CTX, ss store.Store) { @@ -5445,3 +5740,63 @@ func groupTestDistinctGroupMemberCountForSource(t *testing.T, rctx request.CTX, require.NoError(t, err) require.Equal(t, ldapGroupCountBefore+1, ldapGroupCount) } + +func groupTestGroupCountBySource(t *testing.T, rctx request.CTX, ss store.Store) { + // Get initial counts for different sources + customSourceCountBefore, err := ss.Group().GroupCountBySource(model.GroupSourceCustom) + require.NoError(t, err) + ldapSourceCountBefore, err := ss.Group().GroupCountBySource(model.GroupSourceLdap) + require.NoError(t, err) + + // Create groups with different sources + g1 := &model.Group{ + Name: model.NewPointer(model.NewId()), + DisplayName: model.NewId(), + Description: model.NewId(), + Source: model.GroupSourceCustom, + RemoteId: model.NewPointer(model.NewId()), + } + customGroup, err := ss.Group().Create(g1) + require.NoError(t, err) + + g2 := &model.Group{ + Name: model.NewPointer(model.NewId()), + DisplayName: model.NewId(), + Description: model.NewId(), + Source: model.GroupSourceLdap, + RemoteId: model.NewPointer(model.NewId()), + } + ldapGroup, err := ss.Group().Create(g2) + require.NoError(t, err) + + g3 := &model.Group{ + Name: model.NewPointer(model.NewId()), + DisplayName: model.NewId(), + Description: model.NewId(), + Source: model.GroupSourceLdap, + RemoteId: model.NewPointer(model.NewId()), + } + ldapGroup2, err := ss.Group().Create(g3) + require.NoError(t, err) + + defer func() { + ss.Group().Delete(customGroup.Id) + ss.Group().Delete(ldapGroup.Id) + ss.Group().Delete(ldapGroup2.Id) + }() + + // Check counts after creating groups + customSourceCountAfter, err := ss.Group().GroupCountBySource(model.GroupSourceCustom) + require.NoError(t, err) + require.Equal(t, customSourceCountBefore+1, customSourceCountAfter) + + ldapSourceCountAfter, err := ss.Group().GroupCountBySource(model.GroupSourceLdap) + require.NoError(t, err) + require.Equal(t, ldapSourceCountBefore+2, ldapSourceCountAfter) + + // Delete one LDAP group and verify count decreases + ss.Group().Delete(ldapGroup.Id) + ldapSourceCountAfterDelete, err := ss.Group().GroupCountBySource(model.GroupSourceLdap) + require.NoError(t, err) + require.Equal(t, ldapSourceCountAfter-1, ldapSourceCountAfterDelete) +}