diff --git a/app/channel.go b/app/channel.go index 424d26241b..b90db2b1f6 100644 --- a/app/channel.go +++ b/app/channel.go @@ -2482,11 +2482,28 @@ func (a *App) MarkChannelAsUnreadFromPost(postID string, userID string, collapse threadId = post.Id } - threadMembership, _ := a.Srv().Store.Thread().GetMembershipForUser(user.Id, threadId) + var nfErr *store.ErrNotFound + threadMembership, storeErr := a.Srv().Store.Thread().GetMembershipForUser(user.Id, threadId) + if storeErr != nil && !errors.As(storeErr, &nfErr) { + return nil, model.NewAppError("MarkChannelAsUnreadFromPost", "app.channel.update_last_viewed_at_post.app_error", nil, storeErr.Error(), http.StatusInternalServerError) + } // if this post was not followed before, create thread membership and update mention count if threadMembership == nil { - threadMembership, _ = a.Srv().Store.Thread().MaintainMembership(user.Id, threadId, true, true, true, true, false) - threadData, _ := a.Srv().Store.Thread().Get(threadId) + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: true, + UpdateFollowing: true, + UpdateViewedTimestamp: true, + UpdateParticipants: false, + } + threadMembership, storeErr = a.Srv().Store.Thread().MaintainMembership(user.Id, threadId, opts) + if storeErr != nil && !errors.As(storeErr, &nfErr) { + return nil, model.NewAppError("MarkChannelAsUnreadFromPost", "app.channel.update_last_viewed_at_post.app_error", nil, storeErr.Error(), http.StatusInternalServerError) + } + threadData, storeErr := a.Srv().Store.Thread().Get(threadId) + if storeErr != nil && !errors.As(storeErr, &nfErr) { + return nil, model.NewAppError("MarkChannelAsUnreadFromPost", "app.channel.update_last_viewed_at_post.app_error", nil, storeErr.Error(), http.StatusInternalServerError) + } if threadData != nil && threadMembership != nil && threadMembership.Following { channel, nErr := a.Srv().Store.Channel().Get(post.ChannelId, true) if nErr != nil { @@ -2500,7 +2517,10 @@ func (a *App) MarkChannelAsUnreadFromPost(postID string, userID string, collapse if nErr != nil { return nil, model.NewAppError("MarkChannelAsUnreadFromPost", "app.channel.update_last_viewed_at_post.app_error", nil, nErr.Error(), http.StatusInternalServerError) } - thread, _ := a.Srv().Store.Thread().GetThreadForUser(channel.TeamId, threadMembership, true) + thread, nErr := a.Srv().Store.Thread().GetThreadForUser(channel.TeamId, threadMembership, true) + if nErr != nil { + return nil, model.NewAppError("MarkChannelAsUnreadFromPost", "app.channel.update_last_viewed_at_post.app_error", nil, nErr.Error(), http.StatusInternalServerError) + } a.sanitizeProfiles(thread.Participants, false) thread.Post.SanitizeProps() @@ -2586,7 +2606,14 @@ func (a *App) markChannelAsUnreadFromPostCRTUnsupported(postID string, userID st } // Follow thread if we're not already following it if threadMembership == nil { - threadMembership, nErr = a.Srv().Store.Thread().MaintainMembership(user.Id, threadId, true, false, true, false, false) + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + threadMembership, nErr = a.Srv().Store.Thread().MaintainMembership(user.Id, threadId, opts) if nErr != nil { return nil, model.NewAppError("MarkChannelAsUnreadFromPost", "app.channel.update_last_viewed_at_post.app_error", nil, nErr.Error(), http.StatusInternalServerError) } diff --git a/app/notification.go b/app/notification.go index 8bb90974d6..29f10fa1c0 100644 --- a/app/notification.go +++ b/app/notification.go @@ -206,8 +206,15 @@ func (a *App) SendNotifications(post *model.Post, team *model.Team, channel *mod return } } - _, err := a.Srv().Store.Thread().MaintainMembership(userID, post.RootId, true, incrementMentions, *a.Config().ServiceSettings.ThreadAutoFollow, userID == post.UserId, userID == post.UserId) + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: incrementMentions, + UpdateFollowing: *a.Config().ServiceSettings.ThreadAutoFollow, + UpdateViewedTimestamp: userID == post.UserId, + UpdateParticipants: userID == post.UserId, + } + _, err := a.Srv().Store.Thread().MaintainMembership(userID, post.RootId, opts) if err != nil { mac <- model.NewAppError("SendNotifications", "app.channel.autofollow.app_error", nil, err.Error(), http.StatusInternalServerError) return diff --git a/app/user.go b/app/user.go index 4e9e3e97dc..e2e730fab9 100644 --- a/app/user.go +++ b/app/user.go @@ -2320,7 +2320,14 @@ func (a *App) UpdateThreadsReadForUser(userID, teamID string) *model.AppError { } func (a *App) UpdateThreadFollowForUser(userID, teamID, threadID string, state bool) *model.AppError { - _, err := a.Srv().Store.Thread().MaintainMembership(userID, threadID, state, false, true, state, false) + opts := store.ThreadMembershipOpts{ + Following: state, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: state, + UpdateParticipants: false, + } + _, err := a.Srv().Store.Thread().MaintainMembership(userID, threadID, opts) if err != nil { return model.NewAppError("UpdateThreadFollowForUser", "app.user.update_thread_follow_for_user.app_error", nil, err.Error(), http.StatusInternalServerError) } diff --git a/store/layer_generators/main.go b/store/layer_generators/main.go index 7bc920064c..fcbe9783ca 100644 --- a/store/layer_generators/main.go +++ b/store/layer_generators/main.go @@ -282,7 +282,7 @@ func generateLayer(name, templateFile string) ([]byte, error) { paramsWithType := []string{} for _, param := range params { switch param.Type { - case "ChannelSearchOpts", "UserGetByIdsOpts": + case "ChannelSearchOpts", "UserGetByIdsOpts", "ThreadMembershipOpts": paramsWithType = append(paramsWithType, fmt.Sprintf("%s store.%s", param.Name, param.Type)) case "*UserGetByIdsOpts": paramsWithType = append(paramsWithType, fmt.Sprintf("%s *store.UserGetByIdsOpts", param.Name)) @@ -296,7 +296,7 @@ func generateLayer(name, templateFile string) ([]byte, error) { paramsWithType := []string{} for _, param := range params { switch param.Type { - case "ChannelSearchOpts", "UserGetByIdsOpts": + case "ChannelSearchOpts", "UserGetByIdsOpts", "ThreadMembershipOpts": paramsWithType = append(paramsWithType, fmt.Sprintf("%s store.%s", param.Name, param.Type)) case "*UserGetByIdsOpts": paramsWithType = append(paramsWithType, fmt.Sprintf("%s *store.UserGetByIdsOpts", param.Name)) diff --git a/store/opentracinglayer/opentracinglayer.go b/store/opentracinglayer/opentracinglayer.go index 2587d5fa9a..2a54613178 100644 --- a/store/opentracinglayer/opentracinglayer.go +++ b/store/opentracinglayer/opentracinglayer.go @@ -8902,7 +8902,7 @@ func (s *OpenTracingLayerThreadStore) GetThreadsForUser(userId string, teamID st return result, err } -func (s *OpenTracingLayerThreadStore) MaintainMembership(userID string, postID string, following bool, incrementMentions bool, updateFollowing bool, updateViewedTimestamp bool, updateParticipants bool) (*model.ThreadMembership, error) { +func (s *OpenTracingLayerThreadStore) MaintainMembership(userID string, postID string, opts store.ThreadMembershipOpts) (*model.ThreadMembership, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.MaintainMembership") s.Root.Store.SetContext(newCtx) @@ -8911,7 +8911,7 @@ func (s *OpenTracingLayerThreadStore) MaintainMembership(userID string, postID s }() defer span.Finish() - result, err := s.ThreadStore.MaintainMembership(userID, postID, following, incrementMentions, updateFollowing, updateViewedTimestamp, updateParticipants) + result, err := s.ThreadStore.MaintainMembership(userID, postID, opts) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) diff --git a/store/retrylayer/retrylayer.go b/store/retrylayer/retrylayer.go index 68f1572058..32365bdf96 100644 --- a/store/retrylayer/retrylayer.go +++ b/store/retrylayer/retrylayer.go @@ -9688,11 +9688,11 @@ func (s *RetryLayerThreadStore) GetThreadsForUser(userId string, teamID string, } -func (s *RetryLayerThreadStore) MaintainMembership(userID string, postID string, following bool, incrementMentions bool, updateFollowing bool, updateViewedTimestamp bool, updateParticipants bool) (*model.ThreadMembership, error) { +func (s *RetryLayerThreadStore) MaintainMembership(userID string, postID string, opts store.ThreadMembershipOpts) (*model.ThreadMembership, error) { tries := 0 for { - result, err := s.ThreadStore.MaintainMembership(userID, postID, following, incrementMentions, updateFollowing, updateViewedTimestamp, updateParticipants) + result, err := s.ThreadStore.MaintainMembership(userID, postID, opts) if err == nil { return result, nil } diff --git a/store/sqlstore/thread_store.go b/store/sqlstore/thread_store.go index ef7d5f1237..4606d93a34 100644 --- a/store/sqlstore/thread_store.go +++ b/store/sqlstore/thread_store.go @@ -9,6 +9,7 @@ import ( "time" sq "github.com/Masterminds/squirrel" + "github.com/mattermost/gorp" "github.com/pkg/errors" "github.com/mattermost/mattermost-server/v5/model" @@ -88,7 +89,11 @@ func (s *SqlThreadStore) Save(thread *model.Thread) (*model.Thread, error) { } func (s *SqlThreadStore) Update(thread *model.Thread) (*model.Thread, error) { - if _, err := s.GetMaster().Update(thread); err != nil { + return s.update(s.GetMaster(), thread) +} + +func (s *SqlThreadStore) update(ex gorp.SqlExecutor, thread *model.Thread) (*model.Thread, error) { + if _, err := ex.Update(thread); err != nil { return nil, errors.Wrapf(err, "failed to update thread with id=%s", thread.PostId) } @@ -96,9 +101,13 @@ func (s *SqlThreadStore) Update(thread *model.Thread) (*model.Thread, error) { } func (s *SqlThreadStore) Get(id string) (*model.Thread, error) { + return s.get(s.GetReplica(), id) +} + +func (s *SqlThreadStore) get(ex gorp.SqlExecutor, id string) (*model.Thread, error) { var thread model.Thread query, args, _ := s.getQueryBuilder().Select("*").From("Threads").Where(sq.Eq{"PostId": id}).ToSql() - err := s.GetReplica().SelectOne(&thread, query, args...) + err := ex.SelectOne(&thread, query, args...) if err != nil { if err == sql.ErrNoRows { return nil, store.NewErrNotFound("Thread", id) @@ -474,7 +483,11 @@ func (s *SqlThreadStore) Delete(threadId string) error { } func (s *SqlThreadStore) SaveMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { - if err := s.GetMaster().Insert(membership); err != nil { + return s.saveMembership(s.GetMaster(), membership) +} + +func (s *SqlThreadStore) saveMembership(ex gorp.SqlExecutor, membership *model.ThreadMembership) (*model.ThreadMembership, error) { + if err := ex.Insert(membership); err != nil { return nil, errors.Wrapf(err, "failed to save thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) } @@ -482,7 +495,11 @@ func (s *SqlThreadStore) SaveMembership(membership *model.ThreadMembership) (*mo } func (s *SqlThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { - if _, err := s.GetMaster().Update(membership); err != nil { + return s.updateMembership(s.GetMaster(), membership) +} + +func (s *SqlThreadStore) updateMembership(ex gorp.SqlExecutor, membership *model.ThreadMembership) (*model.ThreadMembership, error) { + if _, err := ex.Update(membership); err != nil { return nil, errors.Wrapf(err, "failed to update thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) } @@ -509,8 +526,12 @@ func (s *SqlThreadStore) GetMembershipsForUser(userId, teamId string) ([]*model. } func (s *SqlThreadStore) GetMembershipForUser(userId, postId string) (*model.ThreadMembership, error) { + return s.getMembershipForUser(s.GetReplica(), userId, postId) +} + +func (s *SqlThreadStore) getMembershipForUser(ex gorp.SqlExecutor, userId, postId string) (*model.ThreadMembership, error) { var membership model.ThreadMembership - err := s.GetReplica().SelectOne(&membership, "SELECT * from ThreadMemberships WHERE UserId = :UserId AND PostId = :PostId", map[string]interface{}{"UserId": userId, "PostId": postId}) + err := ex.SelectOne(&membership, "SELECT * from ThreadMemberships WHERE UserId = :UserId AND PostId = :PostId", map[string]interface{}{"UserId": userId, "PostId": postId}) if err != nil { if err == sql.ErrNoRows { return nil, store.NewErrNotFound("Thread", postId) @@ -528,66 +549,89 @@ func (s *SqlThreadStore) DeleteMembershipForUser(userId string, postId string) e return nil } -func (s *SqlThreadStore) MaintainMembership(userId, postId string, following, incrementMentions, updateFollowing, updateViewedTimestamp, updateParticipants bool) (*model.ThreadMembership, error) { - membership, err := s.GetMembershipForUser(userId, postId) +// MaintainMembership creates or updates a thread membership for the given user +// and post. This method is used to update the state of a membership in response +// to some events like: +// - post creation (mentions handling) +// - channel marked unread +// - user explicitly following a thread +func (s *SqlThreadStore) MaintainMembership(userId, postId string, opts store.ThreadMembershipOpts) (*model.ThreadMembership, error) { + trx, err := s.GetMaster().Begin() + if err != nil { + return nil, errors.Wrap(err, "begin_transaction") + } + defer finalizeTransaction(trx) + + membership, err := s.getMembershipForUser(trx, userId, postId) now := utils.MillisFromTime(time.Now()) // if memebership exists, update it if: // a. user started/stopped following a thread // b. mention count changed // c. user viewed a thread if err == nil { - followingNeedsUpdate := (updateFollowing && !membership.Following || membership.Following != following) - if followingNeedsUpdate || incrementMentions || updateViewedTimestamp { + followingNeedsUpdate := (opts.UpdateFollowing && !membership.Following || membership.Following != opts.Following) + if followingNeedsUpdate || opts.IncrementMentions || opts.UpdateViewedTimestamp { if followingNeedsUpdate { - membership.Following = following + membership.Following = opts.Following } - if updateViewedTimestamp { + if opts.UpdateViewedTimestamp { membership.LastViewed = now } membership.LastUpdated = now - if incrementMentions { + if opts.IncrementMentions { membership.UnreadMentions += 1 } - _, err = s.UpdateMembership(membership) + if _, err = s.updateMembership(trx, membership); err != nil { + return nil, err + } } - return nil, err + + if err = trx.Commit(); err != nil { + return nil, errors.Wrap(err, "commit_transaction") + } + + return membership, err } var nfErr *store.ErrNotFound - if !errors.As(err, &nfErr) { return nil, errors.Wrap(err, "failed to get thread membership") } - mentions := 0 - if incrementMentions { - mentions = 1 + + membership = &model.ThreadMembership{ + PostId: postId, + UserId: userId, + Following: opts.Following, + LastUpdated: now, } - var lastViewed int64 - if updateViewedTimestamp { - lastViewed = now + if opts.IncrementMentions { + membership.UnreadMentions = 1 } - membership, err = s.SaveMembership(&model.ThreadMembership{ - PostId: postId, - UserId: userId, - Following: following, - LastViewed: lastViewed, - LastUpdated: now, - UnreadMentions: int64(mentions), - }) + if opts.UpdateViewedTimestamp { + membership.LastViewed = now + } + membership, err = s.saveMembership(trx, membership) if err != nil { return nil, err } - if updateParticipants { - thread, err2 := s.Get(postId) - if err2 != nil { - return nil, err2 + if opts.UpdateParticipants { + thread, getErr := s.get(trx, postId) + if getErr != nil { + return nil, getErr } if !thread.Participants.Contains(userId) { thread.Participants = append(thread.Participants, userId) - _, err = s.Update(thread) + if _, err = s.update(trx, thread); err != nil { + return nil, err + } } } + + if err = trx.Commit(); err != nil { + return nil, errors.Wrap(err, "commit_transaction") + } + return membership, err } diff --git a/store/store.go b/store/store.go index 4fbbfc2fcb..b46f0aa7c5 100644 --- a/store/store.go +++ b/store/store.go @@ -296,7 +296,7 @@ type ThreadStore interface { GetMembershipsForUser(userId, teamID string) ([]*model.ThreadMembership, error) GetMembershipForUser(userId, postID string) (*model.ThreadMembership, error) DeleteMembershipForUser(userId, postID string) error - MaintainMembership(userID, postID string, following, incrementMentions, updateFollowing, updateViewedTimestamp, updateParticipants bool) (*model.ThreadMembership, error) + MaintainMembership(userID, postID string, opts ThreadMembershipOpts) (*model.ThreadMembership, error) CollectThreadsWithNewerReplies(userId string, channelIds []string, timestamp int64) ([]string, error) UpdateUnreadsByChannel(userId string, changedThreads []string, timestamp int64, updateViewedTimestamp bool) error } @@ -912,3 +912,21 @@ type UserGetByIdsOpts struct { // Since filters the users based on their UpdateAt timestamp. Since int64 } + +// ThreadMembershipOpts defines some properties to be passed to +// ThreadStore.MaintainMembership() +type ThreadMembershipOpts struct { + // Following indicates whether or not the user is following the thread. + Following bool + // IncrementMentions indicates whether or not the mentions count for + // the thread should be incremented. + IncrementMentions bool + // UpdateFollowing indicates whether or not a membership update should be forced. + UpdateFollowing bool + // UpdateViewedTimestamp indicates whether or not the LastViewed field of the + // membership should be updated. + UpdateViewedTimestamp bool + // UpdateParticipants indicates whether or not the thread's participants list + // should be updated. + UpdateParticipants bool +} diff --git a/store/storetest/mocks/ThreadStore.go b/store/storetest/mocks/ThreadStore.go index 339a761808..d6c4b76764 100644 --- a/store/storetest/mocks/ThreadStore.go +++ b/store/storetest/mocks/ThreadStore.go @@ -6,6 +6,7 @@ package mocks import ( model "github.com/mattermost/mattermost-server/v5/model" + store "github.com/mattermost/mattermost-server/v5/store" mock "github.com/stretchr/testify/mock" ) @@ -226,13 +227,13 @@ func (_m *ThreadStore) GetThreadsForUser(userId string, teamID string, opts mode return r0, r1 } -// MaintainMembership provides a mock function with given fields: userID, postID, following, incrementMentions, updateFollowing, updateViewedTimestamp, updateParticipants -func (_m *ThreadStore) MaintainMembership(userID string, postID string, following bool, incrementMentions bool, updateFollowing bool, updateViewedTimestamp bool, updateParticipants bool) (*model.ThreadMembership, error) { - ret := _m.Called(userID, postID, following, incrementMentions, updateFollowing, updateViewedTimestamp, updateParticipants) +// MaintainMembership provides a mock function with given fields: userID, postID, opts +func (_m *ThreadStore) MaintainMembership(userID string, postID string, opts store.ThreadMembershipOpts) (*model.ThreadMembership, error) { + ret := _m.Called(userID, postID, opts) var r0 *model.ThreadMembership - if rf, ok := ret.Get(0).(func(string, string, bool, bool, bool, bool, bool) *model.ThreadMembership); ok { - r0 = rf(userID, postID, following, incrementMentions, updateFollowing, updateViewedTimestamp, updateParticipants) + if rf, ok := ret.Get(0).(func(string, string, store.ThreadMembershipOpts) *model.ThreadMembership); ok { + r0 = rf(userID, postID, opts) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(*model.ThreadMembership) @@ -240,8 +241,8 @@ func (_m *ThreadStore) MaintainMembership(userID string, postID string, followin } var r1 error - if rf, ok := ret.Get(1).(func(string, string, bool, bool, bool, bool, bool) error); ok { - r1 = rf(userID, postID, following, incrementMentions, updateFollowing, updateViewedTimestamp, updateParticipants) + if rf, ok := ret.Get(1).(func(string, string, store.ThreadMembershipOpts) error); ok { + r1 = rf(userID, postID, opts) } else { r1 = ret.Error(1) } diff --git a/store/storetest/thread_store.go b/store/storetest/thread_store.go index df1775cc06..d8f6d3a8be 100644 --- a/store/storetest/thread_store.go +++ b/store/storetest/thread_store.go @@ -235,7 +235,14 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) { t.Run("Thread last updated is changed when channel is updated after UpdateLastViewedAtPost", func(t *testing.T) { newPosts := makeSomePosts() - _, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, true, false, true, false, false) + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + _, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, opts) require.NoError(t, e) m, err1 := ss.Thread().GetMembershipForUser(newPosts[0].UserId, newPosts[0].Id) require.NoError(t, err1) @@ -256,7 +263,14 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) { t.Run("Thread last updated is changed when channel is updated after IncrementMentionCount", func(t *testing.T) { newPosts := makeSomePosts() - _, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, true, false, true, false, false) + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + _, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, opts) require.NoError(t, e) m, err1 := ss.Thread().GetMembershipForUser(newPosts[0].UserId, newPosts[0].Id) require.NoError(t, err1) @@ -276,7 +290,14 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) { t.Run("Thread last updated is changed when channel is updated after UpdateLastViewedAt", func(t *testing.T) { newPosts := makeSomePosts() - _, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, true, false, true, false, false) + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + _, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, opts) require.NoError(t, e) m, err1 := ss.Thread().GetMembershipForUser(newPosts[0].UserId, newPosts[0].Id) require.NoError(t, err1) @@ -297,11 +318,19 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) { t.Run("Thread membership 'viewed' timestamp is updated properly", func(t *testing.T) { newPosts := makeSomePosts() - tm, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, true, false, true, false, false) + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + tm, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, opts) require.NoError(t, e) require.Equal(t, int64(0), tm.LastViewed) - _, e = ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, true, false, true, true, false) + opts.UpdateViewedTimestamp = true + _, e = ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, opts) require.NoError(t, e) m2, err2 := ss.Thread().GetMembershipForUser(newPosts[0].UserId, newPosts[0].Id) require.NoError(t, err2) @@ -310,14 +339,29 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) { t.Run("Thread membership 'viewed' timestamp is updated properly for new membership", func(t *testing.T) { newPosts := makeSomePosts() - tm, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, true, false, false, true, false) + + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: false, + UpdateViewedTimestamp: true, + UpdateParticipants: false, + } + tm, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, opts) require.NoError(t, e) require.NotEqual(t, int64(0), tm.LastViewed) }) t.Run("Thread last updated is changed when channel is updated after UpdateLastViewedAtPost for mark unread", func(t *testing.T) { newPosts := makeSomePosts() - _, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, true, false, true, false, false) + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + _, e := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, opts) require.NoError(t, e) m, err1 := ss.Thread().GetMembershipForUser(newPosts[0].UserId, newPosts[0].Id) require.NoError(t, err1) @@ -337,7 +381,14 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) { t.Run("Updating post does not make thread unread", func(t *testing.T) { newPosts := makeSomePosts() - m, err := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, true, false, true, false, false) + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + m, err := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, opts) require.NoError(t, err) th, err := ss.Thread().GetThreadForUser("", m, false) require.NoError(t, err) diff --git a/store/timerlayer/timerlayer.go b/store/timerlayer/timerlayer.go index 71c5f4cc28..bb0f7ef985 100644 --- a/store/timerlayer/timerlayer.go +++ b/store/timerlayer/timerlayer.go @@ -8022,10 +8022,10 @@ func (s *TimerLayerThreadStore) GetThreadsForUser(userId string, teamID string, return result, err } -func (s *TimerLayerThreadStore) MaintainMembership(userID string, postID string, following bool, incrementMentions bool, updateFollowing bool, updateViewedTimestamp bool, updateParticipants bool) (*model.ThreadMembership, error) { +func (s *TimerLayerThreadStore) MaintainMembership(userID string, postID string, opts store.ThreadMembershipOpts) (*model.ThreadMembership, error) { start := timemodule.Now() - result, err := s.ThreadStore.MaintainMembership(userID, postID, following, incrementMentions, updateFollowing, updateViewedTimestamp, updateParticipants) + result, err := s.ThreadStore.MaintainMembership(userID, postID, opts) elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) if s.Root.Metrics != nil {