diff --git a/app/channel.go b/app/channel.go index 66c9cd3047..3eaf2f5ef6 100644 --- a/app/channel.go +++ b/app/channel.go @@ -2503,7 +2503,7 @@ func (a *App) MarkChannelAsUnreadFromPost(postID string, userID string, collapse if threadMembership == nil { opts := store.ThreadMembershipOpts{ Following: followThread, - IncrementMentions: true, + IncrementMentions: false, UpdateFollowing: true, UpdateViewedTimestamp: true, UpdateParticipants: false, diff --git a/app/channel_test.go b/app/channel_test.go index 5aacfba1db..69796ec606 100644 --- a/app/channel_test.go +++ b/app/channel_test.go @@ -2212,37 +2212,136 @@ func TestMarkUnreadWithThreads(t *testing.T) { *cfg.ServiceSettings.CollapsedThreads = model.COLLAPSED_THREADS_DEFAULT_ON }) - rootPost, appErr := th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) - require.Nil(t, appErr) - replyPost, appErr := th.App.CreatePost(th.Context, &model.Post{RootId: rootPost.Id, UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) - require.Nil(t, appErr) - threads, appErr := th.App.GetThreadsForUser(th.BasicUser.Id, th.BasicTeam.Id, model.GetUserThreadsOpts{}) - require.Nil(t, appErr) - require.Zero(t, threads.Total) + t.Run("Follow threads only if specified", func(t *testing.T) { + rootPost, appErr := th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + replyPost, appErr := th.App.CreatePost(th.Context, &model.Post{RootId: rootPost.Id, UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + threads, appErr := th.App.GetThreadsForUser(th.BasicUser.Id, th.BasicTeam.Id, model.GetUserThreadsOpts{}) + require.Nil(t, appErr) + require.Zero(t, threads.Total) - _, appErr = th.App.MarkChannelAsUnreadFromPost(replyPost.Id, th.BasicUser.Id, true, true) - require.Nil(t, appErr) + _, appErr = th.App.MarkChannelAsUnreadFromPost(replyPost.Id, th.BasicUser.Id, true, true) + require.Nil(t, appErr) - threads, appErr = th.App.GetThreadsForUser(th.BasicUser.Id, th.BasicTeam.Id, model.GetUserThreadsOpts{}) - require.Nil(t, appErr) - require.NotZero(t, threads.Total) + threads, appErr = th.App.GetThreadsForUser(th.BasicUser.Id, th.BasicTeam.Id, model.GetUserThreadsOpts{}) + require.Nil(t, appErr) + require.NotZero(t, threads.Total) - threadMembership, appErr := th.App.GetThreadMembershipForUser(th.BasicUser.Id, replyPost.RootId) - require.Nil(t, appErr) - require.NotNil(t, threadMembership) - assert.True(t, threadMembership.Following) + threadMembership, appErr := th.App.GetThreadMembershipForUser(th.BasicUser.Id, replyPost.RootId) + require.Nil(t, appErr) + require.NotNil(t, threadMembership) + assert.True(t, threadMembership.Following) - // Create a new thread - rootPost, appErr = th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi2"}, th.BasicChannel, false, false) - require.Nil(t, appErr) - replyPost, appErr = th.App.CreatePost(th.Context, &model.Post{RootId: rootPost.Id, UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi2"}, th.BasicChannel, false, false) - require.Nil(t, appErr) + // Create a new thread + rootPost, appErr = th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi2"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + replyPost, appErr = th.App.CreatePost(th.Context, &model.Post{RootId: rootPost.Id, UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi2"}, th.BasicChannel, false, false) + require.Nil(t, appErr) - _, appErr = th.App.MarkChannelAsUnreadFromPost(replyPost.Id, th.BasicUser.Id, true, false) - require.Nil(t, appErr) + _, appErr = th.App.MarkChannelAsUnreadFromPost(replyPost.Id, th.BasicUser.Id, true, false) + require.Nil(t, appErr) - threadMembership, appErr = th.App.GetThreadMembershipForUser(th.BasicUser.Id, replyPost.RootId) - require.Nil(t, appErr) - require.NotNil(t, threadMembership) - assert.False(t, threadMembership.Following) + threadMembership, appErr = th.App.GetThreadMembershipForUser(th.BasicUser.Id, replyPost.RootId) + require.Nil(t, appErr) + require.NotNil(t, threadMembership) + assert.False(t, threadMembership.Following) + }) + + t.Run("Set unread mentions correctly", func(t *testing.T) { + t.Run("Never followed root post with no replies or mentions", func(t *testing.T) { + rootPost, appErr := th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + _, appErr = th.App.MarkChannelAsUnreadFromPost(rootPost.Id, th.BasicUser.Id, true, true) + require.Nil(t, appErr) + + threadMembership, appErr := th.App.GetThreadMembershipForUser(th.BasicUser.Id, rootPost.Id) + require.Nil(t, appErr) + require.NotNil(t, threadMembership) + assert.Zero(t, threadMembership.UnreadMentions) + }) + + t.Run("Never followed root post with replies and no mentions", func(t *testing.T) { + rootPost, appErr := th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + _, appErr = th.App.CreatePost(th.Context, &model.Post{RootId: rootPost.Id, UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + _, appErr = th.App.MarkChannelAsUnreadFromPost(rootPost.Id, th.BasicUser.Id, true, true) + require.Nil(t, appErr) + + threadMembership, appErr := th.App.GetThreadMembershipForUser(th.BasicUser.Id, rootPost.Id) + require.Nil(t, appErr) + require.NotNil(t, threadMembership) + assert.Zero(t, threadMembership.UnreadMentions) + }) + + t.Run("Never followed root post with replies and mentions", func(t *testing.T) { + rootPost, appErr := th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + _, appErr = th.App.CreatePost(th.Context, &model.Post{RootId: rootPost.Id, UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi @" + th.BasicUser.Username}, th.BasicChannel, false, false) + require.Nil(t, appErr) + _, appErr = th.App.MarkChannelAsUnreadFromPost(rootPost.Id, th.BasicUser.Id, true, true) + require.Nil(t, appErr) + + threadMembership, appErr := th.App.GetThreadMembershipForUser(th.BasicUser.Id, rootPost.Id) + require.Nil(t, appErr) + require.NotNil(t, threadMembership) + assert.Equal(t, int64(1), threadMembership.UnreadMentions) + }) + + t.Run("Previously followed root post with no replies or mentions", func(t *testing.T) { + rootPost, appErr := th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + appErr = th.App.UpdateThreadFollowForUser(th.BasicUser.Id, th.BasicTeam.Id, rootPost.Id, true) + require.Nil(t, appErr) + appErr = th.App.UpdateThreadFollowForUser(th.BasicUser.Id, th.BasicTeam.Id, rootPost.Id, false) + require.Nil(t, appErr) + + _, appErr = th.App.MarkChannelAsUnreadFromPost(rootPost.Id, th.BasicUser.Id, true, true) + require.Nil(t, appErr) + + threadMembership, appErr := th.App.GetThreadMembershipForUser(th.BasicUser.Id, rootPost.Id) + require.Nil(t, appErr) + require.NotNil(t, threadMembership) + assert.Zero(t, threadMembership.UnreadMentions) + }) + + t.Run("Previously followed root post with replies and no mentions", func(t *testing.T) { + rootPost, appErr := th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + _, appErr = th.App.CreatePost(th.Context, &model.Post{RootId: rootPost.Id, UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + appErr = th.App.UpdateThreadFollowForUser(th.BasicUser.Id, th.BasicTeam.Id, rootPost.Id, true) + require.Nil(t, appErr) + appErr = th.App.UpdateThreadFollowForUser(th.BasicUser.Id, th.BasicTeam.Id, rootPost.Id, false) + require.Nil(t, appErr) + + _, appErr = th.App.MarkChannelAsUnreadFromPost(rootPost.Id, th.BasicUser.Id, true, true) + require.Nil(t, appErr) + + threadMembership, appErr := th.App.GetThreadMembershipForUser(th.BasicUser.Id, rootPost.Id) + require.Nil(t, appErr) + require.NotNil(t, threadMembership) + assert.Zero(t, threadMembership.UnreadMentions) + }) + + t.Run("Previously followed root post with replies and mentions", func(t *testing.T) { + rootPost, appErr := th.App.CreatePost(th.Context, &model.Post{UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi"}, th.BasicChannel, false, false) + require.Nil(t, appErr) + _, appErr = th.App.CreatePost(th.Context, &model.Post{RootId: rootPost.Id, UserId: th.BasicUser2.Id, CreateAt: model.GetMillis(), ChannelId: th.BasicChannel.Id, Message: "hi @" + th.BasicUser.Username}, th.BasicChannel, false, false) + require.Nil(t, appErr) + appErr = th.App.UpdateThreadFollowForUser(th.BasicUser.Id, th.BasicTeam.Id, rootPost.Id, true) + require.Nil(t, appErr) + appErr = th.App.UpdateThreadFollowForUser(th.BasicUser.Id, th.BasicTeam.Id, rootPost.Id, false) + require.Nil(t, appErr) + + _, appErr = th.App.MarkChannelAsUnreadFromPost(rootPost.Id, th.BasicUser.Id, true, true) + require.Nil(t, appErr) + + threadMembership, appErr := th.App.GetThreadMembershipForUser(th.BasicUser.Id, rootPost.Id) + require.Nil(t, appErr) + require.NotNil(t, threadMembership) + assert.Equal(t, int64(1), threadMembership.UnreadMentions) + }) + }) } diff --git a/app/user.go b/app/user.go index 4eb072e583..a733f65cb6 100644 --- a/app/user.go +++ b/app/user.go @@ -2228,14 +2228,17 @@ func (a *App) UpdateThreadFollowForUser(userID, teamID, threadID string, state b return model.NewAppError("UpdateThreadFollowForUser", "app.user.update_thread_follow_for_user.app_error", nil, err.Error(), http.StatusInternalServerError) } thread, err := a.Srv().Store.Thread().Get(threadID) - if err != nil { return model.NewAppError("UpdateThreadFollowForUser", "app.user.update_thread_follow_for_user.app_error", nil, err.Error(), http.StatusInternalServerError) } + replyCount := int64(0) + if thread != nil { + replyCount = thread.ReplyCount + } message := model.NewWebSocketEvent(model.WEBSOCKET_EVENT_THREAD_FOLLOW_CHANGED, teamID, "", userID, nil) message.Add("thread_id", threadID) message.Add("state", state) - message.Add("reply_count", thread.ReplyCount) + message.Add("reply_count", replyCount) a.Publish(message) return nil } diff --git a/store/sqlstore/thread_store.go b/store/sqlstore/thread_store.go index 5fb43b41d6..5eae3c8564 100644 --- a/store/sqlstore/thread_store.go +++ b/store/sqlstore/thread_store.go @@ -110,7 +110,7 @@ func (s *SqlThreadStore) get(ex gorp.SqlExecutor, id string) (*model.Thread, err err := ex.SelectOne(&thread, query, args...) if err != nil { if err == sql.ErrNoRows { - return nil, store.NewErrNotFound("Thread", id) + return nil, nil } return nil, errors.Wrapf(err, "failed to get thread with id=%s", id) @@ -626,7 +626,7 @@ func (s *SqlThreadStore) MaintainMembership(userId, postId string, opts store.Th if getErr != nil { return nil, getErr } - if !thread.Participants.Contains(userId) { + if thread != nil && !thread.Participants.Contains(userId) { thread.Participants = append(thread.Participants, userId) if _, err = s.update(trx, thread); err != nil { return nil, err diff --git a/store/storetest/thread_store.go b/store/storetest/thread_store.go index 46e380f282..1eaba518b5 100644 --- a/store/storetest/thread_store.go +++ b/store/storetest/thread_store.go @@ -469,8 +469,9 @@ func testThreadStorePermanentDeleteBatchForRetentionPolicies(t *testing.T, ss st nowMillis := thread.LastReplyAt + *channelPolicy.PostDuration*24*60*60*1000 + 1 _, _, err = ss.Thread().PermanentDeleteBatchForRetentionPolicies(nowMillis, 0, limit, model.RetentionPolicyCursor{}) require.NoError(t, err) - _, err = ss.Thread().Get(post.Id) - require.Error(t, err, "thread should have been deleted by channel policy") + thread, err = ss.Thread().Get(post.Id) + assert.NoError(t, err) + assert.Nil(t, thread, "thread should have been deleted by channel policy") // create a new thread threadStoreCreateReply(t, ss, channel.Id, post.Id, 2000) @@ -498,8 +499,9 @@ func testThreadStorePermanentDeleteBatchForRetentionPolicies(t *testing.T, ss st require.NoError(t, err) _, _, err = ss.Thread().PermanentDeleteBatchForRetentionPolicies(nowMillis, 0, limit, model.RetentionPolicyCursor{}) require.NoError(t, err) - _, err = ss.Thread().Get(post.Id) - require.Error(t, err, "thread should have been deleted by team policy") + thread, err = ss.Thread().Get(post.Id) + assert.NoError(t, err) + assert.Nil(t, thread, "thread should have been deleted by team policy") } func testThreadStorePermanentDeleteBatchThreadMembershipsForRetentionPolicies(t *testing.T, ss store.Store) {