diff --git a/api4/apitestlib.go b/api4/apitestlib.go index f42cb76c92..5cbe9399a1 100644 --- a/api4/apitestlib.go +++ b/api4/apitestlib.go @@ -750,9 +750,9 @@ func (me *TestHelper) MakeUserChannelAdmin(user *model.User, channel *model.Chan 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 { + if _, err = me.App.Srv.Store.Channel().UpdateMember(cm); err != nil { utils.EnableDebugLogForTest() - panic(sr.Err) + panic(err) } } else { utils.EnableDebugLogForTest() diff --git a/app/channel.go b/app/channel.go index 5da2a65d18..83193e5be4 100644 --- a/app/channel.go +++ b/app/channel.go @@ -679,7 +679,8 @@ func (a *App) UpdateChannelMemberRoles(channelId string, userId string, newRoles member.SchemeAdmin = false for _, roleName := range strings.Fields(newRoles) { - role, err := a.GetRoleByName(roleName) + var role *model.Role + role, err = a.GetRoleByName(roleName) if err != nil { err.StatusCode = http.StatusBadRequest return nil, err @@ -714,11 +715,10 @@ func (a *App) UpdateChannelMemberRoles(channelId string, userId string, newRoles member.ExplicitRoles = strings.Join(newExplicitRoles, " ") - result := <-a.Srv.Store.Channel().UpdateMember(member) - if result.Err != nil { - return nil, result.Err + member, err = a.Srv.Store.Channel().UpdateMember(member) + if err != nil { + return nil, err } - member = result.Data.(*model.ChannelMember) a.InvalidateCacheForUser(userId) return member, nil @@ -743,11 +743,10 @@ func (a *App) UpdateChannelMemberSchemeRoles(channelId string, userId string, is member.ExplicitRoles = RemoveRoles([]string{model.CHANNEL_GUEST_ROLE_ID, model.CHANNEL_USER_ROLE_ID, model.CHANNEL_ADMIN_ROLE_ID}, member.ExplicitRoles) } - result := <-a.Srv.Store.Channel().UpdateMember(member) - if result.Err != nil { - return nil, result.Err + member, err = a.Srv.Store.Channel().UpdateMember(member) + if err != nil { + return nil, err } - member = result.Data.(*model.ChannelMember) a.InvalidateCacheForUser(userId) return member, nil @@ -781,9 +780,9 @@ func (a *App) UpdateChannelMemberNotifyProps(data map[string]string, channelId s member.NotifyProps[model.IGNORE_CHANNEL_MENTIONS_NOTIFY_PROP] = ignoreChannelMentions } - result := <-a.Srv.Store.Channel().UpdateMember(member) - if result.Err != nil { - return nil, result.Err + member, err = a.Srv.Store.Channel().UpdateMember(member) + if err != nil { + return nil, err } a.InvalidateCacheForUser(userId) @@ -2023,7 +2022,7 @@ func (a *App) ToggleMuteChannel(channelId string, userId string) *model.ChannelM member.NotifyProps[model.MARK_UNREAD_NOTIFY_PROP] = model.CHANNEL_NOTIFY_MENTION } - <-a.Srv.Store.Channel().UpdateMember(member) + a.Srv.Store.Channel().UpdateMember(member) return member } diff --git a/app/email_batching_test.go b/app/email_batching_test.go index a1e7c91e41..a6ffff478b 100644 --- a/app/email_batching_test.go +++ b/app/email_batching_test.go @@ -9,7 +9,6 @@ import ( "time" "github.com/mattermost/mattermost-server/model" - "github.com/mattermost/mattermost-server/store" "github.com/stretchr/testify/require" ) @@ -113,7 +112,8 @@ func TestCheckPendingNotifications(t *testing.T) { 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)) + _, err = th.App.Srv.Store.Channel().UpdateMember(channelMember) + require.Nil(t, err) err = th.App.Srv.Store.Preference().Save(&model.Preferences{{ UserId: th.BasicUser.Id, @@ -134,7 +134,8 @@ func TestCheckPendingNotifications(t *testing.T) { 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)) + _, err = th.App.Srv.Store.Channel().UpdateMember(channelMember) + require.Nil(t, err) job.checkPendingNotifications(time.Unix(10002, 0), func(string, []*batchedNotification) {}) @@ -215,7 +216,8 @@ func TestCheckPendingNotificationsDefaultInterval(t *testing.T) { 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)) + _, err = th.App.Srv.Store.Channel().UpdateMember(channelMember) + require.Nil(t, err) job.pendingNotifications[th.BasicUser.Id] = []*batchedNotification{ { @@ -254,7 +256,8 @@ func TestCheckPendingNotificationsCantParseInterval(t *testing.T) { 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)) + _, err = th.App.Srv.Store.Channel().UpdateMember(channelMember) + require.Nil(t, err) // preference value is not an integer, so we'll fall back to the default 15min value err = th.App.Srv.Store.Preference().Save(&model.Preferences{{ diff --git a/store/sqlstore/channel_store.go b/store/sqlstore/channel_store.go index c2ccd54bc1..f6146a2a4b 100644 --- a/store/sqlstore/channel_store.go +++ b/store/sqlstore/channel_store.go @@ -1368,31 +1368,26 @@ func (s SqlChannelStore) saveMemberT(transaction *gorp.Transaction, member *mode return result } -func (s SqlChannelStore) UpdateMember(member *model.ChannelMember) store.StoreChannel { - return store.Do(func(result *store.StoreResult) { - member.PreUpdate() +func (s SqlChannelStore) UpdateMember(member *model.ChannelMember) (*model.ChannelMember, *model.AppError) { + member.PreUpdate() - if result.Err = member.IsValid(); result.Err != nil { - return + if err := member.IsValid(); err != nil { + return nil, err + } + + if _, err := s.GetMaster().Update(NewChannelMemberFromModel(member)); err != nil { + return nil, model.NewAppError("SqlChannelStore.UpdateMember", "store.sql_channel.update_member.app_error", nil, "channel_id="+member.ChannelId+", "+"user_id="+member.UserId+", "+err.Error(), http.StatusInternalServerError) + } + + 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": member.ChannelId, "UserId": member.UserId}); err != nil { + if err == sql.ErrNoRows { + return nil, model.NewAppError("SqlChannelStore.GetMember", store.MISSING_CHANNEL_MEMBER_ERROR, nil, "channel_id="+member.ChannelId+"user_id="+member.UserId+","+err.Error(), http.StatusNotFound) } - - if _, err := s.GetMaster().Update(NewChannelMemberFromModel(member)); err != nil { - result.Err = model.NewAppError("SqlChannelStore.UpdateMember", "store.sql_channel.update_member.app_error", nil, "channel_id="+member.ChannelId+", "+"user_id="+member.UserId+", "+err.Error(), http.StatusInternalServerError) - return - } - - 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": member.ChannelId, "UserId": member.UserId}); err != nil { - if err == sql.ErrNoRows { - result.Err = model.NewAppError("SqlChannelStore.GetMember", store.MISSING_CHANNEL_MEMBER_ERROR, nil, "channel_id="+member.ChannelId+"user_id="+member.UserId+","+err.Error(), http.StatusNotFound) - return - } - result.Err = model.NewAppError("SqlChannelStore.GetMember", "store.sql_channel.get_member.app_error", nil, "channel_id="+member.ChannelId+"user_id="+member.UserId+","+err.Error(), http.StatusInternalServerError) - return - } - result.Data = dbMember.ToModel() - }) + return nil, model.NewAppError("SqlChannelStore.GetMember", "store.sql_channel.get_member.app_error", nil, "channel_id="+member.ChannelId+"user_id="+member.UserId+","+err.Error(), http.StatusInternalServerError) + } + return dbMember.ToModel(), nil } func (s SqlChannelStore) GetMembers(channelId string, offset, limit int) (*model.ChannelMembers, *model.AppError) { diff --git a/store/store.go b/store/store.go index d00bcbfe49..f391eed420 100644 --- a/store/store.go +++ b/store/store.go @@ -159,7 +159,7 @@ type ChannelStore interface { GetChannelsByIds(channelIds []string) ([]*model.Channel, *model.AppError) GetForPost(postId string) (*model.Channel, *model.AppError) SaveMember(member *model.ChannelMember) StoreChannel - UpdateMember(member *model.ChannelMember) StoreChannel + UpdateMember(member *model.ChannelMember) (*model.ChannelMember, *model.AppError) GetMembers(channelId string, offset, limit int) (*model.ChannelMembers, *model.AppError) GetMember(channelId string, userId string) (*model.ChannelMember, *model.AppError) GetChannelMembersTimezones(channelId string) ([]model.StringMap, *model.AppError) diff --git a/store/storetest/channel_store.go b/store/storetest/channel_store.go index 2cc309e4e9..e4845fa058 100644 --- a/store/storetest/channel_store.go +++ b/store/storetest/channel_store.go @@ -1771,12 +1771,12 @@ func testUpdateChannelMember(t *testing.T, ss store.Store) { store.Must(ss.Channel().SaveMember(m1)) m1.NotifyProps["test"] = "sometext" - if result := <-ss.Channel().UpdateMember(m1); result.Err != nil { - t.Fatal(result.Err) + if _, err := ss.Channel().UpdateMember(m1); err != nil { + t.Fatal(err) } m1.UserId = "" - if result := <-ss.Channel().UpdateMember(m1); result.Err == nil { + if _, err := ss.Channel().UpdateMember(m1); err == nil { t.Fatal("bad user id - should fail") } } diff --git a/store/storetest/mocks/ChannelStore.go b/store/storetest/mocks/ChannelStore.go index 9bd01328d6..37dc1e35f6 100644 --- a/store/storetest/mocks/ChannelStore.go +++ b/store/storetest/mocks/ChannelStore.go @@ -1493,19 +1493,28 @@ func (_m *ChannelStore) UpdateLastViewedAt(channelIds []string, userId string) ( } // UpdateMember provides a mock function with given fields: member -func (_m *ChannelStore) UpdateMember(member *model.ChannelMember) store.StoreChannel { +func (_m *ChannelStore) UpdateMember(member *model.ChannelMember) (*model.ChannelMember, *model.AppError) { ret := _m.Called(member) - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func(*model.ChannelMember) store.StoreChannel); ok { + var r0 *model.ChannelMember + if rf, ok := ret.Get(0).(func(*model.ChannelMember) *model.ChannelMember); ok { r0 = rf(member) } 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(*model.ChannelMember) *model.AppError); ok { + r1 = rf(member) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 } // UserBelongsToChannels provides a mock function with given fields: userId, channelIds