From 12c50eb83018de88212b78c7e0dbc934dccfbd51 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jes=C3=BAs=20Espino?= Date: Mon, 15 Apr 2019 22:53:52 +0200 Subject: [PATCH] Initial migration of the store to be sync (#10592) * Migrating audit store * Final migration example for the audit store * async example * Ending migration * Removing Async helper * Fixing tests * Fixing govet problems with the StoreResult instanstiation --- api4/apitestlib.go | 5 +- app/audit.go | 12 +- app/bot.go | 20 +- app/channel.go | 126 ++++++------ app/command.go | 8 +- app/email_batching.go | 12 +- app/email_batching_test.go | 13 +- app/export.go | 7 +- app/oauth.go | 11 +- app/post.go | 18 +- app/session.go | 16 +- app/team.go | 36 +++- app/user.go | 26 +-- app/webhook.go | 7 +- cmd/mattermost/commands/userargs.go | 4 +- store/sqlstore/audit_store.go | 99 +++++----- store/sqlstore/channel_store.go | 20 +- store/sqlstore/user_store.go | 31 ++- store/store.go | 12 +- store/storetest/audit_store.go | 62 +++--- store/storetest/channel_store.go | 63 +++--- store/storetest/mocks/AuditStore.go | 55 ++++-- store/storetest/mocks/ChannelStore.go | 19 +- store/storetest/mocks/GroupStore.go | 64 +++--- .../mocks/LayeredStoreDatabaseLayer.go | 184 +++++++++--------- store/storetest/mocks/LayeredStoreSupplier.go | 184 +++++++++--------- store/storetest/mocks/UserStore.go | 19 +- store/storetest/team_store.go | 3 +- store/storetest/user_store.go | 36 ++-- web/context.go | 8 +- 30 files changed, 597 insertions(+), 583 deletions(-) diff --git a/api4/apitestlib.go b/api4/apitestlib.go index ab47b9789c..77a7ecd24b 100644 --- a/api4/apitestlib.go +++ b/api4/apitestlib.go @@ -733,8 +733,7 @@ func (me *TestHelper) cleanupTestFile(info *model.FileInfo) error { func (me *TestHelper) MakeUserChannelAdmin(user *model.User, channel *model.Channel) { utils.DisableDebugLogForTest() - if cmr := <-me.App.Srv.Store.Channel().GetMember(channel.Id, user.Id); cmr.Err == nil { - cm := cmr.Data.(*model.ChannelMember) + if cm, err := me.App.Srv.Store.Channel().GetMember(channel.Id, user.Id); err == nil { cm.SchemeAdmin = true if sr := <-me.App.Srv.Store.Channel().UpdateMember(cm); sr.Err != nil { utils.EnableDebugLogForTest() @@ -742,7 +741,7 @@ func (me *TestHelper) MakeUserChannelAdmin(user *model.User, channel *model.Chan } } else { utils.EnableDebugLogForTest() - panic(cmr.Err) + panic(err) } utils.EnableDebugLogForTest() diff --git a/app/audit.go b/app/audit.go index 857c2fa462..9382372ccf 100644 --- a/app/audit.go +++ b/app/audit.go @@ -8,17 +8,9 @@ import ( ) func (a *App) GetAudits(userId string, limit int) (model.Audits, *model.AppError) { - result := <-a.Srv.Store.Audit().Get(userId, 0, limit) - if result.Err != nil { - return nil, result.Err - } - return result.Data.(model.Audits), nil + return a.Srv.Store.Audit().Get(userId, 0, limit) } func (a *App) GetAuditsPage(userId string, page int, perPage int) (model.Audits, *model.AppError) { - result := <-a.Srv.Store.Audit().Get(userId, page*perPage, perPage) - if result.Err != nil { - return nil, result.Err - } - return result.Data.(model.Audits), nil + return a.Srv.Store.Audit().Get(userId, page*perPage, perPage) } diff --git a/app/bot.go b/app/bot.go index fac6d5747b..371ead7081 100644 --- a/app/bot.go +++ b/app/bot.go @@ -34,22 +34,21 @@ func (a *App) PatchBot(botUserId string, botPatch *model.BotPatch) (*model.Bot, bot.Patch(botPatch) - result := <-a.Srv.Store.User().Get(botUserId) - if result.Err != nil { - return nil, result.Err + user, err := a.Srv.Store.User().Get(botUserId) + if err != nil { + return nil, err } - user := result.Data.(*model.User) patchedUser := model.UserFromBot(bot) user.Id = patchedUser.Id user.Username = patchedUser.Username user.Email = patchedUser.Email user.FirstName = patchedUser.FirstName - if result = <-a.Srv.Store.User().Update(user, true); result.Err != nil { + if result := <-a.Srv.Store.User().Update(user, true); result.Err != nil { return nil, result.Err } - result = <-a.Srv.Store.Bot().Update(bot) + result := <-a.Srv.Store.Bot().Update(bot) if result.Err != nil { return nil, result.Err } @@ -79,17 +78,16 @@ func (a *App) GetBots(options *model.BotGetOptions) (model.BotList, *model.AppEr // UpdateBotActive marks a bot as active or inactive, along with its corresponding user. func (a *App) UpdateBotActive(botUserId string, active bool) (*model.Bot, *model.AppError) { - result := <-a.Srv.Store.User().Get(botUserId) - if result.Err != nil { - return nil, result.Err + user, err := a.Srv.Store.User().Get(botUserId) + if err != nil { + return nil, err } - user := result.Data.(*model.User) if _, err := a.UpdateActive(user, active); err != nil { return nil, err } - result = <-a.Srv.Store.Bot().Get(botUserId, true) + result := <-a.Srv.Store.Bot().Get(botUserId, true) if result.Err != nil { return nil, result.Err } diff --git a/app/channel.go b/app/channel.go index 72a79376c5..3e64957668 100644 --- a/app/channel.go +++ b/app/channel.go @@ -36,11 +36,11 @@ func (a *App) CreateDefaultChannels(teamId string) ([]*model.Channel, *model.App func (a *App) JoinDefaultChannels(teamId string, user *model.User, shouldBeAdmin bool, userRequestorId string) *model.AppError { var requestor *model.User if userRequestorId != "" { - u := <-a.Srv.Store.User().Get(userRequestorId) - if u.Err != nil { - return u.Err + var err *model.AppError + requestor, err = a.Srv.Store.User().Get(userRequestorId) + if err != nil { + return err } - requestor = u.Data.(*model.User) } defaultChannelList := []string{"town-square"} @@ -312,8 +312,18 @@ func (a *App) GetOrCreateDirectChannel(userId, otherUserId string) (*model.Chann } func (a *App) createDirectChannel(userId string, otherUserId string) (*model.Channel, *model.AppError) { - uc1 := a.Srv.Store.User().Get(userId) - uc2 := a.Srv.Store.User().Get(otherUserId) + uc1 := make(chan store.StoreResult, 1) + uc2 := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(userId) + uc1 <- store.StoreResult{Data: user, Err: err} + close(uc1) + }() + go func() { + user, err := a.Srv.Store.User().Get(otherUserId) + uc2 <- store.StoreResult{Data: user, Err: err} + close(uc2) + }() if result := <-uc1; result.Err != nil { return nil, model.NewAppError("CreateDirectChannel", "api.channel.create_direct_channel.invalid_user.app_error", nil, userId, http.StatusBadRequest) @@ -354,15 +364,15 @@ func (a *App) WaitForChannelMembership(channelId string, userId string) { time.Sleep(100 * time.Millisecond) - result := <-a.Srv.Store.Channel().GetMember(channelId, userId) + _, err := a.Srv.Store.Channel().GetMember(channelId, userId) // If the membership was found then return - if result.Err == nil { + if err == nil { return } // If we received a error but it wasn't a missing channel member then return - if result.Err.Id != store.MISSING_CHANNEL_MEMBER_ERROR { + if err.Id != store.MISSING_CHANNEL_MEMBER_ERROR { return } } @@ -761,12 +771,11 @@ func (a *App) DeleteChannel(channel *model.Channel, userId string) *model.AppErr var user *model.User if userId != "" { - uc := a.Srv.Store.User().Get(userId) - uresult := <-uc - if uresult.Err != nil { - return uresult.Err + var err *model.AppError + user, err = a.Srv.Store.User().Get(userId) + if err != nil { + return err } - user = uresult.Data.(*model.User) } ihcresult := <-ihc @@ -844,14 +853,12 @@ func (a *App) addUserToChannel(user *model.User, channel *model.Channel, teamMem return nil, model.NewAppError("AddUserToChannel", "api.channel.add_user_to_channel.type.app_error", nil, "", http.StatusBadRequest) } - cmchan := a.Srv.Store.Channel().GetMember(channel.Id, user.Id) - - if result := <-cmchan; result.Err != nil { - if result.Err.Id != store.MISSING_CHANNEL_MEMBER_ERROR { - return nil, result.Err + channelMember, err := a.Srv.Store.Channel().GetMember(channel.Id, user.Id) + if err != nil { + if err.Id != store.MISSING_CHANNEL_MEMBER_ERROR { + return nil, err } } else { - channelMember := result.Data.(*model.ChannelMember) return channelMember, nil } @@ -904,12 +911,12 @@ func (a *App) AddUserToChannel(user *model.User, channel *model.Channel) (*model } func (a *App) AddChannelMember(userId string, channel *model.Channel, userRequestorId string, postRootId string, currentSessionId string) (*model.ChannelMember, *model.AppError) { - if result := <-a.Srv.Store.Channel().GetMember(channel.Id, userId); result.Err != nil { - if result.Err.Id != store.MISSING_CHANNEL_MEMBER_ERROR { - return nil, result.Err + if member, err := a.Srv.Store.Channel().GetMember(channel.Id, userId); err != nil { + if err.Id != store.MISSING_CHANNEL_MEMBER_ERROR { + return nil, err } } else { - return result.Data.(*model.ChannelMember), nil + return member, nil } var user *model.User @@ -1003,15 +1010,11 @@ func (a *App) AddDirectChannels(teamId string, user *model.User) *model.AppError } func (a *App) PostUpdateChannelHeaderMessage(userId string, channel *model.Channel, oldChannelHeader, newChannelHeader string) *model.AppError { - uc := a.Srv.Store.User().Get(userId) - - uresult := <-uc - if uresult.Err != nil { - return model.NewAppError("PostUpdateChannelHeaderMessage", "api.channel.post_update_channel_header_message_and_forget.retrieve_user.error", nil, uresult.Err.Error(), http.StatusBadRequest) + user, err := a.Srv.Store.User().Get(userId) + if err != nil { + return model.NewAppError("PostUpdateChannelHeaderMessage", "api.channel.post_update_channel_header_message_and_forget.retrieve_user.error", nil, err.Error(), http.StatusBadRequest) } - user := uresult.Data.(*model.User) - var message string if oldChannelHeader == "" { message = fmt.Sprintf(utils.T("api.channel.post_update_channel_header_message_and_forget.updated_to"), user.Username, newChannelHeader) @@ -1041,15 +1044,11 @@ func (a *App) PostUpdateChannelHeaderMessage(userId string, channel *model.Chann } func (a *App) PostUpdateChannelPurposeMessage(userId string, channel *model.Channel, oldChannelPurpose string, newChannelPurpose string) *model.AppError { - uc := a.Srv.Store.User().Get(userId) - - uresult := <-uc - if uresult.Err != nil { - return model.NewAppError("PostUpdateChannelPurposeMessage", "app.channel.post_update_channel_purpose_message.retrieve_user.error", nil, uresult.Err.Error(), http.StatusBadRequest) + user, err := a.Srv.Store.User().Get(userId) + if err != nil { + return model.NewAppError("PostUpdateChannelPurposeMessage", "app.channel.post_update_channel_purpose_message.retrieve_user.error", nil, err.Error(), http.StatusBadRequest) } - user := uresult.Data.(*model.User) - var message string if oldChannelPurpose == "" { message = fmt.Sprintf(utils.T("app.channel.post_update_channel_purpose_message.updated_to"), user.Username, newChannelPurpose) @@ -1078,15 +1077,11 @@ func (a *App) PostUpdateChannelPurposeMessage(userId string, channel *model.Chan } func (a *App) PostUpdateChannelDisplayNameMessage(userId string, channel *model.Channel, oldChannelDisplayName, newChannelDisplayName string) *model.AppError { - uc := a.Srv.Store.User().Get(userId) - - uresult := <-uc - if uresult.Err != nil { - return model.NewAppError("PostUpdateChannelDisplayNameMessage", "api.channel.post_update_channel_displayname_message_and_forget.retrieve_user.error", nil, uresult.Err.Error(), http.StatusBadRequest) + user, err := a.Srv.Store.User().Get(userId) + if err != nil { + return model.NewAppError("PostUpdateChannelDisplayNameMessage", "api.channel.post_update_channel_displayname_message_and_forget.retrieve_user.error", nil, err.Error(), http.StatusBadRequest) } - user := uresult.Data.(*model.User) - message := fmt.Sprintf(utils.T("api.channel.post_update_channel_displayname_message_and_forget.updated_from"), user.Username, oldChannelDisplayName, newChannelDisplayName) post := &model.Post{ @@ -1234,11 +1229,7 @@ func (a *App) GetPublicChannelsForTeam(teamId string, offset int, limit int) (*m } func (a *App) GetChannelMember(channelId string, userId string) (*model.ChannelMember, *model.AppError) { - result := <-a.Srv.Store.Channel().GetMember(channelId, userId) - if result.Err != nil { - return nil, result.Err - } - return result.Data.(*model.ChannelMember), nil + return a.Srv.Store.Channel().GetMember(channelId, userId) } func (a *App) GetChannelMembersPage(channelId string, page, perPage int) (*model.ChannelMembers, *model.AppError) { @@ -1330,8 +1321,18 @@ func (a *App) GetChannelUnread(channelId, userId string) (*model.ChannelUnread, } func (a *App) JoinChannel(channel *model.Channel, userId string) *model.AppError { - userChan := a.Srv.Store.User().Get(userId) - memberChan := a.Srv.Store.Channel().GetMember(channel.Id, userId) + userChan := make(chan store.StoreResult, 1) + memberChan := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(userId) + userChan <- store.StoreResult{Data: user, Err: err} + close(userChan) + }() + go func() { + member, err := a.Srv.Store.Channel().GetMember(channel.Id, userId) + memberChan <- store.StoreResult{Data: member, Err: err} + close(memberChan) + }() uresult := <-userChan if uresult.Err != nil { @@ -1419,7 +1420,12 @@ func (a *App) postJoinTeamMessage(user *model.User, channel *model.Channel) *mod func (a *App) LeaveChannel(channelId string, userId string) *model.AppError { sc := a.Srv.Store.Channel().Get(channelId, true) - uc := a.Srv.Store.User().Get(userId) + uc := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(userId) + uc <- store.StoreResult{Data: user, Err: err} + close(uc) + }() ccm := a.Srv.Store.Channel().GetMemberCount(channelId, false) cresult := <-sc @@ -1759,12 +1765,11 @@ func (a *App) MarkChannelsAsViewed(channelIds []string, userId string, currentSe } channel := chanResult.Data.(*model.Channel) - result := <-a.Srv.Store.Channel().GetMember(channelId, userId) - if result.Err != nil { - mlog.Warn(fmt.Sprintf("Failed to get membership %v", result.Err)) + member, err := a.Srv.Store.Channel().GetMember(channelId, userId) + if err != nil { + mlog.Warn(fmt.Sprintf("Failed to get membership %v", err)) continue } - member := result.Data.(*model.ChannelMember) notify := member.NotifyProps[model.PUSH_NOTIFY_PROP] if notify == model.CHANNEL_NOTIFY_DEFAULT { @@ -1951,14 +1956,11 @@ func (a *App) GetPinnedPosts(channelId string) (*model.PostList, *model.AppError } func (a *App) ToggleMuteChannel(channelId string, userId string) *model.ChannelMember { - result := <-a.Srv.Store.Channel().GetMember(channelId, userId) - - if result.Err != nil { + member, err := a.Srv.Store.Channel().GetMember(channelId, userId) + if err != nil { return nil } - member := result.Data.(*model.ChannelMember) - if member.NotifyProps[model.MARK_UNREAD_NOTIFY_PROP] == model.CHANNEL_NOTIFY_MENTION { member.NotifyProps[model.MARK_UNREAD_NOTIFY_PROP] = model.CHANNEL_MARK_UNREAD_ALL } else { diff --git a/app/command.go b/app/command.go index 12a3d3269b..4fef960e10 100644 --- a/app/command.go +++ b/app/command.go @@ -13,6 +13,7 @@ import ( "github.com/mattermost/mattermost-server/mlog" "github.com/mattermost/mattermost-server/model" + "github.com/mattermost/mattermost-server/store" "github.com/mattermost/mattermost-server/utils" goi18n "github.com/nicksnyder/go-i18n/i18n" ) @@ -219,7 +220,12 @@ func (a *App) tryExecuteCustomCommand(args *model.CommandArgs, trigger string, m chanChan := a.Srv.Store.Channel().Get(args.ChannelId, true) teamChan := a.Srv.Store.Team().Get(args.TeamId) - userChan := a.Srv.Store.User().Get(args.UserId) + userChan := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(args.UserId) + userChan <- store.StoreResult{Data: user, Err: err} + close(userChan) + }() result := <-a.Srv.Store.Command().GetByTeam(args.TeamId) if result.Err != nil { diff --git a/app/email_batching.go b/app/email_batching.go index 30b82b3675..b55122fcc9 100644 --- a/app/email_batching.go +++ b/app/email_batching.go @@ -199,26 +199,24 @@ func (job *EmailBatchingJob) checkPendingNotifications(now time.Time, handler fu } func (s *Server) sendBatchedEmailNotification(userId string, notifications []*batchedNotification) { - result := <-s.Store.User().Get(userId) - if result.Err != nil { + user, err := s.Store.User().Get(userId) + if err != nil { mlog.Warn("Unable to find recipient for batched email notification") return } - user := result.Data.(*model.User) translateFunc := utils.GetUserTranslations(user.Locale) displayNameFormat := *s.Config().TeamSettings.TeammateNameDisplay var contents string for _, notification := range notifications { - result := <-s.Store.User().Get(notification.post.UserId) - if result.Err != nil { + sender, err := s.Store.User().Get(notification.post.UserId) + if err != nil { mlog.Warn("Unable to find sender of post for batched email notification") continue } - sender := result.Data.(*model.User) - result = <-s.Store.Channel().Get(notification.post.ChannelId, true) + result := <-s.Store.Channel().Get(notification.post.ChannelId, true) if result.Err != nil { mlog.Warn("Unable to find channel of post for batched email notification") continue diff --git a/app/email_batching_test.go b/app/email_batching_test.go index a8d9d6be1c..8badd8a7a3 100644 --- a/app/email_batching_test.go +++ b/app/email_batching_test.go @@ -10,6 +10,7 @@ import ( "github.com/mattermost/mattermost-server/model" "github.com/mattermost/mattermost-server/store" + "github.com/stretchr/testify/require" ) func TestHandleNewNotifications(t *testing.T) { @@ -109,7 +110,8 @@ func TestCheckPendingNotifications(t *testing.T) { }, } - channelMember := store.Must(th.App.Srv.Store.Channel().GetMember(th.BasicChannel.Id, th.BasicUser.Id)).(*model.ChannelMember) + channelMember, err := th.App.Srv.Store.Channel().GetMember(th.BasicChannel.Id, th.BasicUser.Id) + require.Nil(t, err) channelMember.LastViewedAt = 9999999 store.Must(th.App.Srv.Store.Channel().UpdateMember(channelMember)) @@ -128,7 +130,8 @@ func TestCheckPendingNotifications(t *testing.T) { } // test that notifications are cleared if the user has acted - channelMember = store.Must(th.App.Srv.Store.Channel().GetMember(th.BasicChannel.Id, th.BasicUser.Id)).(*model.ChannelMember) + channelMember, err = th.App.Srv.Store.Channel().GetMember(th.BasicChannel.Id, th.BasicUser.Id) + require.Nil(t, err) channelMember.LastViewedAt = 10001000 store.Must(th.App.Srv.Store.Channel().UpdateMember(channelMember)) @@ -208,7 +211,8 @@ func TestCheckPendingNotificationsDefaultInterval(t *testing.T) { job := NewEmailBatchingJob(th.Server, 128) // bypasses recent user activity check - channelMember := store.Must(th.App.Srv.Store.Channel().GetMember(th.BasicChannel.Id, th.BasicUser.Id)).(*model.ChannelMember) + channelMember, err := th.App.Srv.Store.Channel().GetMember(th.BasicChannel.Id, th.BasicUser.Id) + require.Nil(t, err) channelMember.LastViewedAt = 9999000 store.Must(th.App.Srv.Store.Channel().UpdateMember(channelMember)) @@ -246,7 +250,8 @@ func TestCheckPendingNotificationsCantParseInterval(t *testing.T) { job := NewEmailBatchingJob(th.Server, 128) // bypasses recent user activity check - channelMember := store.Must(th.App.Srv.Store.Channel().GetMember(th.BasicChannel.Id, th.BasicUser.Id)).(*model.ChannelMember) + channelMember, err := th.App.Srv.Store.Channel().GetMember(th.BasicChannel.Id, th.BasicUser.Id) + require.Nil(t, err) channelMember.LastViewedAt = 9999000 store.Must(th.App.Srv.Store.Channel().UpdateMember(channelMember)) diff --git a/app/export.go b/app/export.go index 3b7f827e5e..2190027c88 100644 --- a/app/export.go +++ b/app/export.go @@ -415,11 +415,10 @@ func (a *App) BuildPostReactions(postId string) (*[]ReactionImportData, *model.A reactions := result.Data.([]*model.Reaction) for _, reaction := range reactions { - result := <-a.Srv.Store.User().Get(reaction.UserId) - if result.Err != nil { - return nil, result.Err + user, err := a.Srv.Store.User().Get(reaction.UserId) + if err != nil { + return nil, err } - user := result.Data.(*model.User) reactionsOfPost = append(reactionsOfPost, *ImportReactionFromPost(user, reaction)) } diff --git a/app/oauth.go b/app/oauth.go index 9678ccc508..624022a4f7 100644 --- a/app/oauth.go +++ b/app/oauth.go @@ -259,11 +259,11 @@ func (a *App) GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectUri, c return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.redirect_uri.app_error", nil, "", http.StatusBadRequest) } - result = <-a.Srv.Store.User().Get(authData.UserId) - if result.Err != nil { + var err *model.AppError + user, err = a.Srv.Store.User().Get(authData.UserId) + if err != nil { return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.internal_user.app_error", nil, "", http.StatusNotFound) } - user = result.Data.(*model.User) result = <-a.Srv.Store.OAuth().GetPreviousAccessData(user.Id, clientId) if result.Err != nil { @@ -318,11 +318,10 @@ func (a *App) GetOAuthAccessTokenForCodeFlow(clientId, grantType, redirectUri, c } accessData = result.Data.(*model.AccessData) - result = <-a.Srv.Store.User().Get(accessData.UserId) - if result.Err != nil { + user, err := a.Srv.Store.User().Get(accessData.UserId) + if err != nil { return nil, model.NewAppError("GetOAuthAccessToken", "api.oauth.get_access_token.internal_user.app_error", nil, "", http.StatusNotFound) } - user = result.Data.(*model.User) access, err := a.newSessionUpdateToken(oauthApp.Name, accessData, user) if err != nil { diff --git a/app/post.go b/app/post.go index 33a45350fd..57ec3ee5fc 100644 --- a/app/post.go +++ b/app/post.go @@ -50,11 +50,10 @@ func (a *App) CreatePostAsUser(post *model.Post, currentSessionId string) (*mode } if err.Id == "api.post.create_post.town_square_read_only" { - result := <-a.Srv.Store.User().Get(post.UserId) - if result.Err != nil { - return nil, result.Err + user, userErr := a.Srv.Store.User().Get(post.UserId) + if userErr != nil { + return nil, userErr } - user := result.Data.(*model.User) T := utils.GetUserTranslations(user.Locale) a.SendEphemeralPost( @@ -164,11 +163,10 @@ func (a *App) CreatePost(post *model.Post, channel *model.Channel, triggerWebhoo pchan = a.Srv.Store.Post().Get(post.RootId) } - result := <-a.Srv.Store.User().Get(post.UserId) - if result.Err != nil { - return nil, result.Err + user, err := a.Srv.Store.User().Get(post.UserId) + if err != nil { + return nil, err } - user := result.Data.(*model.User) if a.License() != nil && *a.Config().TeamSettings.ExperimentalTownSquareIsReadOnly && !post.IsSystemMessage() && @@ -180,7 +178,7 @@ func (a *App) CreatePost(post *model.Post, channel *model.Channel, triggerWebhoo // Verify the parent/child relationships are correct var parentPostList *model.PostList if pchan != nil { - result = <-pchan + result := <-pchan if result.Err != nil { return nil, model.NewAppError("createPost", "api.post.create_post.root_id.app_error", nil, "", http.StatusBadRequest) } @@ -245,7 +243,7 @@ func (a *App) CreatePost(post *model.Post, channel *model.Channel, triggerWebhoo } } - result = <-a.Srv.Store.Post().Save(post) + result := <-a.Srv.Store.Post().Save(post) if result.Err != nil { return nil, result.Err } diff --git a/app/session.go b/app/session.go index a95d9f66d5..c36cb84f6c 100644 --- a/app/session.go +++ b/app/session.go @@ -9,6 +9,7 @@ import ( "github.com/mattermost/mattermost-server/mlog" "github.com/mattermost/mattermost-server/model" + "github.com/mattermost/mattermost-server/store" ) func (a *App) CreateSession(session *model.Session) (*model.Session, *model.AppError) { @@ -247,7 +248,12 @@ func (a *App) CreateUserAccessToken(token *model.UserAccessToken) (*model.UserAc token.Token = model.NewId() - uchan := a.Srv.Store.User().Get(token.UserId) + uchan := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(token.UserId) + uchan <- store.StoreResult{Data: user, Err: err} + close(uchan) + }() result := <-a.Srv.Store.UserAccessToken().Save(token) if result.Err != nil { @@ -284,12 +290,10 @@ func (a *App) createSessionForUserAccessToken(tokenString string) (*model.Sessio return nil, model.NewAppError("createSessionForUserAccessToken", "app.user_access_token.invalid_or_missing", nil, "inactive_token", http.StatusUnauthorized) } - var user *model.User - result = <-a.Srv.Store.User().Get(token.UserId) - if result.Err != nil { - return nil, result.Err + user, err := a.Srv.Store.User().Get(token.UserId) + if err != nil { + return nil, err } - user = result.Data.(*model.User) if user.DeleteAt != 0 { return nil, model.NewAppError("createSessionForUserAccessToken", "app.user_access_token.invalid_or_missing", nil, "inactive_user_id="+user.Id, http.StatusUnauthorized) diff --git a/app/team.go b/app/team.go index cf81cb877f..d316407356 100644 --- a/app/team.go +++ b/app/team.go @@ -18,6 +18,7 @@ import ( "github.com/mattermost/mattermost-server/mlog" "github.com/mattermost/mattermost-server/model" "github.com/mattermost/mattermost-server/plugin" + "github.com/mattermost/mattermost-server/store" "github.com/mattermost/mattermost-server/utils" ) @@ -327,7 +328,12 @@ func (a *App) sendUpdatedMemberRoleEvent(userId string, member *model.TeamMember func (a *App) AddUserToTeam(teamId string, userId string, userRequestorId string) (*model.Team, *model.AppError) { tchan := a.Srv.Store.Team().Get(teamId) - uchan := a.Srv.Store.User().Get(userId) + uchan := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(userId) + uchan <- store.StoreResult{Data: user, Err: err} + close(uchan) + }() result := <-tchan if result.Err != nil { @@ -375,7 +381,12 @@ func (a *App) AddUserToTeamByToken(userId string, tokenId string) (*model.Team, tokenData := model.MapFromJson(strings.NewReader(token.Extra)) tchan := a.Srv.Store.Team().Get(tokenData["teamId"]) - uchan := a.Srv.Store.User().Get(userId) + uchan := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(userId) + uchan <- store.StoreResult{Data: user, Err: err} + close(uchan) + }() result = <-tchan if result.Err != nil { @@ -402,7 +413,12 @@ func (a *App) AddUserToTeamByToken(userId string, tokenId string) (*model.Team, func (a *App) AddUserToTeamByInviteId(inviteId string, userId string) (*model.Team, *model.AppError) { tchan := a.Srv.Store.Team().GetByInviteId(inviteId) - uchan := a.Srv.Store.User().Get(userId) + uchan := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(userId) + uchan <- store.StoreResult{Data: user, Err: err} + close(uchan) + }() result := <-tchan if result.Err != nil { @@ -761,7 +777,12 @@ func (a *App) GetTeamUnread(teamId, userId string) (*model.TeamUnread, *model.Ap func (a *App) RemoveUserFromTeam(teamId string, userId string, requestorId string) *model.AppError { tchan := a.Srv.Store.Team().Get(teamId) - uchan := a.Srv.Store.User().Get(userId) + uchan := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(userId) + uchan <- store.StoreResult{Data: user, Err: err} + close(uchan) + }() result := <-tchan if result.Err != nil { @@ -926,7 +947,12 @@ func (a *App) InviteNewUsersToTeam(emailList []string, teamId, senderId string) } tchan := a.Srv.Store.Team().Get(teamId) - uchan := a.Srv.Store.User().Get(senderId) + uchan := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(senderId) + uchan <- store.StoreResult{Data: user, Err: err} + close(uchan) + }() result := <-tchan if result.Err != nil { diff --git a/app/user.go b/app/user.go index 06c56b087e..4b07c49318 100644 --- a/app/user.go +++ b/app/user.go @@ -398,11 +398,7 @@ func (a *App) IsUsernameTaken(name string) bool { } func (a *App) GetUser(userId string) (*model.User, *model.AppError) { - result := <-a.Srv.Store.User().Get(userId) - if result.Err != nil { - return nil, result.Err - } - return result.Data.(*model.User), nil + return a.Srv.Store.User().Get(userId) } func (a *App) GetUserByUsername(username string) (*model.User, *model.AppError) { @@ -655,11 +651,10 @@ func (a *App) GenerateMfaSecret(userId string) (*model.MfaSecret, *model.AppErro } func (a *App) ActivateMfa(userId, token string) *model.AppError { - result := <-a.Srv.Store.User().Get(userId) - if result.Err != nil { - return result.Err + user, err := a.Srv.Store.User().Get(userId) + if err != nil { + return err } - user := result.Data.(*model.User) if len(user.AuthService) > 0 && user.AuthService != model.USER_AUTH_SERVICE_LDAP { return model.NewAppError("ActivateMfa", "api.user.activate_mfa.email_and_ldap_only.app_error", nil, "", http.StatusBadRequest) @@ -1086,11 +1081,10 @@ func (a *App) sendUpdatedUserEvent(user model.User) { } func (a *App) UpdateUser(user *model.User, sendNotifications bool) (*model.User, *model.AppError) { - result := <-a.Srv.Store.User().Get(user.Id) - if result.Err != nil { - return nil, result.Err + prev, err := a.Srv.Store.User().Get(user.Id) + if err != nil { + return nil, err } - prev := result.Data.(*model.User) if !CheckUserDomain(user, *a.Config().TeamSettings.RestrictCreationToDomains) { if !prev.IsLDAPUser() && !prev.IsSAMLUser() && user.Email != prev.Email { @@ -1112,7 +1106,7 @@ func (a *App) UpdateUser(user *model.User, sendNotifications bool) (*model.User, user.Email = prev.Email } - result = <-a.Srv.Store.User().Update(user, false) + result := <-a.Srv.Store.User().Update(user, false) if result.Err != nil { return nil, result.Err } @@ -1482,8 +1476,8 @@ func (a *App) PermanentDeleteUser(user *model.User) *model.AppError { return result.Err } - if result := <-a.Srv.Store.Audit().PermanentDeleteByUser(user.Id); result.Err != nil { - return result.Err + if err := a.Srv.Store.Audit().PermanentDeleteByUser(user.Id); err != nil { + return err } if result := <-a.Srv.Store.Team().RemoveAllMembersByUser(user.Id); result.Err != nil { diff --git a/app/webhook.go b/app/webhook.go index fc6338bd62..69be77cd06 100644 --- a/app/webhook.go +++ b/app/webhook.go @@ -610,7 +610,12 @@ func (a *App) HandleIncomingWebhook(hookId string, req *model.IncomingWebhookReq hook = result.Data.(*model.IncomingWebhook) } - uchan := a.Srv.Store.User().Get(hook.UserId) + uchan := make(chan store.StoreResult, 1) + go func() { + user, err := a.Srv.Store.User().Get(hook.UserId) + uchan <- store.StoreResult{Data: user, Err: err} + close(uchan) + }() if len(req.Props) == 0 { req.Props = make(model.StringInterface) diff --git a/cmd/mattermost/commands/userargs.go b/cmd/mattermost/commands/userargs.go index ddeed64604..9295ef3bda 100644 --- a/cmd/mattermost/commands/userargs.go +++ b/cmd/mattermost/commands/userargs.go @@ -30,9 +30,7 @@ func getUserFromUserArg(a *app.App, userArg string) *model.User { } if user == nil { - if result := <-a.Srv.Store.User().Get(userArg); result.Err == nil { - user = result.Data.(*model.User) - } + user, _ = a.Srv.Store.User().Get(userArg) } return user diff --git a/store/sqlstore/audit_store.go b/store/sqlstore/audit_store.go index eb30058c74..258b59e26d 100644 --- a/store/sqlstore/audit_store.go +++ b/store/sqlstore/audit_store.go @@ -34,71 +34,60 @@ func (s SqlAuditStore) CreateIndexesIfNotExists() { s.CreateIndexIfNotExists("idx_audits_user_id", "Audits", "UserId") } -func (s SqlAuditStore) Save(audit *model.Audit) store.StoreChannel { - return store.Do(func(result *store.StoreResult) { - audit.Id = model.NewId() - audit.CreateAt = model.GetMillis() +func (s SqlAuditStore) Save(audit *model.Audit) *model.AppError { + audit.Id = model.NewId() + audit.CreateAt = model.GetMillis() - if err := s.GetMaster().Insert(audit); err != nil { - result.Err = model.NewAppError("SqlAuditStore.Save", "store.sql_audit.save.saving.app_error", nil, "user_id="+audit.UserId+" action="+audit.Action, http.StatusInternalServerError) - } - }) + if err := s.GetMaster().Insert(audit); err != nil { + return model.NewAppError("SqlAuditStore.Save", "store.sql_audit.save.saving.app_error", nil, "user_id="+audit.UserId+" action="+audit.Action, http.StatusInternalServerError) + } + return nil } -func (s SqlAuditStore) Get(user_id string, offset int, limit int) store.StoreChannel { - return store.Do(func(result *store.StoreResult) { - if limit > 1000 { - limit = 1000 - result.Err = model.NewAppError("SqlAuditStore.Get", "store.sql_audit.get.limit.app_error", nil, "user_id="+user_id, http.StatusBadRequest) - return - } +func (s SqlAuditStore) Get(user_id string, offset int, limit int) (model.Audits, *model.AppError) { + if limit > 1000 { + return nil, model.NewAppError("SqlAuditStore.Get", "store.sql_audit.get.limit.app_error", nil, "user_id="+user_id, http.StatusBadRequest) + } - query := "SELECT * FROM Audits" + query := "SELECT * FROM Audits" - if len(user_id) != 0 { - query += " WHERE UserId = :user_id" - } + if len(user_id) != 0 { + query += " WHERE UserId = :user_id" + } - query += " ORDER BY CreateAt DESC LIMIT :limit OFFSET :offset" + query += " ORDER BY CreateAt DESC LIMIT :limit OFFSET :offset" - var audits model.Audits - if _, err := s.GetReplica().Select(&audits, query, map[string]interface{}{"user_id": user_id, "limit": limit, "offset": offset}); err != nil { - result.Err = model.NewAppError("SqlAuditStore.Get", "store.sql_audit.get.finding.app_error", nil, "user_id="+user_id, http.StatusInternalServerError) - } else { - result.Data = audits - } - }) + var audits model.Audits + if _, err := s.GetReplica().Select(&audits, query, map[string]interface{}{"user_id": user_id, "limit": limit, "offset": offset}); err != nil { + return nil, model.NewAppError("SqlAuditStore.Get", "store.sql_audit.get.finding.app_error", nil, "user_id="+user_id, http.StatusInternalServerError) + } + return audits, nil } -func (s SqlAuditStore) PermanentDeleteByUser(userId string) store.StoreChannel { - return store.Do(func(result *store.StoreResult) { - if _, err := s.GetMaster().Exec("DELETE FROM Audits WHERE UserId = :userId", - map[string]interface{}{"userId": userId}); err != nil { - result.Err = model.NewAppError("SqlAuditStore.Delete", "store.sql_audit.permanent_delete_by_user.app_error", nil, "user_id="+userId, http.StatusInternalServerError) - } - }) +func (s SqlAuditStore) PermanentDeleteByUser(userId string) *model.AppError { + if _, err := s.GetMaster().Exec("DELETE FROM Audits WHERE UserId = :userId", + map[string]interface{}{"userId": userId}); err != nil { + return model.NewAppError("SqlAuditStore.Delete", "store.sql_audit.permanent_delete_by_user.app_error", nil, "user_id="+userId, http.StatusInternalServerError) + } + return nil } -func (s SqlAuditStore) PermanentDeleteBatch(endTime int64, limit int64) store.StoreChannel { - return store.Do(func(result *store.StoreResult) { - var query string - if s.DriverName() == "postgres" { - query = "DELETE from Audits WHERE Id = any (array (SELECT Id FROM Audits WHERE CreateAt < :EndTime LIMIT :Limit))" - } else { - query = "DELETE from Audits WHERE CreateAt < :EndTime LIMIT :Limit" - } +func (s SqlAuditStore) PermanentDeleteBatch(endTime int64, limit int64) (int64, *model.AppError) { + var query string + if s.DriverName() == "postgres" { + query = "DELETE from Audits WHERE Id = any (array (SELECT Id FROM Audits WHERE CreateAt < :EndTime LIMIT :Limit))" + } else { + query = "DELETE from Audits WHERE CreateAt < :EndTime LIMIT :Limit" + } - sqlResult, err := s.GetMaster().Exec(query, map[string]interface{}{"EndTime": endTime, "Limit": limit}) - if err != nil { - result.Err = model.NewAppError("SqlAuditStore.PermanentDeleteBatch", "store.sql_audit.permanent_delete_batch.app_error", nil, ""+err.Error(), http.StatusInternalServerError) - } else { - rowsAffected, err1 := sqlResult.RowsAffected() - if err1 != nil { - result.Err = model.NewAppError("SqlAuditStore.PermanentDeleteBatch", "store.sql_audit.permanent_delete_batch.app_error", nil, ""+err.Error(), http.StatusInternalServerError) - result.Data = int64(0) - } else { - result.Data = rowsAffected - } - } - }) + sqlResult, err := s.GetMaster().Exec(query, map[string]interface{}{"EndTime": endTime, "Limit": limit}) + if err != nil { + return 0, model.NewAppError("SqlAuditStore.PermanentDeleteBatch", "store.sql_audit.permanent_delete_batch.app_error", nil, ""+err.Error(), http.StatusInternalServerError) + } + + rowsAffected, err1 := sqlResult.RowsAffected() + if err1 != nil { + return 0, model.NewAppError("SqlAuditStore.PermanentDeleteBatch", "store.sql_audit.permanent_delete_batch.app_error", nil, ""+err.Error(), http.StatusInternalServerError) + } + return rowsAffected, nil } diff --git a/store/sqlstore/channel_store.go b/store/sqlstore/channel_store.go index dc167c795d..bdb2d36044 100644 --- a/store/sqlstore/channel_store.go +++ b/store/sqlstore/channel_store.go @@ -1410,21 +1410,17 @@ func (s SqlChannelStore) GetChannelMembersTimezones(channelId string) store.Stor }) } -func (s SqlChannelStore) GetMember(channelId string, userId string) store.StoreChannel { - return store.Do(func(result *store.StoreResult) { - var dbMember channelMemberWithSchemeRoles +func (s SqlChannelStore) GetMember(channelId string, userId string) (*model.ChannelMember, *model.AppError) { + var dbMember channelMemberWithSchemeRoles - if err := s.GetReplica().SelectOne(&dbMember, CHANNEL_MEMBERS_WITH_SCHEME_SELECT_QUERY+"WHERE ChannelMembers.ChannelId = :ChannelId AND ChannelMembers.UserId = :UserId", map[string]interface{}{"ChannelId": channelId, "UserId": userId}); err != nil { - if err == sql.ErrNoRows { - result.Err = model.NewAppError("SqlChannelStore.GetMember", store.MISSING_CHANNEL_MEMBER_ERROR, nil, "channel_id="+channelId+"user_id="+userId+","+err.Error(), http.StatusNotFound) - return - } - result.Err = model.NewAppError("SqlChannelStore.GetMember", "store.sql_channel.get_member.app_error", nil, "channel_id="+channelId+"user_id="+userId+","+err.Error(), http.StatusInternalServerError) - return + if err := s.GetReplica().SelectOne(&dbMember, CHANNEL_MEMBERS_WITH_SCHEME_SELECT_QUERY+"WHERE ChannelMembers.ChannelId = :ChannelId AND ChannelMembers.UserId = :UserId", map[string]interface{}{"ChannelId": channelId, "UserId": userId}); err != nil { + if err == sql.ErrNoRows { + return nil, model.NewAppError("SqlChannelStore.GetMember", store.MISSING_CHANNEL_MEMBER_ERROR, nil, "channel_id="+channelId+"user_id="+userId+","+err.Error(), http.StatusNotFound) } + return nil, model.NewAppError("SqlChannelStore.GetMember", "store.sql_channel.get_member.app_error", nil, "channel_id="+channelId+"user_id="+userId+","+err.Error(), http.StatusInternalServerError) + } - result.Data = dbMember.ToModel() - }) + return dbMember.ToModel(), nil } func (s SqlChannelStore) InvalidateAllChannelMembersForUser(userId string) { diff --git a/store/sqlstore/user_store.go b/store/sqlstore/user_store.go index 4e9c3f948d..f7aec2045d 100644 --- a/store/sqlstore/user_store.go +++ b/store/sqlstore/user_store.go @@ -329,27 +329,22 @@ func (us SqlUserStore) UpdateMfaActive(userId string, active bool) store.StoreCh }) } -func (us SqlUserStore) Get(id string) store.StoreChannel { - return store.Do(func(result *store.StoreResult) { - query := us.usersQuery.Where("Id = ?", id) +func (us SqlUserStore) Get(id string) (*model.User, *model.AppError) { + query := us.usersQuery.Where("Id = ?", id) - queryString, args, err := query.ToSql() - if err != nil { - result.Err = model.NewAppError("SqlUserStore.Get", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) - return - } + queryString, args, err := query.ToSql() + if err != nil { + return nil, model.NewAppError("SqlUserStore.Get", "store.sql_user.app_error", nil, err.Error(), http.StatusInternalServerError) + } - user := &model.User{} - if err := us.GetReplica().SelectOne(user, queryString, args...); err == sql.ErrNoRows { - result.Err = model.NewAppError("SqlUserStore.Get", store.MISSING_ACCOUNT_ERROR, nil, "user_id="+id, http.StatusNotFound) - return - } else if err != nil { - result.Err = model.NewAppError("SqlUserStore.Get", "store.sql_user.get.app_error", nil, "user_id="+id+", "+err.Error(), http.StatusInternalServerError) - return - } + user := &model.User{} + if err := us.GetReplica().SelectOne(user, queryString, args...); err == sql.ErrNoRows { + return nil, model.NewAppError("SqlUserStore.Get", store.MISSING_ACCOUNT_ERROR, nil, "user_id="+id, http.StatusNotFound) + } else if err != nil { + return nil, model.NewAppError("SqlUserStore.Get", "store.sql_user.get.app_error", nil, "user_id="+id+", "+err.Error(), http.StatusInternalServerError) + } - result.Data = user - }) + return user, nil } func (us SqlUserStore) GetAll() store.StoreChannel { diff --git a/store/store.go b/store/store.go index e4e7ff5a18..2de7cb1142 100644 --- a/store/store.go +++ b/store/store.go @@ -156,7 +156,7 @@ type ChannelStore interface { SaveMember(member *model.ChannelMember) StoreChannel UpdateMember(member *model.ChannelMember) StoreChannel GetMembers(channelId string, offset, limit int) StoreChannel - GetMember(channelId string, userId string) StoreChannel + GetMember(channelId string, userId string) (*model.ChannelMember, *model.AppError) GetChannelMembersTimezones(channelId string) StoreChannel GetAllChannelMembersForUser(userId string, allowFromCache bool, includeDeleted bool) StoreChannel InvalidateAllChannelMembersForUser(userId string) @@ -248,7 +248,7 @@ type UserStore interface { UpdateAuthData(userId string, service string, authData *string, email string, resetMfa bool) StoreChannel UpdateMfaSecret(userId, secret string) StoreChannel UpdateMfaActive(userId string, active bool) StoreChannel - Get(id string) StoreChannel + Get(id string) (*model.User, *model.AppError) GetAll() StoreChannel ClearCaches() InvalidateProfilesInChannelCacheByUser(userId string) @@ -322,10 +322,10 @@ type SessionStore interface { } type AuditStore interface { - Save(audit *model.Audit) StoreChannel - Get(user_id string, offset int, limit int) StoreChannel - PermanentDeleteByUser(userId string) StoreChannel - PermanentDeleteBatch(endTime int64, limit int64) StoreChannel + Save(audit *model.Audit) *model.AppError + Get(user_id string, offset int, limit int) (model.Audits, *model.AppError) + PermanentDeleteByUser(userId string) *model.AppError + PermanentDeleteBatch(endTime int64, limit int64) (int64, *model.AppError) } type ClusterDiscoveryStore interface { diff --git a/store/storetest/audit_store.go b/store/storetest/audit_store.go index 2e107dba8c..de7c18927c 100644 --- a/store/storetest/audit_store.go +++ b/store/storetest/audit_store.go @@ -9,6 +9,8 @@ import ( "github.com/mattermost/mattermost-server/model" "github.com/mattermost/mattermost-server/store" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestAuditStore(t *testing.T, ss store.Store) { @@ -18,74 +20,58 @@ func TestAuditStore(t *testing.T, ss store.Store) { func testAuditStore(t *testing.T, ss store.Store) { audit := &model.Audit{UserId: model.NewId(), IpAddress: "ipaddress", Action: "Action"} - store.Must(ss.Audit().Save(audit)) + require.Nil(t, ss.Audit().Save(audit)) time.Sleep(100 * time.Millisecond) - store.Must(ss.Audit().Save(audit)) + require.Nil(t, ss.Audit().Save(audit)) time.Sleep(100 * time.Millisecond) - store.Must(ss.Audit().Save(audit)) + require.Nil(t, ss.Audit().Save(audit)) time.Sleep(100 * time.Millisecond) audit.ExtraInfo = "extra" time.Sleep(100 * time.Millisecond) - store.Must(ss.Audit().Save(audit)) + require.Nil(t, ss.Audit().Save(audit)) time.Sleep(100 * time.Millisecond) - c := ss.Audit().Get(audit.UserId, 0, 100) - result := <-c - audits := result.Data.(model.Audits) + audits, err := ss.Audit().Get(audit.UserId, 0, 100) + require.Nil(t, err) - if len(audits) != 4 { - t.Fatal("Failed to save and retrieve 4 audit logs") - } + assert.Len(t, audits, 4) - if audits[0].ExtraInfo != "extra" { - t.Fatal("Failed to save property for extra info") - } + assert.Equal(t, "extra", audits[0].ExtraInfo) - c = ss.Audit().Get("missing", 0, 100) - result = <-c - audits = result.Data.(model.Audits) + audits, err = ss.Audit().Get("missing", 0, 100) - if len(audits) != 0 { - t.Fatal("Should have returned empty because user_id is missing") - } + assert.Len(t, audits, 0) - c = ss.Audit().Get("", 0, 100) - result = <-c - audits = result.Data.(model.Audits) + audits, err = ss.Audit().Get("", 0, 100) if len(audits) < 4 { t.Fatal("Failed to save and retrieve 4 audit logs") } - if r2 := <-ss.Audit().PermanentDeleteByUser(audit.UserId); r2.Err != nil { - t.Fatal(r2.Err) - } + require.Nil(t, ss.Audit().PermanentDeleteByUser(audit.UserId)) } func testAuditStorePermanentDeleteBatch(t *testing.T, ss store.Store) { a1 := &model.Audit{UserId: model.NewId(), IpAddress: "ipaddress", Action: "Action"} - store.Must(ss.Audit().Save(a1)) + require.Nil(t, ss.Audit().Save(a1)) time.Sleep(10 * time.Millisecond) a2 := &model.Audit{UserId: a1.UserId, IpAddress: "ipaddress", Action: "Action"} - store.Must(ss.Audit().Save(a2)) + require.Nil(t, ss.Audit().Save(a2)) time.Sleep(10 * time.Millisecond) cutoff := model.GetMillis() time.Sleep(10 * time.Millisecond) a3 := &model.Audit{UserId: a1.UserId, IpAddress: "ipaddress", Action: "Action"} - store.Must(ss.Audit().Save(a3)) + require.Nil(t, ss.Audit().Save(a3)) - if r := <-ss.Audit().Get(a1.UserId, 0, 100); len(r.Data.(model.Audits)) != 3 { - t.Fatal("Expected 3 audits. Got ", len(r.Data.(model.Audits))) - } + audits, err := ss.Audit().Get(a1.UserId, 0, 100) + assert.Len(t, audits, 3) - store.Must(ss.Audit().PermanentDeleteBatch(cutoff, 1000000)) + _, err = ss.Audit().PermanentDeleteBatch(cutoff, 1000000) + require.Nil(t, err) - if r := <-ss.Audit().Get(a1.UserId, 0, 100); len(r.Data.(model.Audits)) != 1 { - t.Fatal("Expected 1 audit. Got ", len(r.Data.(model.Audits))) - } + audits, err = ss.Audit().Get(a1.UserId, 0, 100) + assert.Len(t, audits, 1) - if r2 := <-ss.Audit().PermanentDeleteByUser(a1.UserId); r2.Err != nil { - t.Fatal(r2.Err) - } + require.Nil(t, ss.Audit().PermanentDeleteByUser(a1.UserId)) } diff --git a/store/storetest/channel_store.go b/store/storetest/channel_store.go index c68638ac17..1d72fe12b2 100644 --- a/store/storetest/channel_store.go +++ b/store/storetest/channel_store.go @@ -889,7 +889,7 @@ func testChannelMemberStore(t *testing.T, ss store.Store) { c1t3 := (<-ss.Channel().Get(c1.Id, false)).Data.(*model.Channel) assert.EqualValues(t, 0, c1t3.ExtraUpdateAt, "ExtraUpdateAt should be 0") - member := (<-ss.Channel().GetMember(o1.ChannelId, o1.UserId)).Data.(*model.ChannelMember) + member, _ := ss.Channel().GetMember(o1.ChannelId, o1.UserId) if member.ChannelId != o1.ChannelId { t.Fatal("should have go member") } @@ -1587,12 +1587,14 @@ func testChannelStoreUpdateLastViewedAt(t *testing.T, ss store.Store) { t.Fatal("last viewed at time incorrect") } - rm1 := store.Must(ss.Channel().GetMember(m1.ChannelId, m1.UserId)).(*model.ChannelMember) + rm1, err := ss.Channel().GetMember(m1.ChannelId, m1.UserId) + assert.Nil(t, err) assert.Equal(t, rm1.LastViewedAt, o1.LastPostAt) assert.Equal(t, rm1.LastUpdateAt, o1.LastPostAt) assert.Equal(t, rm1.MsgCount, o1.TotalMsgCount) - rm2 := store.Must(ss.Channel().GetMember(m2.ChannelId, m2.UserId)).(*model.ChannelMember) + rm2, err := ss.Channel().GetMember(m2.ChannelId, m2.UserId) + assert.Nil(t, err) assert.Equal(t, rm2.LastViewedAt, o2.LastPostAt) assert.Equal(t, rm2.LastUpdateAt, o2.LastPostAt) assert.Equal(t, rm2.MsgCount, o2.TotalMsgCount) @@ -1700,25 +1702,25 @@ func testGetMember(t *testing.T, ss store.Store) { } store.Must(ss.Channel().SaveMember(m2)) - if result := <-ss.Channel().GetMember(model.NewId(), userId); result.Err == nil { + if _, err := ss.Channel().GetMember(model.NewId(), userId); err == nil { t.Fatal("should've failed to get member for non-existent channel") } - if result := <-ss.Channel().GetMember(c1.Id, model.NewId()); result.Err == nil { + if _, err := ss.Channel().GetMember(c1.Id, model.NewId()); err == nil { t.Fatal("should've failed to get member for non-existent user") } - if result := <-ss.Channel().GetMember(c1.Id, userId); result.Err != nil { - t.Fatal("shouldn't have errored when getting member", result.Err) - } else if member := result.Data.(*model.ChannelMember); member.ChannelId != c1.Id { + if member, err := ss.Channel().GetMember(c1.Id, userId); err != nil { + t.Fatal("shouldn't have errored when getting member", err) + } else if member.ChannelId != c1.Id { t.Fatal("should've gotten member of channel 1") } else if member.UserId != userId { t.Fatal("should've gotten member for user") } - if result := <-ss.Channel().GetMember(c2.Id, userId); result.Err != nil { - t.Fatal("shouldn't have errored when getting member", result.Err) - } else if member := result.Data.(*model.ChannelMember); member.ChannelId != c2.Id { + if member, err := ss.Channel().GetMember(c2.Id, userId); err != nil { + t.Fatal("shouldn't have errored when getting member", err) + } else if member.ChannelId != c2.Id { t.Fatal("should've gotten member of channel 2") } else if member.UserId != userId { t.Fatal("should've gotten member for user") @@ -2815,23 +2817,20 @@ func testChannelStoreMigrateChannelMembers(t *testing.T, ss store.Store) { ss.Channel().ClearCaches() - res1 := <-ss.Channel().GetMember(cm1.ChannelId, cm1.UserId) - assert.Nil(t, res1.Err) - cm1b := res1.Data.(*model.ChannelMember) + cm1b, err := ss.Channel().GetMember(cm1.ChannelId, cm1.UserId) + assert.Nil(t, err) assert.Equal(t, "", cm1b.ExplicitRoles) assert.True(t, cm1b.SchemeUser) assert.True(t, cm1b.SchemeAdmin) - res2 := <-ss.Channel().GetMember(cm2.ChannelId, cm2.UserId) - assert.Nil(t, res2.Err) - cm2b := res2.Data.(*model.ChannelMember) + cm2b, err := ss.Channel().GetMember(cm2.ChannelId, cm2.UserId) + assert.Nil(t, err) assert.Equal(t, "", cm2b.ExplicitRoles) assert.True(t, cm2b.SchemeUser) assert.False(t, cm2b.SchemeAdmin) - res3 := <-ss.Channel().GetMember(cm3.ChannelId, cm3.UserId) - assert.Nil(t, res3.Err) - cm3b := res3.Data.(*model.ChannelMember) + cm3b, err := ss.Channel().GetMember(cm3.ChannelId, cm3.UserId) + assert.Nil(t, err) assert.Equal(t, "something_else", cm3b.ExplicitRoles) assert.False(t, cm3b.SchemeUser) assert.False(t, cm3b.SchemeAdmin) @@ -2920,21 +2919,21 @@ func testChannelStoreClearAllCustomRoleAssignments(t *testing.T, ss store.Store) require.Nil(t, (<-ss.Channel().ClearAllCustomRoleAssignments()).Err) - r1 := <-ss.Channel().GetMember(m1.ChannelId, m1.UserId) - require.Nil(t, r1.Err) - assert.Equal(t, m1.ExplicitRoles, r1.Data.(*model.ChannelMember).Roles) + member, err := ss.Channel().GetMember(m1.ChannelId, m1.UserId) + require.Nil(t, err) + assert.Equal(t, m1.ExplicitRoles, member.Roles) - r2 := <-ss.Channel().GetMember(m2.ChannelId, m2.UserId) - require.Nil(t, r2.Err) - assert.Equal(t, "channel_user channel_admin", r2.Data.(*model.ChannelMember).Roles) + member, err = ss.Channel().GetMember(m2.ChannelId, m2.UserId) + require.Nil(t, err) + assert.Equal(t, "channel_user channel_admin", member.Roles) - r3 := <-ss.Channel().GetMember(m3.ChannelId, m3.UserId) - require.Nil(t, r3.Err) - assert.Equal(t, m3.ExplicitRoles, r3.Data.(*model.ChannelMember).Roles) + member, err = ss.Channel().GetMember(m3.ChannelId, m3.UserId) + require.Nil(t, err) + assert.Equal(t, m3.ExplicitRoles, member.Roles) - r4 := <-ss.Channel().GetMember(m4.ChannelId, m4.UserId) - require.Nil(t, r4.Err) - assert.Equal(t, "", r4.Data.(*model.ChannelMember).Roles) + member, err = ss.Channel().GetMember(m4.ChannelId, m4.UserId) + require.Nil(t, err) + assert.Equal(t, "", member.Roles) } // testMaterializedPublicChannels tests edge cases involving the triggers and stored procedures diff --git a/store/storetest/mocks/AuditStore.go b/store/storetest/mocks/AuditStore.go index d1ee9082ec..dd23d9c1ae 100644 --- a/store/storetest/mocks/AuditStore.go +++ b/store/storetest/mocks/AuditStore.go @@ -6,7 +6,6 @@ package mocks import mock "github.com/stretchr/testify/mock" import model "github.com/mattermost/mattermost-server/model" -import store "github.com/mattermost/mattermost-server/store" // AuditStore is an autogenerated mock type for the AuditStore type type AuditStore struct { @@ -14,47 +13,63 @@ type AuditStore struct { } // Get provides a mock function with given fields: user_id, offset, limit -func (_m *AuditStore) Get(user_id string, offset int, limit int) store.StoreChannel { +func (_m *AuditStore) Get(user_id string, offset int, limit int) (model.Audits, *model.AppError) { ret := _m.Called(user_id, offset, limit) - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func(string, int, int) store.StoreChannel); ok { + var r0 model.Audits + if rf, ok := ret.Get(0).(func(string, int, int) model.Audits); ok { r0 = rf(user_id, offset, limit) } else { if ret.Get(0) != nil { - r0 = ret.Get(0).(store.StoreChannel) + r0 = ret.Get(0).(model.Audits) } } - return r0 + var r1 *model.AppError + if rf, ok := ret.Get(1).(func(string, int, int) *model.AppError); ok { + r1 = rf(user_id, offset, limit) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 } // PermanentDeleteBatch provides a mock function with given fields: endTime, limit -func (_m *AuditStore) PermanentDeleteBatch(endTime int64, limit int64) store.StoreChannel { +func (_m *AuditStore) PermanentDeleteBatch(endTime int64, limit int64) (int64, *model.AppError) { ret := _m.Called(endTime, limit) - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func(int64, int64) store.StoreChannel); ok { + var r0 int64 + if rf, ok := ret.Get(0).(func(int64, int64) int64); ok { r0 = rf(endTime, limit) } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(store.StoreChannel) + r0 = ret.Get(0).(int64) + } + + var r1 *model.AppError + if rf, ok := ret.Get(1).(func(int64, int64) *model.AppError); ok { + r1 = rf(endTime, limit) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) } } - return r0 + return r0, r1 } // PermanentDeleteByUser provides a mock function with given fields: userId -func (_m *AuditStore) PermanentDeleteByUser(userId string) store.StoreChannel { +func (_m *AuditStore) PermanentDeleteByUser(userId string) *model.AppError { ret := _m.Called(userId) - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func(string) store.StoreChannel); ok { + var r0 *model.AppError + if rf, ok := ret.Get(0).(func(string) *model.AppError); ok { r0 = rf(userId) } else { if ret.Get(0) != nil { - r0 = ret.Get(0).(store.StoreChannel) + r0 = ret.Get(0).(*model.AppError) } } @@ -62,15 +77,15 @@ func (_m *AuditStore) PermanentDeleteByUser(userId string) store.StoreChannel { } // Save provides a mock function with given fields: audit -func (_m *AuditStore) Save(audit *model.Audit) store.StoreChannel { +func (_m *AuditStore) Save(audit *model.Audit) *model.AppError { ret := _m.Called(audit) - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func(*model.Audit) store.StoreChannel); ok { + var r0 *model.AppError + if rf, ok := ret.Get(0).(func(*model.Audit) *model.AppError); ok { r0 = rf(audit) } else { if ret.Get(0) != nil { - r0 = ret.Get(0).(store.StoreChannel) + r0 = ret.Get(0).(*model.AppError) } } diff --git a/store/storetest/mocks/ChannelStore.go b/store/storetest/mocks/ChannelStore.go index ee32173052..ad7b3cef98 100644 --- a/store/storetest/mocks/ChannelStore.go +++ b/store/storetest/mocks/ChannelStore.go @@ -483,19 +483,28 @@ func (_m *ChannelStore) GetFromMaster(id string) store.StoreChannel { } // GetMember provides a mock function with given fields: channelId, userId -func (_m *ChannelStore) GetMember(channelId string, userId string) store.StoreChannel { +func (_m *ChannelStore) GetMember(channelId string, userId string) (*model.ChannelMember, *model.AppError) { ret := _m.Called(channelId, userId) - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func(string, string) store.StoreChannel); ok { + var r0 *model.ChannelMember + if rf, ok := ret.Get(0).(func(string, string) *model.ChannelMember); ok { r0 = rf(channelId, userId) } else { if ret.Get(0) != nil { - r0 = ret.Get(0).(store.StoreChannel) + r0 = ret.Get(0).(*model.ChannelMember) } } - return r0 + var r1 *model.AppError + if rf, ok := ret.Get(1).(func(string, string) *model.AppError); ok { + r1 = rf(channelId, userId) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 } // GetMemberCount provides a mock function with given fields: channelId, allowFromCache diff --git a/store/storetest/mocks/GroupStore.go b/store/storetest/mocks/GroupStore.go index b66202bf4b..f352ed1642 100644 --- a/store/storetest/mocks/GroupStore.go +++ b/store/storetest/mocks/GroupStore.go @@ -13,6 +13,38 @@ type GroupStore struct { mock.Mock } +// ChannelMembersToAdd provides a mock function with given fields: since +func (_m *GroupStore) ChannelMembersToAdd(since int64) store.StoreChannel { + ret := _m.Called(since) + + var r0 store.StoreChannel + if rf, ok := ret.Get(0).(func(int64) store.StoreChannel); ok { + r0 = rf(since) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(store.StoreChannel) + } + } + + return r0 +} + +// ChannelMembersToRemove provides a mock function with given fields: +func (_m *GroupStore) ChannelMembersToRemove() store.StoreChannel { + ret := _m.Called() + + var r0 store.StoreChannel + if rf, ok := ret.Get(0).(func() store.StoreChannel); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(store.StoreChannel) + } + } + + return r0 +} + // Create provides a mock function with given fields: group func (_m *GroupStore) Create(group *model.Group) store.StoreChannel { ret := _m.Called(group) @@ -269,22 +301,6 @@ func (_m *GroupStore) GetMemberUsersPage(groupID string, offset int, limit int) return r0 } -// ChannelMembersToAdd provides a mock function with given fields: since -func (_m *GroupStore) ChannelMembersToAdd(since int64) store.StoreChannel { - ret := _m.Called(since) - - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func(int64) store.StoreChannel); ok { - r0 = rf(since) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(store.StoreChannel) - } - } - - return r0 -} - // TeamMembersToAdd provides a mock function with given fields: since func (_m *GroupStore) TeamMembersToAdd(since int64) store.StoreChannel { ret := _m.Called(since) @@ -301,22 +317,6 @@ func (_m *GroupStore) TeamMembersToAdd(since int64) store.StoreChannel { return r0 } -// ChannelMembersToRemove provides a mock function with given fields: -func (_m *GroupStore) ChannelMembersToRemove() store.StoreChannel { - ret := _m.Called() - - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func() store.StoreChannel); ok { - r0 = rf() - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(store.StoreChannel) - } - } - - return r0 -} - // TeamMembersToRemove provides a mock function with given fields: func (_m *GroupStore) TeamMembersToRemove() store.StoreChannel { ret := _m.Called() diff --git a/store/storetest/mocks/LayeredStoreDatabaseLayer.go b/store/storetest/mocks/LayeredStoreDatabaseLayer.go index fe803cbc2c..9d700539c9 100644 --- a/store/storetest/mocks/LayeredStoreDatabaseLayer.go +++ b/store/storetest/mocks/LayeredStoreDatabaseLayer.go @@ -78,6 +78,52 @@ func (_m *LayeredStoreDatabaseLayer) ChannelMemberHistory() store.ChannelMemberH return r0 } +// ChannelMembersToAdd provides a mock function with given fields: ctx, since, hints +func (_m *LayeredStoreDatabaseLayer) ChannelMembersToAdd(ctx context.Context, since int64, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { + _va := make([]interface{}, len(hints)) + for _i := range hints { + _va[_i] = hints[_i] + } + var _ca []interface{} + _ca = append(_ca, ctx, since) + _ca = append(_ca, _va...) + ret := _m.Called(_ca...) + + var r0 *store.LayeredStoreSupplierResult + if rf, ok := ret.Get(0).(func(context.Context, int64, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { + r0 = rf(ctx, since, hints...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) + } + } + + return r0 +} + +// ChannelMembersToRemove provides a mock function with given fields: ctx, hints +func (_m *LayeredStoreDatabaseLayer) ChannelMembersToRemove(ctx context.Context, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { + _va := make([]interface{}, len(hints)) + for _i := range hints { + _va[_i] = hints[_i] + } + var _ca []interface{} + _ca = append(_ca, ctx) + _ca = append(_ca, _va...) + ret := _m.Called(_ca...) + + var r0 *store.LayeredStoreSupplierResult + if rf, ok := ret.Get(0).(func(context.Context, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { + r0 = rf(ctx, hints...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) + } + } + + return r0 +} + // Close provides a mock function with given fields: func (_m *LayeredStoreDatabaseLayer) Close() { _m.Called() @@ -704,98 +750,6 @@ func (_m *LayeredStoreDatabaseLayer) OAuth() store.OAuthStore { return r0 } -// ChannelMembersToAdd provides a mock function with given fields: ctx, since, hints -func (_m *LayeredStoreDatabaseLayer) ChannelMembersToAdd(ctx context.Context, since int64, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { - _va := make([]interface{}, len(hints)) - for _i := range hints { - _va[_i] = hints[_i] - } - var _ca []interface{} - _ca = append(_ca, ctx, since) - _ca = append(_ca, _va...) - ret := _m.Called(_ca...) - - var r0 *store.LayeredStoreSupplierResult - if rf, ok := ret.Get(0).(func(context.Context, int64, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { - r0 = rf(ctx, since, hints...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) - } - } - - return r0 -} - -// TeamMembersToAdd provides a mock function with given fields: ctx, since, hints -func (_m *LayeredStoreDatabaseLayer) TeamMembersToAdd(ctx context.Context, since int64, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { - _va := make([]interface{}, len(hints)) - for _i := range hints { - _va[_i] = hints[_i] - } - var _ca []interface{} - _ca = append(_ca, ctx, since) - _ca = append(_ca, _va...) - ret := _m.Called(_ca...) - - var r0 *store.LayeredStoreSupplierResult - if rf, ok := ret.Get(0).(func(context.Context, int64, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { - r0 = rf(ctx, since, hints...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) - } - } - - return r0 -} - -// ChannelMembersToRemove provides a mock function with given fields: ctx, hints -func (_m *LayeredStoreDatabaseLayer) ChannelMembersToRemove(ctx context.Context, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { - _va := make([]interface{}, len(hints)) - for _i := range hints { - _va[_i] = hints[_i] - } - var _ca []interface{} - _ca = append(_ca, ctx) - _ca = append(_ca, _va...) - ret := _m.Called(_ca...) - - var r0 *store.LayeredStoreSupplierResult - if rf, ok := ret.Get(0).(func(context.Context, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { - r0 = rf(ctx, hints...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) - } - } - - return r0 -} - -// TeamMembersToRemove provides a mock function with given fields: ctx, hints -func (_m *LayeredStoreDatabaseLayer) TeamMembersToRemove(ctx context.Context, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { - _va := make([]interface{}, len(hints)) - for _i := range hints { - _va[_i] = hints[_i] - } - var _ca []interface{} - _ca = append(_ca, ctx) - _ca = append(_ca, _va...) - ret := _m.Called(_ca...) - - var r0 *store.LayeredStoreSupplierResult - if rf, ok := ret.Get(0).(func(context.Context, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { - r0 = rf(ctx, hints...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) - } - } - - return r0 -} - // Plugin provides a mock function with given fields: func (_m *LayeredStoreDatabaseLayer) Plugin() store.PluginStore { ret := _m.Called() @@ -1398,6 +1352,52 @@ func (_m *LayeredStoreDatabaseLayer) Team() store.TeamStore { return r0 } +// TeamMembersToAdd provides a mock function with given fields: ctx, since, hints +func (_m *LayeredStoreDatabaseLayer) TeamMembersToAdd(ctx context.Context, since int64, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { + _va := make([]interface{}, len(hints)) + for _i := range hints { + _va[_i] = hints[_i] + } + var _ca []interface{} + _ca = append(_ca, ctx, since) + _ca = append(_ca, _va...) + ret := _m.Called(_ca...) + + var r0 *store.LayeredStoreSupplierResult + if rf, ok := ret.Get(0).(func(context.Context, int64, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { + r0 = rf(ctx, since, hints...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) + } + } + + return r0 +} + +// TeamMembersToRemove provides a mock function with given fields: ctx, hints +func (_m *LayeredStoreDatabaseLayer) TeamMembersToRemove(ctx context.Context, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { + _va := make([]interface{}, len(hints)) + for _i := range hints { + _va[_i] = hints[_i] + } + var _ca []interface{} + _ca = append(_ca, ctx) + _ca = append(_ca, _va...) + ret := _m.Called(_ca...) + + var r0 *store.LayeredStoreSupplierResult + if rf, ok := ret.Get(0).(func(context.Context, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { + r0 = rf(ctx, hints...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) + } + } + + return r0 +} + // TermsOfService provides a mock function with given fields: func (_m *LayeredStoreDatabaseLayer) TermsOfService() store.TermsOfServiceStore { ret := _m.Called() diff --git a/store/storetest/mocks/LayeredStoreSupplier.go b/store/storetest/mocks/LayeredStoreSupplier.go index 070f441cf3..78378f44c4 100644 --- a/store/storetest/mocks/LayeredStoreSupplier.go +++ b/store/storetest/mocks/LayeredStoreSupplier.go @@ -14,6 +14,52 @@ type LayeredStoreSupplier struct { mock.Mock } +// ChannelMembersToAdd provides a mock function with given fields: ctx, since, hints +func (_m *LayeredStoreSupplier) ChannelMembersToAdd(ctx context.Context, since int64, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { + _va := make([]interface{}, len(hints)) + for _i := range hints { + _va[_i] = hints[_i] + } + var _ca []interface{} + _ca = append(_ca, ctx, since) + _ca = append(_ca, _va...) + ret := _m.Called(_ca...) + + var r0 *store.LayeredStoreSupplierResult + if rf, ok := ret.Get(0).(func(context.Context, int64, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { + r0 = rf(ctx, since, hints...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) + } + } + + return r0 +} + +// ChannelMembersToRemove provides a mock function with given fields: ctx, hints +func (_m *LayeredStoreSupplier) ChannelMembersToRemove(ctx context.Context, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { + _va := make([]interface{}, len(hints)) + for _i := range hints { + _va[_i] = hints[_i] + } + var _ca []interface{} + _ca = append(_ca, ctx) + _ca = append(_ca, _va...) + ret := _m.Called(_ca...) + + var r0 *store.LayeredStoreSupplierResult + if rf, ok := ret.Get(0).(func(context.Context, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { + r0 = rf(ctx, hints...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) + } + } + + return r0 +} + // GetGroupsByChannel provides a mock function with given fields: ctx, channelId, page, perPage, hints func (_m *LayeredStoreSupplier) GetGroupsByChannel(ctx context.Context, channelId string, page int, perPage int, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { _va := make([]interface{}, len(hints)) @@ -444,98 +490,6 @@ func (_m *LayeredStoreSupplier) Next() store.LayeredStoreSupplier { return r0 } -// ChannelMembersToAdd provides a mock function with given fields: ctx, since, hints -func (_m *LayeredStoreSupplier) ChannelMembersToAdd(ctx context.Context, since int64, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { - _va := make([]interface{}, len(hints)) - for _i := range hints { - _va[_i] = hints[_i] - } - var _ca []interface{} - _ca = append(_ca, ctx, since) - _ca = append(_ca, _va...) - ret := _m.Called(_ca...) - - var r0 *store.LayeredStoreSupplierResult - if rf, ok := ret.Get(0).(func(context.Context, int64, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { - r0 = rf(ctx, since, hints...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) - } - } - - return r0 -} - -// TeamMembersToAdd provides a mock function with given fields: ctx, since, hints -func (_m *LayeredStoreSupplier) TeamMembersToAdd(ctx context.Context, since int64, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { - _va := make([]interface{}, len(hints)) - for _i := range hints { - _va[_i] = hints[_i] - } - var _ca []interface{} - _ca = append(_ca, ctx, since) - _ca = append(_ca, _va...) - ret := _m.Called(_ca...) - - var r0 *store.LayeredStoreSupplierResult - if rf, ok := ret.Get(0).(func(context.Context, int64, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { - r0 = rf(ctx, since, hints...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) - } - } - - return r0 -} - -// ChannelMembersToRemove provides a mock function with given fields: ctx, hints -func (_m *LayeredStoreSupplier) ChannelMembersToRemove(ctx context.Context, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { - _va := make([]interface{}, len(hints)) - for _i := range hints { - _va[_i] = hints[_i] - } - var _ca []interface{} - _ca = append(_ca, ctx) - _ca = append(_ca, _va...) - ret := _m.Called(_ca...) - - var r0 *store.LayeredStoreSupplierResult - if rf, ok := ret.Get(0).(func(context.Context, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { - r0 = rf(ctx, hints...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) - } - } - - return r0 -} - -// TeamMembersToRemove provides a mock function with given fields: ctx, hints -func (_m *LayeredStoreSupplier) TeamMembersToRemove(ctx context.Context, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { - _va := make([]interface{}, len(hints)) - for _i := range hints { - _va[_i] = hints[_i] - } - var _ca []interface{} - _ca = append(_ca, ctx) - _ca = append(_ca, _va...) - ret := _m.Called(_ca...) - - var r0 *store.LayeredStoreSupplierResult - if rf, ok := ret.Get(0).(func(context.Context, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { - r0 = rf(ctx, hints...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) - } - } - - return r0 -} - // ReactionDelete provides a mock function with given fields: ctx, reaction, hints func (_m *LayeredStoreSupplier) ReactionDelete(ctx context.Context, reaction *model.Reaction, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { _va := make([]interface{}, len(hints)) @@ -977,3 +931,49 @@ func (_m *LayeredStoreSupplier) SchemeSave(ctx context.Context, scheme *model.Sc func (_m *LayeredStoreSupplier) SetChainNext(_a0 store.LayeredStoreSupplier) { _m.Called(_a0) } + +// TeamMembersToAdd provides a mock function with given fields: ctx, since, hints +func (_m *LayeredStoreSupplier) TeamMembersToAdd(ctx context.Context, since int64, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { + _va := make([]interface{}, len(hints)) + for _i := range hints { + _va[_i] = hints[_i] + } + var _ca []interface{} + _ca = append(_ca, ctx, since) + _ca = append(_ca, _va...) + ret := _m.Called(_ca...) + + var r0 *store.LayeredStoreSupplierResult + if rf, ok := ret.Get(0).(func(context.Context, int64, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { + r0 = rf(ctx, since, hints...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) + } + } + + return r0 +} + +// TeamMembersToRemove provides a mock function with given fields: ctx, hints +func (_m *LayeredStoreSupplier) TeamMembersToRemove(ctx context.Context, hints ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult { + _va := make([]interface{}, len(hints)) + for _i := range hints { + _va[_i] = hints[_i] + } + var _ca []interface{} + _ca = append(_ca, ctx) + _ca = append(_ca, _va...) + ret := _m.Called(_ca...) + + var r0 *store.LayeredStoreSupplierResult + if rf, ok := ret.Get(0).(func(context.Context, ...store.LayeredStoreHint) *store.LayeredStoreSupplierResult); ok { + r0 = rf(ctx, hints...) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*store.LayeredStoreSupplierResult) + } + } + + return r0 +} diff --git a/store/storetest/mocks/UserStore.go b/store/storetest/mocks/UserStore.go index 50451dd02c..2a8c24f336 100644 --- a/store/storetest/mocks/UserStore.go +++ b/store/storetest/mocks/UserStore.go @@ -99,19 +99,28 @@ func (_m *UserStore) Count(options model.UserCountOptions) store.StoreChannel { } // Get provides a mock function with given fields: id -func (_m *UserStore) Get(id string) store.StoreChannel { +func (_m *UserStore) Get(id string) (*model.User, *model.AppError) { ret := _m.Called(id) - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func(string) store.StoreChannel); ok { + var r0 *model.User + if rf, ok := ret.Get(0).(func(string) *model.User); ok { r0 = rf(id) } else { if ret.Get(0) != nil { - r0 = ret.Get(0).(store.StoreChannel) + r0 = ret.Get(0).(*model.User) } } - return r0 + var r1 *model.AppError + if rf, ok := ret.Get(1).(func(string) *model.AppError); ok { + r1 = rf(id) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 } // GetAll provides a mock function with given fields: diff --git a/store/storetest/team_store.go b/store/storetest/team_store.go index 29bf72fa24..efffefda77 100644 --- a/store/storetest/team_store.go +++ b/store/storetest/team_store.go @@ -1012,7 +1012,8 @@ func testSaveTeamMemberMaxMembers(t *testing.T, ss store.Store) { } // Deactivating a user should make them stop counting against max members - user2 := store.Must(ss.User().Get(userIds[1])).(*model.User) + user2, err := ss.User().Get(userIds[1]) + require.Nil(t, err) user2.DeleteAt = 1234 store.Must(ss.User().Update(user2, true)) diff --git a/store/storetest/user_store.go b/store/storetest/user_store.go index 6baac34e5e..bba10c178b 100644 --- a/store/storetest/user_store.go +++ b/store/storetest/user_store.go @@ -220,12 +220,10 @@ func testUserStoreUpdateUpdateAt(t *testing.T, ss store.Store) { t.Fatal(err) } - if r1 := <-ss.User().Get(u1.Id); r1.Err != nil { - t.Fatal(r1.Err) - } else { - if r1.Data.(*model.User).UpdateAt <= u1.UpdateAt { - t.Fatal("UpdateAt not updated correctly") - } + user, err := ss.User().Get(u1.Id) + require.Nil(t, err) + if user.UpdateAt <= u1.UpdateAt { + t.Fatal("UpdateAt not updated correctly") } } @@ -241,14 +239,11 @@ func testUserStoreUpdateFailedPasswordAttempts(t *testing.T, ss store.Store) { t.Fatal(err) } - if r1 := <-ss.User().Get(u1.Id); r1.Err != nil { - t.Fatal(r1.Err) - } else { - if r1.Data.(*model.User).FailedAttempts != 3 { - t.Fatal("FailedAttempts not updated correctly") - } + user, err := ss.User().Get(u1.Id) + require.Nil(t, err) + if user.FailedAttempts != 3 { + t.Fatal("FailedAttempts not updated correctly") } - } func testUserStoreGet(t *testing.T, ss store.Store) { @@ -274,23 +269,20 @@ func testUserStoreGet(t *testing.T, ss store.Store) { store.Must(ss.Team().SaveMember(&model.TeamMember{TeamId: model.NewId(), UserId: u1.Id}, -1)) t.Run("fetch empty id", func(t *testing.T) { - require.NotNil(t, (<-ss.User().Get("")).Err) + _, err := ss.User().Get("") + require.NotNil(t, err) }) t.Run("fetch user 1", func(t *testing.T) { - result := <-ss.User().Get(u1.Id) - require.Nil(t, result.Err) - - actual := result.Data.(*model.User) + actual, err := ss.User().Get(u1.Id) + require.Nil(t, err) require.Equal(t, u1, actual) require.False(t, actual.IsBot) }) t.Run("fetch user 2, also a bot", func(t *testing.T) { - result := <-ss.User().Get(u2.Id) - require.Nil(t, result.Err) - - actual := result.Data.(*model.User) + actual, err := ss.User().Get(u2.Id) + require.Nil(t, err) require.Equal(t, u2, actual) require.True(t, actual.IsBot) }) diff --git a/web/context.go b/web/context.go index 174674ae13..1d56eb5547 100644 --- a/web/context.go +++ b/web/context.go @@ -25,8 +25,8 @@ type Context struct { func (c *Context) LogAudit(extraInfo string) { audit := &model.Audit{UserId: c.App.Session.UserId, IpAddress: c.App.IpAddress, Action: c.App.Path, ExtraInfo: extraInfo, SessionId: c.App.Session.Id} - if r := <-c.App.Srv.Store.Audit().Save(audit); r.Err != nil { - c.LogError(r.Err) + if err := c.App.Srv.Store.Audit().Save(audit); err != nil { + c.LogError(err) } } @@ -37,8 +37,8 @@ func (c *Context) LogAuditWithUserId(userId, extraInfo string) { } audit := &model.Audit{UserId: userId, IpAddress: c.App.IpAddress, Action: c.App.Path, ExtraInfo: extraInfo, SessionId: c.App.Session.Id} - if r := <-c.App.Srv.Store.Audit().Save(audit); r.Err != nil { - c.LogError(r.Err) + if err := c.App.Srv.Store.Audit().Save(audit); err != nil { + c.LogError(err) } }