From f3eee28f569ebdb58f5a79515bb74a6517ea6a0c Mon Sep 17 00:00:00 2001 From: Ben Schumacher Date: Thu, 3 Oct 2024 16:26:53 +0200 Subject: [PATCH] [MM-60253] Avoid unnecessary cache clearing during LDAP sync (#28300) --- server/channels/app/channel.go | 20 ++- server/channels/app/syncables.go | 45 ++++-- server/channels/app/team.go | 8 +- server/channels/app/user.go | 4 +- .../opentracinglayer/opentracinglayer.go | 12 +- .../channels/store/retrylayer/retrylayer.go | 20 +-- .../channels/store/sqlstore/channel_store.go | 95 ++++++++--- server/channels/store/sqlstore/team_store.go | 66 +++++++- server/channels/store/store.go | 6 +- .../channels/store/storetest/channel_store.go | 14 ++ .../channels/store/storetest/group_store.go | 149 ++++++++++++------ .../store/storetest/mocks/ChannelStore.go | 22 ++- .../store/storetest/mocks/TeamStore.go | 28 +++- .../channels/store/timerlayer/timerlayer.go | 12 +- 14 files changed, 360 insertions(+), 141 deletions(-) diff --git a/server/channels/app/channel.go b/server/channels/app/channel.go index 3a763e7633..8664c33c2e 100644 --- a/server/channels/app/channel.go +++ b/server/channels/app/channel.go @@ -1298,7 +1298,7 @@ func (a *App) UpdateChannelMemberNotifyProps(c request.CTX, data map[string]stri a.invalidateCacheForChannelMembersNotifyProps(member.ChannelId) // Notify the clients that the member notify props changed - err = a.sendUpdateChannelMemberNotifyPropsEvent(member) + err = a.sendUpdateChannelMemberEvent(member) if err != nil { return nil, model.NewAppError("UpdateChannelMemberNotifyProps", "api.marshal_error", nil, "", http.StatusInternalServerError).Wrap(err) } @@ -1339,7 +1339,7 @@ func (a *App) PatchChannelMembersNotifyProps(c request.CTX, members []*model.Cha // Notify clients that their notify props have changed for _, member := range updated { - err := a.sendUpdateChannelMemberNotifyPropsEvent(member) + err := a.sendUpdateChannelMemberEvent(member) if err != nil { c.Logger().Warn("Failed to send WebSocket event for updated channel member notify props", mlog.Err(err)) } @@ -1348,7 +1348,7 @@ func (a *App) PatchChannelMembersNotifyProps(c request.CTX, members []*model.Cha return updated, nil } -func (a *App) sendUpdateChannelMemberNotifyPropsEvent(member *model.ChannelMember) error { +func (a *App) sendUpdateChannelMemberEvent(member *model.ChannelMember) error { evt := model.NewWebSocketEvent(model.WebsocketEventChannelMemberUpdated, "", "", member.UserId, nil, "") memberJSON, jsonErr := json.Marshal(member) if jsonErr != nil { @@ -1962,9 +1962,11 @@ func (a *App) GetAllChannels(c request.CTX, page, perPage int, opts model.Channe opts.ExcludeChannelNames = a.DefaultChannelNames(c) } storeOpts := store.ChannelSearchOpts{ - ExcludeChannelNames: opts.ExcludeChannelNames, NotAssociatedToGroup: opts.NotAssociatedToGroup, IncludeDeleted: opts.IncludeDeleted, + ExcludeChannelNames: opts.ExcludeChannelNames, + GroupConstrained: opts.GroupConstrained, + ExcludeGroupConstrained: opts.ExcludeGroupConstrained, ExcludePolicyConstrained: opts.ExcludePolicyConstrained, IncludePolicyID: opts.IncludePolicyID, } @@ -1981,9 +1983,13 @@ func (a *App) GetAllChannelsCount(c request.CTX, opts model.ChannelSearchOpts) ( opts.ExcludeChannelNames = a.DefaultChannelNames(c) } storeOpts := store.ChannelSearchOpts{ - ExcludeChannelNames: opts.ExcludeChannelNames, - NotAssociatedToGroup: opts.NotAssociatedToGroup, - IncludeDeleted: opts.IncludeDeleted, + NotAssociatedToGroup: opts.NotAssociatedToGroup, + IncludeDeleted: opts.IncludeDeleted, + ExcludeChannelNames: opts.ExcludeChannelNames, + GroupConstrained: opts.GroupConstrained, + ExcludeGroupConstrained: opts.ExcludeGroupConstrained, + ExcludePolicyConstrained: opts.ExcludePolicyConstrained, + IncludePolicyID: opts.IncludePolicyID, } count, err := a.Srv().Store().Channel().GetAllChannelsCount(storeOpts) if err != nil { diff --git a/server/channels/app/syncables.go b/server/channels/app/syncables.go index 45d293cd98..d4c3cdeaa5 100644 --- a/server/channels/app/syncables.go +++ b/server/channels/app/syncables.go @@ -234,26 +234,47 @@ func (a *App) SyncSyncableRoles(rctx request.CTX, syncableID string, syncableTyp switch syncableType { case model.GroupSyncableTypeTeam: - nErr := a.Srv().Store().Team().UpdateMembersRole(syncableID, permittedAdmins) - if nErr != nil { - return model.NewAppError("App.SyncSyncableRoles", "app.update_error", nil, "", http.StatusInternalServerError).Wrap(nErr) + var updatedMembers []*model.TeamMember + updatedMembers, err = a.Srv().Store().Team().UpdateMembersRole(syncableID, permittedAdmins) + if err != nil { + return model.NewAppError("App.SyncSyncableRoles", "app.update_error", nil, "", http.StatusInternalServerError).Wrap(err) + } + + for _, member := range updatedMembers { + a.ClearSessionCacheForUser(member.UserId) + + if appErr := a.sendUpdatedTeamMemberEvent(member); appErr != nil { + rctx.Logger().Warn("Error sending channel member updated websocket event", mlog.Err(appErr)) + } } - return nil case model.GroupSyncableTypeChannel: - nErr := a.Srv().Store().Channel().UpdateMembersRole(syncableID, permittedAdmins) - if nErr != nil { - return model.NewAppError("App.SyncSyncableRoles", "app.update_error", nil, "", http.StatusInternalServerError).Wrap(nErr) + var updatedMembers []*model.ChannelMember + updatedMembers, err = a.Srv().Store().Channel().UpdateMembersRole(syncableID, permittedAdmins) + if err != nil { + return model.NewAppError("App.SyncSyncableRoles", "app.update_error", nil, "", http.StatusInternalServerError).Wrap(err) + } + + for _, member := range updatedMembers { + a.ClearSessionCacheForUser(member.UserId) + + if appErr := a.sendUpdateChannelMemberEvent(member); appErr != nil { + rctx.Logger().Warn("Error sending channel member updated websocket event", mlog.Err(appErr)) + } } - return nil default: return model.NewAppError("App.SyncSyncableRoles", "groups.unsupported_syncable_type", map[string]any{"Value": syncableType}, "", http.StatusInternalServerError) } + + return nil } // SyncRolesAndMembership updates the SchemeAdmin status and membership of all of the members of the given // syncable. func (a *App) SyncRolesAndMembership(rctx request.CTX, syncableID string, syncableType model.GroupSyncableType, includeRemovedMembers bool) { - a.SyncSyncableRoles(rctx, syncableID, syncableType) + appErr := a.SyncSyncableRoles(rctx, syncableID, syncableType) + if appErr != nil { + rctx.Logger().Warn("Error syncing syncable roles", mlog.Err(appErr)) + } lastJob, _ := a.Srv().Store().Job().GetNewestJobByStatusAndType(model.JobStatusSuccess, model.JobTypeLdapSync) var since int64 @@ -272,9 +293,6 @@ func (a *App) SyncRolesAndMembership(rctx request.CTX, syncableID string, syncab if err := a.deleteGroupConstrainedTeamMemberships(rctx, &syncableID); err != nil { rctx.Logger().Warn("Error deleting group constrained team memberships", mlog.Err(err)) } - if err := a.ClearTeamMembersCache(syncableID); err != nil { - rctx.Logger().Warn("Error clearing team members cache", mlog.Err(err)) - } case model.GroupSyncableTypeChannel: params.ScopedChannelID = &syncableID if err := a.createDefaultChannelMemberships(rctx, params); err != nil { @@ -283,8 +301,5 @@ func (a *App) SyncRolesAndMembership(rctx request.CTX, syncableID string, syncab if err := a.deleteGroupConstrainedChannelMemberships(rctx, &syncableID); err != nil { rctx.Logger().Warn("Error deleting group constrained team memberships", mlog.Err(err)) } - if err := a.ClearChannelMembersCache(rctx, syncableID); err != nil { - rctx.Logger().Warn("Error clearing channel members cache", mlog.Err(err)) - } } } diff --git a/server/channels/app/team.go b/server/channels/app/team.go index 415965fe6e..72ce428660 100644 --- a/server/channels/app/team.go +++ b/server/channels/app/team.go @@ -469,7 +469,7 @@ func (a *App) UpdateTeamMemberRoles(c request.CTX, teamID string, userID string, a.ClearSessionCacheForUser(userID) - if appErr := a.sendUpdatedMemberRoleEvent(userID, member); appErr != nil { + if appErr := a.sendUpdatedTeamMemberEvent(member); appErr != nil { return nil, appErr } @@ -512,15 +512,15 @@ func (a *App) UpdateTeamMemberSchemeRoles(c request.CTX, teamID string, userID s a.ClearSessionCacheForUser(userID) - if appErr := a.sendUpdatedMemberRoleEvent(userID, member); appErr != nil { + if appErr := a.sendUpdatedTeamMemberEvent(member); appErr != nil { return nil, appErr } return member, nil } -func (a *App) sendUpdatedMemberRoleEvent(userID string, member *model.TeamMember) *model.AppError { - message := model.NewWebSocketEvent(model.WebsocketEventMemberroleUpdated, "", "", userID, nil, "") +func (a *App) sendUpdatedTeamMemberEvent(member *model.TeamMember) *model.AppError { + message := model.NewWebSocketEvent(model.WebsocketEventMemberroleUpdated, "", "", member.UserId, nil, "") tmJSON, jsonErr := json.Marshal(member) if jsonErr != nil { return model.NewAppError("sendUpdatedMemberRoleEvent", "api.marshal_error", nil, "", http.StatusInternalServerError).Wrap(jsonErr) diff --git a/server/channels/app/user.go b/server/channels/app/user.go index 62f9d0004e..c17155e700 100644 --- a/server/channels/app/user.go +++ b/server/channels/app/user.go @@ -2443,7 +2443,7 @@ func (a *App) PromoteGuestToUser(c request.CTX, user *model.User, requestorId st } for _, member := range teamMembers { - a.sendUpdatedMemberRoleEvent(user.Id, member) + a.sendUpdatedTeamMemberEvent(member) channelMembers, appErr := a.GetChannelMembersForUser(c, member.TeamId, user.Id) if appErr != nil { @@ -2487,7 +2487,7 @@ func (a *App) DemoteUserToGuest(c request.CTX, user *model.User) *model.AppError } for _, member := range teamMembers { - a.sendUpdatedMemberRoleEvent(user.Id, member) + a.sendUpdatedTeamMemberEvent(member) channelMembers, appErr := a.GetChannelMembersForUser(c, member.TeamId, user.Id) if appErr != nil { diff --git a/server/channels/store/opentracinglayer/opentracinglayer.go b/server/channels/store/opentracinglayer/opentracinglayer.go index ec5ff4791e..e05d3ac0cb 100644 --- a/server/channels/store/opentracinglayer/opentracinglayer.go +++ b/server/channels/store/opentracinglayer/opentracinglayer.go @@ -2563,7 +2563,7 @@ func (s *OpenTracingLayerChannelStore) UpdateMemberNotifyProps(channelID string, return result, err } -func (s *OpenTracingLayerChannelStore) UpdateMembersRole(channelID string, userIDs []string) error { +func (s *OpenTracingLayerChannelStore) UpdateMembersRole(channelID string, userIDs []string) ([]*model.ChannelMember, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ChannelStore.UpdateMembersRole") s.Root.Store.SetContext(newCtx) @@ -2572,13 +2572,13 @@ func (s *OpenTracingLayerChannelStore) UpdateMembersRole(channelID string, userI }() defer span.Finish() - err := s.ChannelStore.UpdateMembersRole(channelID, userIDs) + result, err := s.ChannelStore.UpdateMembersRole(channelID, userIDs) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) } - return err + return result, err } func (s *OpenTracingLayerChannelStore) UpdateMultipleMembers(members []*model.ChannelMember) ([]*model.ChannelMember, error) { @@ -10554,7 +10554,7 @@ func (s *OpenTracingLayerTeamStore) UpdateMember(rctx request.CTX, member *model return result, err } -func (s *OpenTracingLayerTeamStore) UpdateMembersRole(teamID string, userIDs []string) error { +func (s *OpenTracingLayerTeamStore) UpdateMembersRole(teamID string, adminIDs []string) ([]*model.TeamMember, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "TeamStore.UpdateMembersRole") s.Root.Store.SetContext(newCtx) @@ -10563,13 +10563,13 @@ func (s *OpenTracingLayerTeamStore) UpdateMembersRole(teamID string, userIDs []s }() defer span.Finish() - err := s.TeamStore.UpdateMembersRole(teamID, userIDs) + result, err := s.TeamStore.UpdateMembersRole(teamID, adminIDs) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) } - return err + return result, err } func (s *OpenTracingLayerTeamStore) UpdateMultipleMembers(members []*model.TeamMember) ([]*model.TeamMember, error) { diff --git a/server/channels/store/retrylayer/retrylayer.go b/server/channels/store/retrylayer/retrylayer.go index e04b8c5a36..3db4b2d39e 100644 --- a/server/channels/store/retrylayer/retrylayer.go +++ b/server/channels/store/retrylayer/retrylayer.go @@ -2840,21 +2840,21 @@ func (s *RetryLayerChannelStore) UpdateMemberNotifyProps(channelID string, userI } -func (s *RetryLayerChannelStore) UpdateMembersRole(channelID string, userIDs []string) error { +func (s *RetryLayerChannelStore) UpdateMembersRole(channelID string, userIDs []string) ([]*model.ChannelMember, error) { tries := 0 for { - err := s.ChannelStore.UpdateMembersRole(channelID, userIDs) + result, err := s.ChannelStore.UpdateMembersRole(channelID, userIDs) if err == nil { - return nil + return result, nil } if !isRepeatableError(err) { - return err + return result, err } tries++ if tries >= 3 { err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") - return err + return result, err } timepkg.Sleep(100 * timepkg.Millisecond) } @@ -12071,21 +12071,21 @@ func (s *RetryLayerTeamStore) UpdateMember(rctx request.CTX, member *model.TeamM } -func (s *RetryLayerTeamStore) UpdateMembersRole(teamID string, userIDs []string) error { +func (s *RetryLayerTeamStore) UpdateMembersRole(teamID string, adminIDs []string) ([]*model.TeamMember, error) { tries := 0 for { - err := s.TeamStore.UpdateMembersRole(teamID, userIDs) + result, err := s.TeamStore.UpdateMembersRole(teamID, adminIDs) if err == nil { - return nil + return result, nil } if !isRepeatableError(err) { - return err + return result, err } tries++ if tries >= 3 { err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") - return err + return result, err } timepkg.Sleep(100 * timepkg.Millisecond) } diff --git a/server/channels/store/sqlstore/channel_store.go b/server/channels/store/sqlstore/channel_store.go index dce8d6ddcf..ec81005e38 100644 --- a/server/channels/store/sqlstore/channel_store.go +++ b/server/channels/store/sqlstore/channel_store.go @@ -7,6 +7,7 @@ import ( "context" "database/sql" "fmt" + "slices" "sort" "strconv" "strings" @@ -1190,6 +1191,15 @@ func (s SqlChannelStore) getAllChannelsQuery(opts store.ChannelSearchOpts, forCo query = query.Where("c.Id NOT IN (SELECT ChannelId FROM GroupChannels WHERE GroupChannels.GroupId = ? AND GroupChannels.DeleteAt = 0)", opts.NotAssociatedToGroup) } + if opts.GroupConstrained { + query = query.Where(sq.Eq{"c.GroupConstrained": true}) + } else if opts.ExcludeGroupConstrained { + query = query.Where(sq.Or{ + sq.NotEq{"c.GroupConstrained": true}, + sq.Eq{"c.GroupConstrained": nil}, + }) + } + if len(opts.ExcludeChannelNames) > 0 { query = query.Where(sq.NotEq{"c.Name": opts.ExcludeChannelNames}) } @@ -4161,27 +4171,76 @@ func (s SqlChannelStore) UserBelongsToChannels(userId string, channelIds []strin return c > 0, nil } -// TODO: parameterize userIDs -func (s SqlChannelStore) UpdateMembersRole(channelID string, userIDs []string) error { - sql := fmt.Sprintf(` - UPDATE - ChannelMembers - SET - SchemeAdmin = CASE WHEN UserId IN ('%s') THEN - TRUE - ELSE - FALSE - END - WHERE - ChannelId = ? - AND (SchemeGuest = false OR SchemeGuest IS NULL) - `, strings.Join(userIDs, "', '")) +// UpdateMembersRole updates all the members of channelID in the adminIDs string array to be admins and sets all other +// users as not being admin. +// It returns the list of userIDs whose roles got updated. +// +// TODO: parameterize adminIDs +func (s SqlChannelStore) UpdateMembersRole(channelID string, adminIDs []string) (_ []*model.ChannelMember, err error) { + transaction, err := s.GetMasterX().Beginx() + if err != nil { + return nil, err + } + defer finalizeTransactionX(transaction, &err) - if _, err := s.GetMasterX().Exec(sql, channelID); err != nil { - return errors.Wrap(err, "failed to update ChannelMembers") + // On MySQL it's not possible to update a table and select from it in the same query. + // A SELECT and a UPDATE query are needed. + // Once we only support PostgreSQL, this can be done in a single query using RETURNING. + query, args, err := s.getQueryBuilder(). + Select("*"). + From("ChannelMembers"). + Where(sq.Eq{"ChannelID": channelID}). + Where(sq.Or{sq.Eq{"SchemeGuest": false}, sq.Expr("SchemeGuest IS NULL")}). + Where( + sq.Or{ + // New admins + sq.And{ + sq.Eq{"SchemeAdmin": false}, + sq.Eq{"UserId": adminIDs}, + }, + // Demoted admins + sq.And{ + sq.Eq{"SchemeAdmin": true}, + sq.NotEq{"UserId": adminIDs}, + }, + }, + ).ToSql() + if err != nil { + return nil, errors.Wrap(err, "channel_tosql") } - return nil + var updatedMembers []*model.ChannelMember + if err = transaction.Select(&updatedMembers, query, args...); err != nil { + return nil, errors.Wrap(err, "failed to get list of updated users") + } + + // Update SchemeAdmin field as the data from the SQL is not updated yet + for _, member := range updatedMembers { + if slices.Contains(adminIDs, member.UserId) { + member.SchemeAdmin = true + } else { + member.SchemeAdmin = false + } + } + + query, args, err = s.getQueryBuilder(). + Update("ChannelMembers"). + Set("SchemeAdmin", sq.Case().When(sq.Eq{"UserId": adminIDs}, "true").Else("false")). + Where(sq.Eq{"ChannelId": channelID}). + Where(sq.Or{sq.Eq{"SchemeGuest": false}, sq.Expr("SchemeGuest IS NULL")}).ToSql() + if err != nil { + return nil, errors.Wrap(err, "team_tosql") + } + + if _, err = transaction.Exec(query, args...); err != nil { + return nil, errors.Wrap(err, "failed to update ChannelMembers") + } + + if err = transaction.Commit(); err != nil { + return nil, errors.Wrap(err, "commit_transaction") + } + + return updatedMembers, nil } func (s SqlChannelStore) GroupSyncedChannelCount() (int64, error) { diff --git a/server/channels/store/sqlstore/team_store.go b/server/channels/store/sqlstore/team_store.go index 92674434c1..f564fd1e62 100644 --- a/server/channels/store/sqlstore/team_store.go +++ b/server/channels/store/sqlstore/team_store.go @@ -6,6 +6,7 @@ package sqlstore import ( "database/sql" "fmt" + "slices" "strings" sq "github.com/mattermost/squirrel" @@ -1591,23 +1592,74 @@ func (s SqlTeamStore) UserBelongsToTeams(userId string, teamIds []string) (bool, return c > 0, nil } -// UpdateMembersRole updates all the members of teamID in the userIds string array to be admins and sets all other +// UpdateMembersRole updates all the members of teamID in the adminIDs string array to be admins and sets all other // users as not being admin. -func (s SqlTeamStore) UpdateMembersRole(teamID string, userIDs []string) error { +// It returns the list of userIDs whose roles got updated. +func (s SqlTeamStore) UpdateMembersRole(teamID string, adminIDs []string) (_ []*model.TeamMember, err error) { + transaction, err := s.GetMasterX().Beginx() + if err != nil { + return nil, err + } + defer finalizeTransactionX(transaction, &err) + + // On MySQL it's not possible to update a table and select from it in the same query. + // A SELECT and a UPDATE query are needed. + // Once we only support PostgreSQL, this can be done in a single query using RETURNING. query, args, err := s.getQueryBuilder(). + Select("*"). + From("TeamMembers"). + Where(sq.Eq{"TeamId": teamID, "DeleteAt": 0}). + Where(sq.Or{sq.Eq{"SchemeGuest": false}, sq.Expr("SchemeGuest IS NULL")}). + Where( + sq.Or{ + // New admins + sq.And{ + sq.Eq{"SchemeAdmin": false}, + sq.Eq{"UserId": adminIDs}, + }, + // Demoted admins + sq.And{ + sq.Eq{"SchemeAdmin": true}, + sq.NotEq{"UserId": adminIDs}, + }, + }, + ).ToSql() + if err != nil { + return nil, errors.Wrap(err, "team_tosql") + } + + var updatedMembers []*model.TeamMember + if err = transaction.Select(&updatedMembers, query, args...); err != nil { + return nil, errors.Wrap(err, "failed to get list of updated users") + } + + // Update SchemeAdmin field as the data from the SQL is not updated yet + for _, member := range updatedMembers { + if slices.Contains(adminIDs, member.UserId) { + member.SchemeAdmin = true + } else { + member.SchemeAdmin = false + } + } + + query, args, err = s.getQueryBuilder(). Update("TeamMembers"). - Set("SchemeAdmin", sq.Case().When(sq.Eq{"UserId": userIDs}, "true").Else("false")). + Set("SchemeAdmin", sq.Case().When(sq.Eq{"UserId": adminIDs}, "true").Else("false")). Where(sq.Eq{"TeamId": teamID, "DeleteAt": 0}). Where(sq.Or{sq.Eq{"SchemeGuest": false}, sq.Expr("SchemeGuest IS NULL")}).ToSql() if err != nil { - return errors.Wrap(err, "team_tosql") + return nil, errors.Wrap(err, "team_tosql") } - if _, err = s.GetMasterX().Exec(query, args...); err != nil { - return errors.Wrap(err, "failed to update TeamMembers") + if _, err = transaction.Exec(query, args...); err != nil { + return nil, errors.Wrap(err, "failed to update TeamMembers") } - return nil + if err = transaction.Commit(); err != nil { + return nil, errors.Wrap(err, "commit_transaction") + } + + return updatedMembers, nil } func applyTeamMemberViewRestrictionsFilter(query sq.SelectBuilder, restrictions *model.ViewUsersRestrictions) sq.SelectBuilder { diff --git a/server/channels/store/store.go b/server/channels/store/store.go index 443363efeb..0463e500fd 100644 --- a/server/channels/store/store.go +++ b/server/channels/store/store.go @@ -168,7 +168,8 @@ type TeamStore interface { // UpdateMembersRole sets all of the given team members to admins and all of the other members of the team to // non-admin members. - UpdateMembersRole(teamID string, userIDs []string) error + // It returns the list of userIDs whose roles got updated. + UpdateMembersRole(teamID string, adminIDs []string) ([]*model.TeamMember, error) // GroupSyncedTeamCount returns the count of non-deleted group-constrained teams. GroupSyncedTeamCount() (int64, error) @@ -300,7 +301,8 @@ type ChannelStore interface { // UpdateMembersRole sets all of the given team members to admins and all of the other members of the team to // non-admin members. - UpdateMembersRole(channelID string, userIDs []string) error + // It returns the list of userIDs whose roles got updated. + UpdateMembersRole(channelID string, userIDs []string) ([]*model.ChannelMember, error) // GroupSyncedChannelCount returns the count of non-deleted group-constrained channels. GroupSyncedChannelCount() (int64, error) diff --git a/server/channels/store/storetest/channel_store.go b/server/channels/store/storetest/channel_store.go index a9877e0eaa..5a226902ad 100644 --- a/server/channels/store/storetest/channel_store.go +++ b/server/channels/store/storetest/channel_store.go @@ -3854,6 +3854,7 @@ func testChannelStoreGetAllChannels(t *testing.T, rctx request.CTX, ss store.Sto c1.DisplayName = "Channel1" + model.NewId() c1.Name = NewTestId() c1.Type = model.ChannelTypeOpen + c1.GroupConstrained = model.NewPointer(true) _, nErr := ss.Channel().Save(rctx, &c1, -1) require.NoError(t, nErr) @@ -3939,6 +3940,19 @@ func testChannelStoreGetAllChannels(t *testing.T, rctx request.CTX, ss store.Sto list, nErr = ss.Channel().GetAllChannels(0, 10, store.ChannelSearchOpts{NotAssociatedToGroup: group.Id}) require.NoError(t, nErr) assert.Len(t, list, 1) + assert.Equal(t, c3.Id, list[0].Id) + + // GroupConstrained + list, nErr = ss.Channel().GetAllChannels(0, 10, store.ChannelSearchOpts{GroupConstrained: true}) + require.NoError(t, nErr) + require.Len(t, list, 1) + assert.Equal(t, c1.Id, list[0].Id) + + // ExcludeGroupConstrained + list, nErr = ss.Channel().GetAllChannels(0, 10, store.ChannelSearchOpts{ExcludeGroupConstrained: true}) + require.NoError(t, nErr) + require.Len(t, list, 1) + assert.Equal(t, c3.Id, list[0].Id) // Exclude channel names list, nErr = ss.Channel().GetAllChannels(0, 10, store.ChannelSearchOpts{ExcludeChannelNames: []string{c1.Name}}) diff --git a/server/channels/store/storetest/group_store.go b/server/channels/store/storetest/group_store.go index 8bb0d465f0..3e8a2f6712 100644 --- a/server/channels/store/storetest/group_store.go +++ b/server/channels/store/storetest/group_store.go @@ -81,7 +81,7 @@ func TestGroupStore(t *testing.T, rctx request.CTX, ss store.Store) { t.Run("AdminRoleGroupsForSyncableMember_Team", func(t *testing.T) { groupTestAdminRoleGroupsForSyncableMemberTeam(t, rctx, ss) }) t.Run("PermittedSyncableAdmins_Team", func(t *testing.T) { groupTestPermittedSyncableAdminsTeam(t, rctx, ss) }) t.Run("PermittedSyncableAdmins_Channel", func(t *testing.T) { groupTestPermittedSyncableAdminsChannel(t, rctx, ss) }) - t.Run("UpdateMembersRole_Team", func(t *testing.T) { groupTestpUpdateMembersRoleTeam(t, rctx, ss) }) + t.Run("UpdateMembersRole_Team", func(t *testing.T) { groupTestUpdateMembersRoleTeam(t, rctx, ss) }) t.Run("UpdateMembersRole_Channel", func(t *testing.T) { groupTestpUpdateMembersRoleChannel(t, rctx, ss) }) t.Run("GroupCount", func(t *testing.T) { groupTestGroupCount(t, rctx, ss) }) @@ -4739,7 +4739,7 @@ func groupTestPermittedSyncableAdminsChannel(t *testing.T, rctx request.CTX, ss require.ElementsMatch(t, []string{user3.Id}, actualUserIDs) } -func groupTestpUpdateMembersRoleTeam(t *testing.T, rctx request.CTX, ss store.Store) { +func groupTestUpdateMembersRoleTeam(t *testing.T, rctx request.CTX, ss store.Store) { team := &model.Team{ DisplayName: "Name", Description: "Some description", @@ -4759,6 +4759,7 @@ func groupTestpUpdateMembersRoleTeam(t *testing.T, rctx request.CTX, ss store.St } user1, err = ss.User().Save(rctx, user1) require.NoError(t, err) + t.Log("Created user1", user1.Id) user2 := &model.User{ Email: MakeEmail(), @@ -4766,6 +4767,7 @@ func groupTestpUpdateMembersRoleTeam(t *testing.T, rctx request.CTX, ss store.St } user2, err = ss.User().Save(rctx, user2) require.NoError(t, err) + t.Log("Created user2", user2.Id) user3 := &model.User{ Email: MakeEmail(), @@ -4773,6 +4775,7 @@ func groupTestpUpdateMembersRoleTeam(t *testing.T, rctx request.CTX, ss store.St } user3, err = ss.User().Save(rctx, user3) require.NoError(t, err) + t.Log("Created user3", user3.Id) user4 := &model.User{ Email: MakeEmail(), @@ -4780,6 +4783,7 @@ func groupTestpUpdateMembersRoleTeam(t *testing.T, rctx request.CTX, ss store.St } user4, err = ss.User().Save(rctx, user4) require.NoError(t, err) + t.Log("Created user4", user4.Id) for _, user := range []*model.User{user1, user2, user3} { _, nErr := ss.Team().SaveMember(rctx, &model.TeamMember{TeamId: team.Id, UserId: user.Id}, 9999) @@ -4790,53 +4794,73 @@ func groupTestpUpdateMembersRoleTeam(t *testing.T, rctx request.CTX, ss store.St require.NoError(t, nErr) tests := []struct { - testName string - inUserIDs []string - targetSchemeAdminValue bool + testName string + newAdmins []string + expectedUpdatedUsers []string }{ { - "Given users are admins", + "Two new admins", + []string{user1.Id, user2.Id}, []string{user1.Id, user2.Id}, - true, }, { - "Given users are members", + "Demote one admin", + []string{user1.Id}, []string{user2.Id}, - false, }, { - "Non-given users are admins", - []string{user2.Id}, - false, + "Operation is idempotent", + []string{user1.Id}, + nil, }, { - "Non-given users are members", - []string{user2.Id}, - false, + "Promote a team member", + []string{user1.Id, user3.Id}, + []string{user3.Id}, + }, + { + "Guests never get promoted", + []string{user1.Id, user3.Id, user4.Id}, + nil, }, } for _, tt := range tests { t.Run(tt.testName, func(t *testing.T) { - err = ss.Team().UpdateMembersRole(team.Id, tt.inUserIDs) + var updatedMembers []*model.TeamMember + updatedMembers, err = ss.Team().UpdateMembersRole(team.Id, tt.newAdmins) require.NoError(t, err) + var updatedUserIDs []string + for _, member := range updatedMembers { + assert.False(t, member.SchemeGuest, fmt.Sprintf("userID: %s", member.UserId)) + + if slices.Contains(tt.newAdmins, member.UserId) { + assert.True(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + } else { + assert.False(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + } + + updatedUserIDs = append(updatedUserIDs, member.UserId) + } + assert.ElementsMatch(t, tt.expectedUpdatedUsers, updatedUserIDs) + members, err := ss.Team().GetMembers(team.Id, 0, 100, nil) require.NoError(t, err) - require.GreaterOrEqual(t, len(members), 4) // sanity check for team membership + assert.GreaterOrEqual(t, len(members), 4) // sanity check for team membership for _, member := range members { - if slices.Contains(tt.inUserIDs, member.UserId) { - require.True(t, member.SchemeAdmin) - } else { - require.False(t, member.SchemeAdmin) - } - // Ensure guest account never changes. if member.UserId == user4.Id { - require.False(t, member.SchemeUser) - require.False(t, member.SchemeAdmin) - require.True(t, member.SchemeGuest) + assert.False(t, member.SchemeUser, fmt.Sprintf("userID: %s", member.UserId)) + assert.False(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + assert.True(t, member.SchemeGuest, fmt.Sprintf("userID: %s", member.UserId)) + } else { + if slices.Contains(tt.newAdmins, member.UserId) { + assert.True(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + } else { + assert.False(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + } } } }) @@ -4859,6 +4883,7 @@ func groupTestpUpdateMembersRoleChannel(t *testing.T, rctx request.CTX, ss store } user1, err = ss.User().Save(rctx, user1) require.NoError(t, err) + t.Log("Created user1", user1.Id) user2 := &model.User{ Email: MakeEmail(), @@ -4866,6 +4891,7 @@ func groupTestpUpdateMembersRoleChannel(t *testing.T, rctx request.CTX, ss store } user2, err = ss.User().Save(rctx, user2) require.NoError(t, err) + t.Log("Created user2", user2.Id) user3 := &model.User{ Email: MakeEmail(), @@ -4873,6 +4899,7 @@ func groupTestpUpdateMembersRoleChannel(t *testing.T, rctx request.CTX, ss store } user3, err = ss.User().Save(rctx, user3) require.NoError(t, err) + t.Log("Created user3", user3.Id) user4 := &model.User{ Email: MakeEmail(), @@ -4880,6 +4907,7 @@ func groupTestpUpdateMembersRoleChannel(t *testing.T, rctx request.CTX, ss store } user4, err = ss.User().Save(rctx, user4) require.NoError(t, err) + t.Log("Created user4", user4.Id) for _, user := range []*model.User{user1, user2, user3} { _, err = ss.Channel().SaveMember(rctx, &model.ChannelMember{ @@ -4899,54 +4927,73 @@ func groupTestpUpdateMembersRoleChannel(t *testing.T, rctx request.CTX, ss store require.NoError(t, err) tests := []struct { - testName string - inUserIDs []string - targetSchemeAdminValue bool + testName string + newAdmins []string + expectedUpdatedUsers []string }{ { - "Given users are admins", + "Two new admins", + []string{user1.Id, user2.Id}, []string{user1.Id, user2.Id}, - true, }, { - "Given users are members", + "Demote one admin", + []string{user1.Id}, []string{user2.Id}, - false, }, { - "Non-given users are admins", - []string{user2.Id}, - false, + "Operation is idempotent", + []string{user1.Id}, + nil, }, { - "Non-given users are members", - []string{user2.Id}, - false, + "Promote a team member", + []string{user1.Id, user3.Id}, + []string{user3.Id}, + }, + { + "Guests never get promoted", + []string{user1.Id, user3.Id, user4.Id}, + nil, }, } for _, tt := range tests { t.Run(tt.testName, func(t *testing.T) { - err = ss.Channel().UpdateMembersRole(channel.Id, tt.inUserIDs) + var updatedMemmbers []*model.ChannelMember + updatedMemmbers, err = ss.Channel().UpdateMembersRole(channel.Id, tt.newAdmins) require.NoError(t, err) + var updatedUserIDs []string + for _, member := range updatedMemmbers { + assert.False(t, member.SchemeGuest, fmt.Sprintf("userID: %s", member.UserId)) + + if slices.Contains(tt.newAdmins, member.UserId) { + assert.True(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + } else { + assert.False(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + } + + updatedUserIDs = append(updatedUserIDs, member.UserId) + } + assert.ElementsMatch(t, tt.expectedUpdatedUsers, updatedUserIDs) + members, err := ss.Channel().GetMembers(channel.Id, 0, 100) require.NoError(t, err) - - require.GreaterOrEqual(t, len(members), 4) // sanity check for channel membership + assert.GreaterOrEqual(t, len(members), 4) // sanity check for channel membership for _, member := range members { - if slices.Contains(tt.inUserIDs, member.UserId) { - require.True(t, member.SchemeAdmin) - } else { - require.False(t, member.SchemeAdmin) - } - // Ensure guest account never changes. if member.UserId == user4.Id { - require.False(t, member.SchemeUser) - require.False(t, member.SchemeAdmin) - require.True(t, member.SchemeGuest) + assert.False(t, member.SchemeUser, fmt.Sprintf("userID: %s", member.UserId)) + assert.False(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + assert.True(t, member.SchemeGuest, fmt.Sprintf("userID: %s", member.UserId)) + } else { + if slices.Contains(tt.newAdmins, member.UserId) { + assert.True(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + } else { + assert.False(t, member.SchemeAdmin, fmt.Sprintf("userID: %s", member.UserId)) + } } } }) diff --git a/server/channels/store/storetest/mocks/ChannelStore.go b/server/channels/store/storetest/mocks/ChannelStore.go index 2661df49ce..489151e9cb 100644 --- a/server/channels/store/storetest/mocks/ChannelStore.go +++ b/server/channels/store/storetest/mocks/ChannelStore.go @@ -2915,21 +2915,33 @@ func (_m *ChannelStore) UpdateMemberNotifyProps(channelID string, userID string, } // UpdateMembersRole provides a mock function with given fields: channelID, userIDs -func (_m *ChannelStore) UpdateMembersRole(channelID string, userIDs []string) error { +func (_m *ChannelStore) UpdateMembersRole(channelID string, userIDs []string) ([]*model.ChannelMember, error) { ret := _m.Called(channelID, userIDs) if len(ret) == 0 { panic("no return value specified for UpdateMembersRole") } - var r0 error - if rf, ok := ret.Get(0).(func(string, []string) error); ok { + var r0 []*model.ChannelMember + var r1 error + if rf, ok := ret.Get(0).(func(string, []string) ([]*model.ChannelMember, error)); ok { + return rf(channelID, userIDs) + } + if rf, ok := ret.Get(0).(func(string, []string) []*model.ChannelMember); ok { r0 = rf(channelID, userIDs) } else { - r0 = ret.Error(0) + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.ChannelMember) + } } - return r0 + if rf, ok := ret.Get(1).(func(string, []string) error); ok { + r1 = rf(channelID, userIDs) + } else { + r1 = ret.Error(1) + } + + return r0, r1 } // UpdateMultipleMembers provides a mock function with given fields: members diff --git a/server/channels/store/storetest/mocks/TeamStore.go b/server/channels/store/storetest/mocks/TeamStore.go index 2c23ba3077..60493056c6 100644 --- a/server/channels/store/storetest/mocks/TeamStore.go +++ b/server/channels/store/storetest/mocks/TeamStore.go @@ -1336,22 +1336,34 @@ func (_m *TeamStore) UpdateMember(rctx request.CTX, member *model.TeamMember) (* return r0, r1 } -// UpdateMembersRole provides a mock function with given fields: teamID, userIDs -func (_m *TeamStore) UpdateMembersRole(teamID string, userIDs []string) error { - ret := _m.Called(teamID, userIDs) +// UpdateMembersRole provides a mock function with given fields: teamID, adminIDs +func (_m *TeamStore) UpdateMembersRole(teamID string, adminIDs []string) ([]*model.TeamMember, error) { + ret := _m.Called(teamID, adminIDs) if len(ret) == 0 { panic("no return value specified for UpdateMembersRole") } - var r0 error - if rf, ok := ret.Get(0).(func(string, []string) error); ok { - r0 = rf(teamID, userIDs) + var r0 []*model.TeamMember + var r1 error + if rf, ok := ret.Get(0).(func(string, []string) ([]*model.TeamMember, error)); ok { + return rf(teamID, adminIDs) + } + if rf, ok := ret.Get(0).(func(string, []string) []*model.TeamMember); ok { + r0 = rf(teamID, adminIDs) } else { - r0 = ret.Error(0) + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.TeamMember) + } } - return r0 + if rf, ok := ret.Get(1).(func(string, []string) error); ok { + r1 = rf(teamID, adminIDs) + } else { + r1 = ret.Error(1) + } + + return r0, r1 } // UpdateMultipleMembers provides a mock function with given fields: members diff --git a/server/channels/store/timerlayer/timerlayer.go b/server/channels/store/timerlayer/timerlayer.go index ef4bf76962..61dde64b5b 100644 --- a/server/channels/store/timerlayer/timerlayer.go +++ b/server/channels/store/timerlayer/timerlayer.go @@ -2366,10 +2366,10 @@ func (s *TimerLayerChannelStore) UpdateMemberNotifyProps(channelID string, userI return result, err } -func (s *TimerLayerChannelStore) UpdateMembersRole(channelID string, userIDs []string) error { +func (s *TimerLayerChannelStore) UpdateMembersRole(channelID string, userIDs []string) ([]*model.ChannelMember, error) { start := time.Now() - err := s.ChannelStore.UpdateMembersRole(channelID, userIDs) + result, err := s.ChannelStore.UpdateMembersRole(channelID, userIDs) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -2379,7 +2379,7 @@ func (s *TimerLayerChannelStore) UpdateMembersRole(channelID string, userIDs []s } s.Root.Metrics.ObserveStoreMethodDuration("ChannelStore.UpdateMembersRole", success, elapsed) } - return err + return result, err } func (s *TimerLayerChannelStore) UpdateMultipleMembers(members []*model.ChannelMember) ([]*model.ChannelMember, error) { @@ -9495,10 +9495,10 @@ func (s *TimerLayerTeamStore) UpdateMember(rctx request.CTX, member *model.TeamM return result, err } -func (s *TimerLayerTeamStore) UpdateMembersRole(teamID string, userIDs []string) error { +func (s *TimerLayerTeamStore) UpdateMembersRole(teamID string, adminIDs []string) ([]*model.TeamMember, error) { start := time.Now() - err := s.TeamStore.UpdateMembersRole(teamID, userIDs) + result, err := s.TeamStore.UpdateMembersRole(teamID, adminIDs) elapsed := float64(time.Since(start)) / float64(time.Second) if s.Root.Metrics != nil { @@ -9508,7 +9508,7 @@ func (s *TimerLayerTeamStore) UpdateMembersRole(teamID string, userIDs []string) } s.Root.Metrics.ObserveStoreMethodDuration("TeamStore.UpdateMembersRole", success, elapsed) } - return err + return result, err } func (s *TimerLayerTeamStore) UpdateMultipleMembers(members []*model.TeamMember) ([]*model.TeamMember, error) {