diff --git a/app/channel.go b/app/channel.go index 362bb763d4..0b5d96ff3e 100644 --- a/app/channel.go +++ b/app/channel.go @@ -2877,7 +2877,17 @@ func (a *App) MarkChannelsAsViewed(channelIDs []string, userID string, currentSe } } } - times, err := a.Srv().Store.Channel().UpdateLastViewedAt(channelIDs, userID, false) + + var err error + updateThreads := *a.Config().ServiceSettings.ThreadAutoFollow && (!collapsedThreadsSupported || !a.IsCRTEnabledForUser(userID)) + if updateThreads { + err = a.Srv().Store.Thread().MarkAllAsReadByChannels(userID, channelIDs) + if err != nil { + return nil, model.NewAppError("MarkChannelsAsViewed", "app.channel.update_last_viewed_at.app_error", nil, err.Error(), http.StatusInternalServerError) + } + } + + times, err := a.Srv().Store.Channel().UpdateLastViewedAt(channelIDs, userID) if err != nil { var invErr *store.ErrInvalidInput switch { @@ -2899,18 +2909,12 @@ func (a *App) MarkChannelsAsViewed(channelIDs []string, userID string, currentSe a.clearPushNotification(currentSessionId, userID, channelID, "") } - if *a.Config().ServiceSettings.ThreadAutoFollow && (!collapsedThreadsSupported || !a.IsCRTEnabledForUser(userID)) { - if err := a.Srv().Store.Thread().MarkAllAsReadInChannels(userID, channelIDs); err != nil { - return nil, model.NewAppError("MarkChannelsAsViewed", "app.channel.update_last_viewed_at.app_error", nil, err.Error(), http.StatusInternalServerError) - } - - if a.IsCRTEnabledForUser(userID) { - timestamp := model.GetMillis() - for _, channelID := range channelIDs { - message := model.NewWebSocketEvent(model.WebsocketEventThreadReadChanged, "", channelID, userID, nil) - message.Add("timestamp", timestamp) - a.Publish(message) - } + if updateThreads && a.IsCRTEnabledForUser(userID) { + timestamp := model.GetMillis() + for _, channelID := range channelIDs { + message := model.NewWebSocketEvent(model.WebsocketEventThreadReadChanged, "", channelID, userID, nil) + message.Add("timestamp", timestamp) + a.Publish(message) } } diff --git a/app/channel_test.go b/app/channel_test.go index 889e8675e3..bb1e9d2b15 100644 --- a/app/channel_test.go +++ b/app/channel_test.go @@ -1978,7 +1978,7 @@ func TestMarkChannelsAsViewedPanic(t *testing.T) { times := map[string]int64{ "userID": 1, } - mockChannelStore.On("UpdateLastViewedAt", []string{"channelID"}, "userID", false).Return(times, nil) + mockChannelStore.On("UpdateLastViewedAt", []string{"channelID"}, "userID").Return(times, nil) mockSessionStore := mocks.SessionStore{} mockOAuthStore := mocks.OAuthStore{} var err error @@ -1992,11 +1992,8 @@ func TestMarkChannelsAsViewedPanic(t *testing.T) { require.NoError(t, err) mockPreferenceStore := mocks.PreferenceStore{} mockPreferenceStore.On("Get", mock.AnythingOfType("string"), mock.AnythingOfType("string"), mock.AnythingOfType("string")).Return(&model.Preference{Value: "test"}, nil) - mockThreadStore := mocks.ThreadStore{} - mockThreadStore.On("MarkAllAsReadInChannels", "userID", []string{"channelID"}).Return(nil) mockStore.On("Channel").Return(&mockChannelStore) mockStore.On("Preference").Return(&mockPreferenceStore) - mockStore.On("Thread").Return(&mockThreadStore) _, appErr := th.App.MarkChannelsAsViewed([]string{"channelID"}, "userID", th.Context.Session().Id, false) require.Nil(t, appErr) @@ -2034,7 +2031,6 @@ func TestMarkChannelAsUnreadFromPostPanic(t *testing.T) { mockPreferenceStore.On("Get", "userID", model.PreferenceCategoryDisplaySettings, model.PreferenceNameCollapsedThreadsEnabled).Return(&model.Preference{Value: "on"}, nil) mockThreadStore := mocks.ThreadStore{} - mockThreadStore.On("MarkAllAsReadInChannels", "userID", []string{"channelID"}).Return(nil) mockThreadStore.On("GetMembershipForUser", "userID", "rootID").Return(nil, nil) mockThreadStore.On("MaintainMembership", "userID", "rootID", mock.AnythingOfType("store.ThreadMembershipOpts")).Return(&model.ThreadMembership{}, nil) mockThreadStore.On("Get", "rootID").Return(nil, errors.New("bad error")) // Returning an error from here causes the panic diff --git a/app/user.go b/app/user.go index dc566165d1..1c90505f77 100644 --- a/app/user.go +++ b/app/user.go @@ -2300,7 +2300,7 @@ func (a *App) GetThreadForUser(teamID string, threadMembership *model.ThreadMemb } func (a *App) UpdateThreadsReadForUser(userID, teamID string) *model.AppError { - nErr := a.Srv().Store.Thread().MarkAllAsRead(userID, teamID) + nErr := a.Srv().Store.Thread().MarkAllAsReadByTeam(userID, teamID) if nErr != nil { return model.NewAppError("UpdateThreadsReadForUser", "app.user.update_threads_read_for_user.app_error", nil, nErr.Error(), http.StatusInternalServerError) } diff --git a/store/opentracinglayer/opentracinglayer.go b/store/opentracinglayer/opentracinglayer.go index 123ae01c56..0f2de5eb36 100644 --- a/store/opentracinglayer/opentracinglayer.go +++ b/store/opentracinglayer/opentracinglayer.go @@ -2284,7 +2284,7 @@ func (s *OpenTracingLayerChannelStore) Update(channel *model.Channel) (*model.Ch return result, err } -func (s *OpenTracingLayerChannelStore) UpdateLastViewedAt(channelIds []string, userID string, updateThreads bool) (map[string]int64, error) { +func (s *OpenTracingLayerChannelStore) UpdateLastViewedAt(channelIds []string, userID string) (map[string]int64, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ChannelStore.UpdateLastViewedAt") s.Root.Store.SetContext(newCtx) @@ -2293,7 +2293,7 @@ func (s *OpenTracingLayerChannelStore) UpdateLastViewedAt(channelIds []string, u }() defer span.Finish() - result, err := s.ChannelStore.UpdateLastViewedAt(channelIds, userID, updateThreads) + result, err := s.ChannelStore.UpdateLastViewedAt(channelIds, userID) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -9213,24 +9213,6 @@ func (s *OpenTracingLayerTermsOfServiceStore) Save(termsOfService *model.TermsOf return result, err } -func (s *OpenTracingLayerThreadStore) CollectThreadsWithNewerReplies(userId string, channelIds []string, timestamp int64) ([]string, error) { - origCtx := s.Root.Store.Context() - span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.CollectThreadsWithNewerReplies") - s.Root.Store.SetContext(newCtx) - defer func() { - s.Root.Store.SetContext(origCtx) - }() - - defer span.Finish() - result, err := s.ThreadStore.CollectThreadsWithNewerReplies(userId, channelIds, timestamp) - if err != nil { - span.LogFields(spanlog.Error(err)) - ext.Error.Set(span, true) - } - - return result, err -} - func (s *OpenTracingLayerThreadStore) DeleteMembershipForUser(userId string, postID string) error { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.DeleteMembershipForUser") @@ -9447,7 +9429,7 @@ func (s *OpenTracingLayerThreadStore) MaintainMembership(userID string, postID s return result, err } -func (s *OpenTracingLayerThreadStore) MarkAllAsRead(userID string, teamID string) error { +func (s *OpenTracingLayerThreadStore) MarkAllAsRead(userID string, threadIds []string) error { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.MarkAllAsRead") s.Root.Store.SetContext(newCtx) @@ -9456,7 +9438,7 @@ func (s *OpenTracingLayerThreadStore) MarkAllAsRead(userID string, teamID string }() defer span.Finish() - err := s.ThreadStore.MarkAllAsRead(userID, teamID) + err := s.ThreadStore.MarkAllAsRead(userID, threadIds) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -9465,16 +9447,34 @@ func (s *OpenTracingLayerThreadStore) MarkAllAsRead(userID string, teamID string return err } -func (s *OpenTracingLayerThreadStore) MarkAllAsReadInChannels(userID string, channelIDs []string) error { +func (s *OpenTracingLayerThreadStore) MarkAllAsReadByChannels(userID string, channelIDs []string) error { origCtx := s.Root.Store.Context() - span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.MarkAllAsReadInChannels") + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.MarkAllAsReadByChannels") s.Root.Store.SetContext(newCtx) defer func() { s.Root.Store.SetContext(origCtx) }() defer span.Finish() - err := s.ThreadStore.MarkAllAsReadInChannels(userID, channelIDs) + err := s.ThreadStore.MarkAllAsReadByChannels(userID, channelIDs) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return err +} + +func (s *OpenTracingLayerThreadStore) MarkAllAsReadByTeam(userID string, teamID string) error { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.MarkAllAsReadByTeam") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + err := s.ThreadStore.MarkAllAsReadByTeam(userID, teamID) if err != nil { span.LogFields(spanlog.Error(err)) ext.Error.Set(span, true) @@ -9537,24 +9537,6 @@ func (s *OpenTracingLayerThreadStore) PermanentDeleteBatchThreadMembershipsForRe return result, resultVar1, err } -func (s *OpenTracingLayerThreadStore) UpdateLastViewedByThreadIds(userId string, threadIds []string, timestamp int64) error { - origCtx := s.Root.Store.Context() - span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.UpdateLastViewedByThreadIds") - s.Root.Store.SetContext(newCtx) - defer func() { - s.Root.Store.SetContext(origCtx) - }() - - defer span.Finish() - err := s.ThreadStore.UpdateLastViewedByThreadIds(userId, threadIds, timestamp) - if err != nil { - span.LogFields(spanlog.Error(err)) - ext.Error.Set(span, true) - } - - return err -} - func (s *OpenTracingLayerThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.UpdateMembership") diff --git a/store/retrylayer/retrylayer.go b/store/retrylayer/retrylayer.go index 174f70c61d..50116b5d62 100644 --- a/store/retrylayer/retrylayer.go +++ b/store/retrylayer/retrylayer.go @@ -2527,11 +2527,11 @@ func (s *RetryLayerChannelStore) Update(channel *model.Channel) (*model.Channel, } -func (s *RetryLayerChannelStore) UpdateLastViewedAt(channelIds []string, userID string, updateThreads bool) (map[string]int64, error) { +func (s *RetryLayerChannelStore) UpdateLastViewedAt(channelIds []string, userID string) (map[string]int64, error) { tries := 0 for { - result, err := s.ChannelStore.UpdateLastViewedAt(channelIds, userID, updateThreads) + result, err := s.ChannelStore.UpdateLastViewedAt(channelIds, userID) if err == nil { return result, nil } @@ -10519,27 +10519,6 @@ func (s *RetryLayerTermsOfServiceStore) Save(termsOfService *model.TermsOfServic } -func (s *RetryLayerThreadStore) CollectThreadsWithNewerReplies(userId string, channelIds []string, timestamp int64) ([]string, error) { - - tries := 0 - for { - result, err := s.ThreadStore.CollectThreadsWithNewerReplies(userId, channelIds, timestamp) - if err == nil { - return result, nil - } - if !isRepeatableError(err) { - return result, err - } - tries++ - if tries >= 3 { - err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") - return result, err - } - timepkg.Sleep(100 * timepkg.Millisecond) - } - -} - func (s *RetryLayerThreadStore) DeleteMembershipForUser(userId string, postID string) error { tries := 0 @@ -10792,11 +10771,11 @@ func (s *RetryLayerThreadStore) MaintainMembership(userID string, postID string, } -func (s *RetryLayerThreadStore) MarkAllAsRead(userID string, teamID string) error { +func (s *RetryLayerThreadStore) MarkAllAsRead(userID string, threadIds []string) error { tries := 0 for { - err := s.ThreadStore.MarkAllAsRead(userID, teamID) + err := s.ThreadStore.MarkAllAsRead(userID, threadIds) if err == nil { return nil } @@ -10813,11 +10792,32 @@ func (s *RetryLayerThreadStore) MarkAllAsRead(userID string, teamID string) erro } -func (s *RetryLayerThreadStore) MarkAllAsReadInChannels(userID string, channelIDs []string) error { +func (s *RetryLayerThreadStore) MarkAllAsReadByChannels(userID string, channelIDs []string) error { tries := 0 for { - err := s.ThreadStore.MarkAllAsReadInChannels(userID, channelIDs) + err := s.ThreadStore.MarkAllAsReadByChannels(userID, channelIDs) + if err == nil { + return nil + } + if !isRepeatableError(err) { + return err + } + tries++ + if tries >= 3 { + err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") + return err + } + timepkg.Sleep(100 * timepkg.Millisecond) + } + +} + +func (s *RetryLayerThreadStore) MarkAllAsReadByTeam(userID string, teamID string) error { + + tries := 0 + for { + err := s.ThreadStore.MarkAllAsReadByTeam(userID, teamID) if err == nil { return nil } @@ -10897,27 +10897,6 @@ func (s *RetryLayerThreadStore) PermanentDeleteBatchThreadMembershipsForRetentio } -func (s *RetryLayerThreadStore) UpdateLastViewedByThreadIds(userId string, threadIds []string, timestamp int64) error { - - tries := 0 - for { - err := s.ThreadStore.UpdateLastViewedByThreadIds(userId, threadIds, timestamp) - if err == nil { - return nil - } - if !isRepeatableError(err) { - return err - } - tries++ - if tries >= 3 { - err = errors.Wrap(err, "giving up after 3 consecutive repeatable transaction failures") - return err - } - timepkg.Sleep(100 * timepkg.Millisecond) - } - -} - func (s *RetryLayerThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { tries := 0 diff --git a/store/sqlstore/channel_store.go b/store/sqlstore/channel_store.go index 9f1c16622a..a3b0c62373 100644 --- a/store/sqlstore/channel_store.go +++ b/store/sqlstore/channel_store.go @@ -2311,17 +2311,7 @@ func (s SqlChannelStore) PermanentDeleteMembersByUser(userId string) error { return nil } -func (s SqlChannelStore) UpdateLastViewedAt(channelIds []string, userId string, updateThreads bool) (map[string]int64, error) { - var threadsToUpdate []string - now := model.GetMillis() - if updateThreads { - var err error - threadsToUpdate, err = s.Thread().CollectThreadsWithNewerReplies(userId, channelIds, now) - if err != nil { - return nil, err - } - } - +func (s SqlChannelStore) UpdateLastViewedAt(channelIds []string, userId string) (map[string]int64, error) { lastPostAtTimes := []struct { Id string LastPostAt int64 @@ -2379,9 +2369,6 @@ func (s SqlChannelStore) UpdateLastViewedAt(channelIds []string, userId string, for _, t := range lastPostAtTimes { times[t.Id] = t.LastPostAt } - if updateThreads { - s.Thread().UpdateLastViewedByThreadIds(userId, threadsToUpdate, now) - } return times, nil } @@ -2424,9 +2411,6 @@ func (s SqlChannelStore) UpdateLastViewedAt(channelIds []string, userId string, return nil, errors.Wrapf(err, "failed to update ChannelMembers with userId=%s and channelId in %v", userId, channelIds) } - if updateThreads { - s.Thread().UpdateLastViewedByThreadIds(userId, threadsToUpdate, now) - } return times, nil } diff --git a/store/sqlstore/thread_store.go b/store/sqlstore/thread_store.go index cd64da9b11..d398222c7c 100644 --- a/store/sqlstore/thread_store.go +++ b/store/sqlstore/thread_store.go @@ -511,44 +511,73 @@ func (s *SqlThreadStore) GetThreadForUser(teamId string, threadMembership *model return result, nil } -// MarkAllAsReadInChannels marks all threads for the given user in the given channels as read from -// the current time. -func (s *SqlThreadStore) MarkAllAsReadInChannels(userID string, channelIDs []string) error { - threadIDs := []string{} - - query, args, _ := s.getQueryBuilder(). - Select("ThreadMemberships.PostId"). - Join("Threads ON Threads.PostId = ThreadMemberships.PostId"). - Join("Channels ON Threads.ChannelId = Channels.Id"). - From("ThreadMemberships"). - Where(sq.Eq{"Threads.ChannelId": channelIDs}). - Where(sq.Eq{"ThreadMemberships.UserId": userID}). - ToSql() - - err := s.GetReplicaX().Select(&threadIDs, query, args...) - if err != nil { - return errors.Wrapf(err, "failed to get thread membership with userid=%s", userID) +// MarkAllAsReadByChannels marks thread membership for the given users in the given channels +// as read. This is used by the application layer to keep threads up-to-date when CRT is disabled +// for the enduser, avoiding an influx of unread threads when first turning the feature on. +func (s *SqlThreadStore) MarkAllAsReadByChannels(userID string, channelIDs []string) error { + if len(channelIDs) == 0 { + return nil } + now := model.GetMillis() + + // TODO: Fork squirrel to include https://github.com/Masterminds/squirrel/pull/256 and + // support FROM in an UPDATE query. + channelIDsSql, channelIDsArgs := constructArrayArgs(channelIDs) + + var query string + if s.DriverName() == model.DatabaseDriverPostgres { + query = ` + UPDATE ThreadMemberships + SET LastViewed = ?, UnreadMentions = ?, LastUpdated = ? + FROM Threads + WHERE ThreadMemberships.UserId = ? + AND Threads.PostId = ThreadMemberships.PostId + AND Threads.ChannelID IN ` + channelIDsSql + ` + AND Threads.LastReplyAt > ThreadMemberships.LastViewed + ` + } else { + query = ` + UPDATE ThreadMemberships, Threads + SET ThreadMemberships.LastViewed = ?, ThreadMemberships.UnreadMentions = ?, ThreadMemberships.LastUpdated = ? + WHERE ThreadMemberships.UserId = ? + AND Threads.PostId = ThreadMemberships.PostId + AND Threads.ChannelID IN ` + channelIDsSql + ` + AND Threads.LastReplyAt > ThreadMemberships.LastViewed + ` + } + + args := []interface{}{now, 0, now, userID} + args = append(args, channelIDsArgs...) + + if _, err := s.GetMasterX().Exec(query, args...); err != nil { + return errors.Wrapf(err, "failed to mark all threads as read by channels for user id=%s", userID) + } + + return nil +} + +func (s *SqlThreadStore) MarkAllAsRead(userId string, threadIds []string) error { timestamp := model.GetMillis() - query, args, _ = s.getQueryBuilder(). + + query, args, _ := s.getQueryBuilder(). Update("ThreadMemberships"). - Where(sq.Eq{"PostId": threadIDs}). - Where(sq.Eq{"UserId": userID}). + Where(sq.Eq{"UserId": userId}). + Where(sq.Eq{"PostId": threadIds}). Set("LastViewed", timestamp). Set("UnreadMentions", 0). Set("LastUpdated", model.GetMillis()). ToSql() if _, err := s.GetMasterX().Exec(query, args...); err != nil { - return errors.Wrapf(err, "failed to update thread read state for user id=%s", userID) + return errors.Wrapf(err, "failed to mark %d threads as read for user id=%s", len(threadIds), userId) } - return nil + return nil } -// MarkAllAsRead marks all threads for the given user in the given team as read from the current -// time. -func (s *SqlThreadStore) MarkAllAsRead(userId, teamId string) error { +// MarkAllAsReadByTeam marks all threads for the given user in the given team as read from the +// current time. +func (s *SqlThreadStore) MarkAllAsReadByTeam(userId, teamId string) error { memberships, err := s.GetMembershipsForUser(userId, teamId) if err != nil { return err @@ -786,50 +815,6 @@ func (s *SqlThreadStore) MaintainMembership(userId, postId string, opts store.Th return membership, err } -func (s *SqlThreadStore) CollectThreadsWithNewerReplies(userId string, channelIds []string, timestamp int64) ([]string, error) { - changedThreads := []string{} - query, args, _ := s.getQueryBuilder(). - Select("Threads.PostId"). - From("Threads"). - LeftJoin("ChannelMembers ON ChannelMembers.ChannelId=Threads.ChannelId"). - Where(sq.And{ - sq.Eq{"Threads.ChannelId": channelIds}, - sq.Eq{"ChannelMembers.UserId": userId}, - sq.Or{ - sq.Expr("Threads.LastReplyAt > ChannelMembers.LastViewedAt"), - sq.Gt{"Threads.LastReplyAt": timestamp}, - }, - }). - ToSql() - if err := s.GetReplicaX().Select(&changedThreads, query, args...); err != nil { - return nil, errors.Wrap(err, "failed to fetch threads") - } - return changedThreads, nil -} - -// UpdateLastViewedByThreadIds marks the given threads as read up to the given timestamp. If there -// are no newer posts, it effectively marks the thread as read. If there are newer posts, say -// because the user explicitly marked a past post as unread, the thread will be considered unread -// past the given timestamp. -func (s *SqlThreadStore) UpdateLastViewedByThreadIds(userId string, threadIds []string, timestamp int64) error { - if len(threadIds) == 0 { - return nil - } - - qb := s.getQueryBuilder(). - Update("ThreadMemberships"). - Where(sq.Eq{"UserId": userId, "PostId": threadIds}). - Set("LastViewed", timestamp). - Set("LastUpdated", model.GetMillis()) - updateQuery, updateArgs, _ := qb.ToSql() - - if _, err := s.GetMasterX().Exec(updateQuery, updateArgs...); err != nil { - return errors.Wrap(err, "failed to update thread membership") - } - - return nil -} - func (s *SqlThreadStore) GetPosts(threadId string, since int64) ([]*model.Post, error) { query, args, _ := s.getQueryBuilder(). Select("*"). diff --git a/store/store.go b/store/store.go index dc5cc33caf..de83e6e497 100644 --- a/store/store.go +++ b/store/store.go @@ -223,7 +223,7 @@ type ChannelStore interface { RemoveMembers(channelID string, userIds []string) error PermanentDeleteMembersByUser(userID string) error PermanentDeleteMembersByChannel(channelID string) error - UpdateLastViewedAt(channelIds []string, userID string, updateThreads bool) (map[string]int64, error) + UpdateLastViewedAt(channelIds []string, userID string) (map[string]int64, error) UpdateLastViewedAtPost(unreadPost *model.Post, userID string, mentionCount, mentionCountRoot int, setUnreadCountRoot bool) (*model.ChannelUnreadAt, error) CountPostsAfter(channelID string, timestamp int64, userID string) (int, int, error) IncrementMentionCount(channelID string, userIDs []string, isRoot bool) error @@ -298,8 +298,9 @@ type ThreadStore interface { GetTeamsUnreadForUser(userID string, teamIDs []string) (map[string]*model.TeamUnread, error) GetPosts(threadID string, since int64) ([]*model.Post, error) - MarkAllAsRead(userID, teamID string) error - MarkAllAsReadInChannels(userID string, channelIDs []string) error + MarkAllAsRead(userID string, threadIds []string) error + MarkAllAsReadByTeam(userID, teamID string) error + MarkAllAsReadByChannels(userID string, channelIDs []string) error MarkAsRead(userID, threadID string, timestamp int64) error UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) @@ -307,8 +308,6 @@ type ThreadStore interface { GetMembershipForUser(userId, postID string) (*model.ThreadMembership, error) DeleteMembershipForUser(userId, postID string) error MaintainMembership(userID, postID string, opts ThreadMembershipOpts) (*model.ThreadMembership, error) - CollectThreadsWithNewerReplies(userId string, channelIds []string, timestamp int64) ([]string, error) - UpdateLastViewedByThreadIds(userId string, threadIds []string, timestamp int64) error PermanentDeleteBatchForRetentionPolicies(now, globalPolicyEndTime, limit int64, cursor model.RetentionPolicyCursor) (int64, model.RetentionPolicyCursor, error) PermanentDeleteBatchThreadMembershipsForRetentionPolicies(now, globalPolicyEndTime, limit int64, cursor model.RetentionPolicyCursor) (int64, model.RetentionPolicyCursor, error) DeleteOrphanedRows(limit int) (deleted int64, err error) diff --git a/store/storetest/channel_store.go b/store/storetest/channel_store.go index 1e79647a68..1449ce739d 100644 --- a/store/storetest/channel_store.go +++ b/store/storetest/channel_store.go @@ -4749,11 +4749,11 @@ func testChannelStoreUpdateLastViewedAt(t *testing.T, ss store.Store) { require.NoError(t, err) var times map[string]int64 - times, err = ss.Channel().UpdateLastViewedAt([]string{m1.ChannelId}, m1.UserId, false) + times, err = ss.Channel().UpdateLastViewedAt([]string{m1.ChannelId}, m1.UserId) require.NoError(t, err, "failed to update ", err) require.Equal(t, o1.LastPostAt, times[o1.Id], "last viewed at time incorrect") - times, err = ss.Channel().UpdateLastViewedAt([]string{m1.ChannelId, m2.ChannelId}, m1.UserId, false) + times, err = ss.Channel().UpdateLastViewedAt([]string{m1.ChannelId, m2.ChannelId}, m1.UserId) require.NoError(t, err, "failed to update ", err) require.Equal(t, o2.LastPostAt, times[o2.Id], "last viewed at time incorrect") @@ -4769,7 +4769,7 @@ func testChannelStoreUpdateLastViewedAt(t *testing.T, ss store.Store) { assert.Equal(t, o2.LastPostAt, rm2.LastUpdateAt) assert.Equal(t, o2.TotalMsgCount, rm2.MsgCount) - _, err = ss.Channel().UpdateLastViewedAt([]string{m1.ChannelId}, "missing id", false) + _, err = ss.Channel().UpdateLastViewedAt([]string{m1.ChannelId}, "missing id") require.NoError(t, err, "failed to update") } diff --git a/store/storetest/mocks/ChannelStore.go b/store/storetest/mocks/ChannelStore.go index 4671ec9a11..8bdd402f35 100644 --- a/store/storetest/mocks/ChannelStore.go +++ b/store/storetest/mocks/ChannelStore.go @@ -2008,13 +2008,13 @@ func (_m *ChannelStore) Update(channel *model.Channel) (*model.Channel, error) { return r0, r1 } -// UpdateLastViewedAt provides a mock function with given fields: channelIds, userID, updateThreads -func (_m *ChannelStore) UpdateLastViewedAt(channelIds []string, userID string, updateThreads bool) (map[string]int64, error) { - ret := _m.Called(channelIds, userID, updateThreads) +// UpdateLastViewedAt provides a mock function with given fields: channelIds, userID +func (_m *ChannelStore) UpdateLastViewedAt(channelIds []string, userID string) (map[string]int64, error) { + ret := _m.Called(channelIds, userID) var r0 map[string]int64 - if rf, ok := ret.Get(0).(func([]string, string, bool) map[string]int64); ok { - r0 = rf(channelIds, userID, updateThreads) + if rf, ok := ret.Get(0).(func([]string, string) map[string]int64); ok { + r0 = rf(channelIds, userID) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(map[string]int64) @@ -2022,8 +2022,8 @@ func (_m *ChannelStore) UpdateLastViewedAt(channelIds []string, userID string, u } var r1 error - if rf, ok := ret.Get(1).(func([]string, string, bool) error); ok { - r1 = rf(channelIds, userID, updateThreads) + if rf, ok := ret.Get(1).(func([]string, string) error); ok { + r1 = rf(channelIds, userID) } else { r1 = ret.Error(1) } diff --git a/store/storetest/mocks/ThreadStore.go b/store/storetest/mocks/ThreadStore.go index 2321d02212..c936c4d2a3 100644 --- a/store/storetest/mocks/ThreadStore.go +++ b/store/storetest/mocks/ThreadStore.go @@ -15,29 +15,6 @@ type ThreadStore struct { mock.Mock } -// CollectThreadsWithNewerReplies provides a mock function with given fields: userId, channelIds, timestamp -func (_m *ThreadStore) CollectThreadsWithNewerReplies(userId string, channelIds []string, timestamp int64) ([]string, error) { - ret := _m.Called(userId, channelIds, timestamp) - - var r0 []string - if rf, ok := ret.Get(0).(func(string, []string, int64) []string); ok { - r0 = rf(userId, channelIds, timestamp) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - - var r1 error - if rf, ok := ret.Get(1).(func(string, []string, int64) error); ok { - r1 = rf(userId, channelIds, timestamp) - } else { - r1 = ret.Error(1) - } - - return r0, r1 -} - // DeleteMembershipForUser provides a mock function with given fields: userId, postID func (_m *ThreadStore) DeleteMembershipForUser(userId string, postID string) error { ret := _m.Called(userId, postID) @@ -301,13 +278,13 @@ func (_m *ThreadStore) MaintainMembership(userID string, postID string, opts sto return r0, r1 } -// MarkAllAsRead provides a mock function with given fields: userID, teamID -func (_m *ThreadStore) MarkAllAsRead(userID string, teamID string) error { - ret := _m.Called(userID, teamID) +// MarkAllAsRead provides a mock function with given fields: userID, threadIds +func (_m *ThreadStore) MarkAllAsRead(userID string, threadIds []string) error { + ret := _m.Called(userID, threadIds) var r0 error - if rf, ok := ret.Get(0).(func(string, string) error); ok { - r0 = rf(userID, teamID) + if rf, ok := ret.Get(0).(func(string, []string) error); ok { + r0 = rf(userID, threadIds) } else { r0 = ret.Error(0) } @@ -315,8 +292,8 @@ func (_m *ThreadStore) MarkAllAsRead(userID string, teamID string) error { return r0 } -// MarkAllAsReadInChannels provides a mock function with given fields: userID, channelIDs -func (_m *ThreadStore) MarkAllAsReadInChannels(userID string, channelIDs []string) error { +// MarkAllAsReadByChannels provides a mock function with given fields: userID, channelIDs +func (_m *ThreadStore) MarkAllAsReadByChannels(userID string, channelIDs []string) error { ret := _m.Called(userID, channelIDs) var r0 error @@ -329,6 +306,20 @@ func (_m *ThreadStore) MarkAllAsReadInChannels(userID string, channelIDs []strin return r0 } +// MarkAllAsReadByTeam provides a mock function with given fields: userID, teamID +func (_m *ThreadStore) MarkAllAsReadByTeam(userID string, teamID string) error { + ret := _m.Called(userID, teamID) + + var r0 error + if rf, ok := ret.Get(0).(func(string, string) error); ok { + r0 = rf(userID, teamID) + } else { + r0 = ret.Error(0) + } + + return r0 +} + // MarkAsRead provides a mock function with given fields: userID, threadID, timestamp func (_m *ThreadStore) MarkAsRead(userID string, threadID string, timestamp int64) error { ret := _m.Called(userID, threadID, timestamp) @@ -399,20 +390,6 @@ func (_m *ThreadStore) PermanentDeleteBatchThreadMembershipsForRetentionPolicies return r0, r1, r2 } -// UpdateLastViewedByThreadIds provides a mock function with given fields: userId, threadIds, timestamp -func (_m *ThreadStore) UpdateLastViewedByThreadIds(userId string, threadIds []string, timestamp int64) error { - ret := _m.Called(userId, threadIds, timestamp) - - var r0 error - if rf, ok := ret.Get(0).(func(string, []string, int64) error); ok { - r0 = rf(userId, threadIds, timestamp) - } else { - r0 = ret.Error(0) - } - - return r0 -} - // UpdateMembership provides a mock function with given fields: membership func (_m *ThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { ret := _m.Called(membership) diff --git a/store/storetest/thread_store.go b/store/storetest/thread_store.go index 7a863e97de..933d78f3b9 100644 --- a/store/storetest/thread_store.go +++ b/store/storetest/thread_store.go @@ -6,7 +6,6 @@ package storetest import ( "context" "testing" - "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -24,6 +23,7 @@ func TestThreadStore(t *testing.T, ss store.Store, s SqlStore) { testThreadStorePermanentDeleteBatchThreadMembershipsForRetentionPolicies(t, ss, s) }) t.Run("GetTeamsUnreadForUser", func(t *testing.T) { testGetTeamsUnreadForUser(t, ss) }) + t.Run("MarkAllAsReadByChannels", func(t *testing.T) { testMarkAllAsReadByChannels(t, ss) }) } func testThreadStorePopulation(t *testing.T, ss store.Store) { @@ -262,33 +262,6 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) { require.Nil(t, thread2) }) - t.Run("Thread last updated is changed when channel is updated after UpdateLastViewedAt", func(t *testing.T) { - newPosts := makeSomePosts() - 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) - m.LastUpdated -= 1000 - _, err := ss.Thread().UpdateMembership(m) - require.NoError(t, err) - - _, err = ss.Channel().UpdateLastViewedAt([]string{newPosts[0].ChannelId}, newPosts[0].UserId, true) - require.NoError(t, err) - - assert.Eventually(t, func() bool { - m2, err2 := ss.Thread().GetMembershipForUser(newPosts[0].UserId, newPosts[0].Id) - require.NoError(t, err2) - return m2.LastUpdated > m.LastUpdated - }, time.Second, 10*time.Millisecond) - }) - t.Run("Thread membership 'viewed' timestamp is updated properly", func(t *testing.T) { newPosts := makeSomePosts() @@ -421,6 +394,8 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) { } func threadStoreCreateReply(t *testing.T, ss store.Store, channelID, postID, userID string, createAt int64) *model.Post { + t.Helper() + reply, err := ss.Post().Save(&model.Post{ ChannelId: channelID, UserId: userID, @@ -702,3 +677,147 @@ func testGetTeamsUnreadForUser(t *testing.T, ss store.Store) { assert.Equal(t, int64(1), teamsUnread[team2.Id].ThreadCount) assert.Equal(t, int64(1), teamsUnread[team2.Id].ThreadMentionCount) } + +func testMarkAllAsReadByChannels(t *testing.T, ss store.Store) { + postingUserId := model.NewId() + userAID := model.NewId() + userBID := model.NewId() + + team1, err := ss.Team().Save(&model.Team{ + DisplayName: "Team1", + Name: "team" + model.NewId(), + Email: MakeEmail(), + Type: model.TeamOpen, + }) + require.NoError(t, err) + + channel1, err := ss.Channel().Save(&model.Channel{ + TeamId: team1.Id, + DisplayName: "Channel1", + Name: "channel1" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + + channel2, err := ss.Channel().Save(&model.Channel{ + TeamId: team1.Id, + DisplayName: "Channel2", + Name: "channel2" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + + createThreadMembership := func(userID, postID string) { + t.Helper() + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + _, err := ss.Thread().MaintainMembership(userID, postID, opts) + require.NoError(t, err) + } + + assertThreadReplyCount := func(t *testing.T, userID string, count int64) { + t.Helper() + + teamsUnread, err := ss.Thread().GetTeamsUnreadForUser(userID, []string{team1.Id}) + require.NoError(t, err) + require.Len(t, teamsUnread, 1, "unexpected unread teams count") + assert.Equal(t, count, teamsUnread[team1.Id].ThreadCount, "unexpected thread count") + } + + t.Run("empty set of channels", func(t *testing.T) { + err := ss.Thread().MarkAllAsReadByChannels(model.NewId(), []string{}) + require.NoError(t, err) + }) + + t.Run("single channel", func(t *testing.T) { + post, err := ss.Post().Save(&model.Post{ + ChannelId: channel1.Id, + UserId: postingUserId, + Message: "Root", + }) + require.NoError(t, err) + + _, err = ss.Post().Save(&model.Post{ + ChannelId: channel1.Id, + UserId: postingUserId, + RootId: post.Id, + Message: "Reply", + }) + require.NoError(t, err) + + createThreadMembership(userAID, post.Id) + createThreadMembership(userBID, post.Id) + + assertThreadReplyCount(t, userAID, 1) + assertThreadReplyCount(t, userBID, 1) + + err = ss.Thread().MarkAllAsReadByChannels(userAID, []string{channel1.Id}) + require.NoError(t, err) + + assertThreadReplyCount(t, userAID, 0) + assertThreadReplyCount(t, userBID, 1) + + err = ss.Thread().MarkAllAsReadByChannels(userBID, []string{channel1.Id}) + require.NoError(t, err) + + assertThreadReplyCount(t, userAID, 0) + assertThreadReplyCount(t, userBID, 0) + }) + + t.Run("multiple channels", func(t *testing.T) { + post1, err := ss.Post().Save(&model.Post{ + ChannelId: channel1.Id, + UserId: postingUserId, + Message: "Root", + }) + require.NoError(t, err) + + _, err = ss.Post().Save(&model.Post{ + ChannelId: channel1.Id, + UserId: postingUserId, + RootId: post1.Id, + Message: "Reply", + }) + require.NoError(t, err) + + post2, err := ss.Post().Save(&model.Post{ + ChannelId: channel2.Id, + UserId: postingUserId, + Message: "Root", + }) + require.NoError(t, err) + + _, err = ss.Post().Save(&model.Post{ + ChannelId: channel2.Id, + UserId: postingUserId, + RootId: post2.Id, + Message: "Reply", + }) + require.NoError(t, err) + + createThreadMembership(userAID, post1.Id) + createThreadMembership(userBID, post1.Id) + createThreadMembership(userAID, post2.Id) + createThreadMembership(userBID, post2.Id) + + assertThreadReplyCount(t, userAID, 2) + assertThreadReplyCount(t, userBID, 2) + + err = ss.Thread().MarkAllAsReadByChannels(userAID, []string{channel1.Id, channel2.Id}) + require.NoError(t, err) + + assertThreadReplyCount(t, userAID, 0) + assertThreadReplyCount(t, userBID, 2) + + err = ss.Thread().MarkAllAsReadByChannels(userBID, []string{channel1.Id, channel2.Id}) + require.NoError(t, err) + + assertThreadReplyCount(t, userAID, 0) + assertThreadReplyCount(t, userBID, 0) + }) +} diff --git a/store/timerlayer/timerlayer.go b/store/timerlayer/timerlayer.go index 66ce872576..c1ae4fbd64 100644 --- a/store/timerlayer/timerlayer.go +++ b/store/timerlayer/timerlayer.go @@ -2110,10 +2110,10 @@ func (s *TimerLayerChannelStore) Update(channel *model.Channel) (*model.Channel, return result, err } -func (s *TimerLayerChannelStore) UpdateLastViewedAt(channelIds []string, userID string, updateThreads bool) (map[string]int64, error) { +func (s *TimerLayerChannelStore) UpdateLastViewedAt(channelIds []string, userID string) (map[string]int64, error) { start := timemodule.Now() - result, err := s.ChannelStore.UpdateLastViewedAt(channelIds, userID, updateThreads) + result, err := s.ChannelStore.UpdateLastViewedAt(channelIds, userID) elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) if s.Root.Metrics != nil { @@ -8295,22 +8295,6 @@ func (s *TimerLayerTermsOfServiceStore) Save(termsOfService *model.TermsOfServic return result, err } -func (s *TimerLayerThreadStore) CollectThreadsWithNewerReplies(userId string, channelIds []string, timestamp int64) ([]string, error) { - start := timemodule.Now() - - result, err := s.ThreadStore.CollectThreadsWithNewerReplies(userId, channelIds, timestamp) - - elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) - if s.Root.Metrics != nil { - success := "false" - if err == nil { - success = "true" - } - s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.CollectThreadsWithNewerReplies", success, elapsed) - } - return result, err -} - func (s *TimerLayerThreadStore) DeleteMembershipForUser(userId string, postID string) error { start := timemodule.Now() @@ -8503,10 +8487,10 @@ func (s *TimerLayerThreadStore) MaintainMembership(userID string, postID string, return result, err } -func (s *TimerLayerThreadStore) MarkAllAsRead(userID string, teamID string) error { +func (s *TimerLayerThreadStore) MarkAllAsRead(userID string, threadIds []string) error { start := timemodule.Now() - err := s.ThreadStore.MarkAllAsRead(userID, teamID) + err := s.ThreadStore.MarkAllAsRead(userID, threadIds) elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) if s.Root.Metrics != nil { @@ -8519,10 +8503,10 @@ func (s *TimerLayerThreadStore) MarkAllAsRead(userID string, teamID string) erro return err } -func (s *TimerLayerThreadStore) MarkAllAsReadInChannels(userID string, channelIDs []string) error { +func (s *TimerLayerThreadStore) MarkAllAsReadByChannels(userID string, channelIDs []string) error { start := timemodule.Now() - err := s.ThreadStore.MarkAllAsReadInChannels(userID, channelIDs) + err := s.ThreadStore.MarkAllAsReadByChannels(userID, channelIDs) elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) if s.Root.Metrics != nil { @@ -8530,7 +8514,23 @@ func (s *TimerLayerThreadStore) MarkAllAsReadInChannels(userID string, channelID if err == nil { success = "true" } - s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.MarkAllAsReadInChannels", success, elapsed) + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.MarkAllAsReadByChannels", success, elapsed) + } + return err +} + +func (s *TimerLayerThreadStore) MarkAllAsReadByTeam(userID string, teamID string) error { + start := timemodule.Now() + + err := s.ThreadStore.MarkAllAsReadByTeam(userID, teamID) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.MarkAllAsReadByTeam", success, elapsed) } return err } @@ -8583,22 +8583,6 @@ func (s *TimerLayerThreadStore) PermanentDeleteBatchThreadMembershipsForRetentio return result, resultVar1, err } -func (s *TimerLayerThreadStore) UpdateLastViewedByThreadIds(userId string, threadIds []string, timestamp int64) error { - start := timemodule.Now() - - err := s.ThreadStore.UpdateLastViewedByThreadIds(userId, threadIds, timestamp) - - elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) - if s.Root.Metrics != nil { - success := "false" - if err == nil { - success = "true" - } - s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.UpdateLastViewedByThreadIds", success, elapsed) - } - return err -} - func (s *TimerLayerThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { start := timemodule.Now()