diff --git a/app/user.go b/app/user.go index 4680c226f2..4eb072e583 100644 --- a/app/user.go +++ b/app/user.go @@ -2245,10 +2245,16 @@ func (a *App) UpdateThreadReadForUser(userID, teamID, threadID string, timestamp if err != nil { return nil, err } - membership, nErr := a.Srv().Store.Thread().GetMembershipForUser(userID, threadID) - if nErr != nil { - return nil, model.NewAppError("UpdateThreadsReadForUser", "app.user.update_threads_read_for_user.app_error", nil, nErr.Error(), http.StatusInternalServerError) + + opts := store.ThreadMembershipOpts{ + Following: true, + UpdateFollowing: false, } + membership, storeErr := a.Srv().Store.Thread().MaintainMembership(userID, threadID, opts) + if storeErr != nil { + return nil, model.NewAppError("UpdateThreadReadForUser", "app.user.update_thread_read_for_user.app_error", nil, err.Error(), http.StatusInternalServerError) + } + post, err := a.GetSinglePost(threadID) if err != nil { return nil, err @@ -2258,9 +2264,9 @@ func (a *App) UpdateThreadReadForUser(userID, teamID, threadID string, timestamp return nil, err } membership.Following = true - _, nErr = a.Srv().Store.Thread().UpdateMembership(membership) + _, nErr := a.Srv().Store.Thread().UpdateMembership(membership) if nErr != nil { - return nil, model.NewAppError("UpdateThreadsReadForUser", "app.user.update_threads_read_for_user.app_error", nil, nErr.Error(), http.StatusInternalServerError) + return nil, model.NewAppError("UpdateThreadReadForUser", "app.user.update_thread_read_for_user.app_error", nil, nErr.Error(), http.StatusInternalServerError) } membership.LastViewed = timestamp diff --git a/app/user_test.go b/app/user_test.go index 44ba323df6..cf0665591f 100644 --- a/app/user_test.go +++ b/app/user_test.go @@ -8,6 +8,7 @@ import ( "context" "encoding/json" "errors" + "os" "strings" "testing" "time" @@ -1461,3 +1462,36 @@ func TestPatchUser(t *testing.T) { require.Nil(t, err) }) } + +func TestUpdateThreadReadForUser(t *testing.T) { + os.Setenv("MM_FEATUREFLAGS_COLLAPSEDTHREADS", "true") + defer os.Unsetenv("MM_FEATUREFLAGS_COLLAPSEDTHREADS") + th := Setup(t).InitBasic() + defer th.TearDown() + th.App.UpdateConfig(func(cfg *model.Config) { + *cfg.ServiceSettings.ThreadAutoFollow = true + *cfg.ServiceSettings.CollapsedThreads = model.COLLAPSED_THREADS_DEFAULT_ON + }) + + t.Run("Ensure thread membership is created and followed", 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.UpdateThreadReadForUser(th.BasicUser.Id, th.BasicChannel.TeamId, rootPost.Id, replyPost.CreateAt) + require.Nil(t, appErr) + + threads, appErr = th.App.GetThreadsForUser(th.BasicUser.Id, th.BasicTeam.Id, model.GetUserThreadsOpts{}) + require.Nil(t, appErr) + assert.NotZero(t, threads.Total) + + threadMembership, appErr := th.App.GetThreadMembershipForUser(th.BasicUser.Id, rootPost.Id) + require.Nil(t, appErr) + require.NotNil(t, threadMembership) + assert.True(t, threadMembership.Following) + }) +}