From d5cc2eb2f607c592620e4f91a9fc7435033f8e78 Mon Sep 17 00:00:00 2001 From: Ibrahim Serdar Acikgoz Date: Thu, 29 Aug 2024 14:06:41 +0200 Subject: [PATCH] [MM-59367] export: enable exporting thread followers for CRT (#27623) --- server/channels/app/export.go | 34 ++ server/channels/app/export_converters.go | 8 + server/channels/app/export_test.go | 102 ++++++ server/channels/app/helper_test.go | 20 ++ server/channels/app/import_functions.go | 166 +++++++++- server/channels/app/import_functions_test.go | 284 +++++++++++++++- server/channels/app/imports/import_types.go | 12 + .../opentracinglayer/opentracinglayer.go | 54 +++ .../channels/store/retrylayer/retrylayer.go | 63 ++++ .../channels/store/sqlstore/thread_store.go | 214 +++++++++--- server/channels/store/store.go | 14 + .../store/storetest/mocks/ThreadStore.go | 90 +++++ .../channels/store/storetest/thread_store.go | 309 ++++++++++++++++++ .../channels/store/timerlayer/timerlayer.go | 48 +++ server/i18n/en.json | 16 + server/public/model/thread.go | 20 ++ server/public/model/thread_test.go | 39 +++ 17 files changed, 1430 insertions(+), 63 deletions(-) create mode 100644 server/public/model/thread_test.go diff --git a/server/channels/app/export.go b/server/channels/app/export.go index a62facc9db..e4e1fa54e8 100644 --- a/server/channels/app/export.go +++ b/server/channels/app/export.go @@ -607,6 +607,15 @@ func (a *App) exportAllPosts(ctx request.CTX, job *model.Job, writer io.Writer, return nil, err } + followers, err := a.buildThreadFollowers(ctx, post.Id) + if err != nil { + return nil, err + } + + if len(followers) > 0 { + postLine.Post.ThreadFollowers = &followers + } + if withAttachments && len(replyAttachments) > 0 { attachments = append(attachments, replyAttachments...) } @@ -674,6 +683,21 @@ func (a *App) buildPostReplies(ctx request.CTX, postID string, withAttachments b return replies, attachments, nil } +func (a *App) buildThreadFollowers(_ request.CTX, postID string) ([]imports.ThreadFollowerImportData, *model.AppError) { + var followers []imports.ThreadFollowerImportData + + threadFollowers, nErr := a.Srv().Store().Thread().GetThreadMembershipsForExport(postID) + if nErr != nil { + return nil, model.NewAppError("buildThreadFollowers", "app.thread.get_threadmembers_for_export.app_error", nil, "", http.StatusInternalServerError).Wrap(nErr) + } + + for _, member := range threadFollowers { + followers = append(followers, *ImportFollowerFromThreadMember(member)) + } + + return followers, nil +} + func (a *App) BuildPostReactions(ctx request.CTX, postID string) (*[]ReactionImportData, *model.AppError) { var reactionsOfPost []imports.ReactionImportData @@ -978,6 +1002,16 @@ func (a *App) exportAllDirectPosts(ctx request.CTX, job *model.Job, writer io.Wr if len(postAttachments) > 0 { postLine.DirectPost.Attachments = &postAttachments } + + followers, err := a.buildThreadFollowers(ctx, post.Id) + if err != nil { + return nil, err + } + + if len(followers) > 0 { + postLine.DirectPost.ThreadFollowers = &followers + } + if err := a.exportWriteLine(writer, postLine); err != nil { return nil, err } diff --git a/server/channels/app/export_converters.go b/server/channels/app/export_converters.go index fe8f804893..5e9c34fbe4 100644 --- a/server/channels/app/export_converters.go +++ b/server/channels/app/export_converters.go @@ -342,3 +342,11 @@ func ImportLineFromScheme(scheme *model.Scheme, rolesMap map[string]*model.Role) Scheme: data, } } + +func ImportFollowerFromThreadMember(threadMember *model.ThreadMembershipForExport) *imports.ThreadFollowerImportData { + return &imports.ThreadFollowerImportData{ + User: &threadMember.Username, + LastViewed: &threadMember.LastViewed, + UnreadMentions: &threadMember.UnreadMentions, + } +} diff --git a/server/channels/app/export_test.go b/server/channels/app/export_test.go index d9328ecded..95db4c40e2 100644 --- a/server/channels/app/export_test.go +++ b/server/channels/app/export_test.go @@ -4,7 +4,9 @@ package app import ( + "bufio" "bytes" + "encoding/json" "fmt" "os" "path/filepath" @@ -16,6 +18,7 @@ import ( "github.com/stretchr/testify/require" "github.com/mattermost/mattermost/server/public/model" + "github.com/mattermost/mattermost/server/v8/channels/app/imports" "github.com/mattermost/mattermost/server/v8/channels/utils" "github.com/mattermost/mattermost/server/v8/channels/utils/fileutils" ) @@ -647,6 +650,105 @@ func TestExportDMPostWithSelf(t *testing.T) { assert.Equal(t, th1.BasicUser.Username, (*posts[0].ChannelMembers)[0]) } +func TestExportPostsWithThread(t *testing.T) { + th1 := Setup(t).InitBasic() + defer th1.TearDown() + + assertThreadFollowers := func(t *testing.T, b *bytes.Buffer, postCreateAt int64, userNames []string) { + scanner := bufio.NewScanner(b) + + usersToAssert := make([]string, 0) + + for scanner.Scan() { + var line imports.LineImportData + err := json.Unmarshal(scanner.Bytes(), &line) + require.NoError(t, err) + + switch line.Type { + case "post": + postLine := line.Post + require.NotNil(t, postLine) + + if postLine.CreateAt != nil && *postLine.CreateAt != postCreateAt { + continue + } + + for _, follower := range *postLine.ThreadFollowers { + if follower.User == nil { + require.Fail(t, "follower.User is nil") + } + + usersToAssert = append(usersToAssert, *follower.User) + } + case "direct_post": + postLine := line.DirectPost + require.NotNil(t, postLine) + + if postLine.CreateAt != nil && *postLine.CreateAt != postCreateAt { + continue + } + + for _, follower := range *postLine.ThreadFollowers { + if follower.User == nil { + require.Fail(t, "follower.User is nil") + } + + usersToAssert = append(usersToAssert, *follower.User) + } + default: + continue + } + } + + require.ElementsMatch(t, userNames, usersToAssert) + } + + t.Run("Export thread followers for a thread (public channel)", func(t *testing.T) { + thread := th1.CreatePost(th1.BasicChannel) + _ = th1.CreatePostReply(thread) + + appErr := th1.App.UpdateThreadFollowForUser(th1.BasicUser2.Id, th1.BasicTeam.Id, thread.Id, true) + require.Nil(t, appErr) + + member1, appErr := th1.App.GetThreadMembershipForUser(th1.BasicUser.Id, thread.Id) + require.Nil(t, appErr) + require.NotNil(t, member1) + + member2, appErr := th1.App.GetThreadMembershipForUser(th1.BasicUser2.Id, thread.Id) + require.Nil(t, appErr) + require.NotNil(t, member2) + + var b bytes.Buffer + err := th1.App.BulkExport(th1.Context, &b, "somePath", nil, model.BulkExportOpts{}) + require.Nil(t, err) + + assertThreadFollowers(t, &b, thread.CreateAt, []string{th1.BasicUser.Username, th1.BasicUser2.Username}) + }) + + t.Run("Export thread followers for a thread (direct messages)", func(t *testing.T) { + dmc := th1.CreateDmChannel(th1.BasicUser2) + + thread := th1.CreatePost(dmc) + _ = th1.CreatePostReply(thread) + + appErr := th1.App.UpdateThreadFollowForUser(th1.BasicUser2.Id, th1.BasicTeam.Id, thread.Id, true) + require.Nil(t, appErr) + + member1, appErr := th1.App.GetThreadMembershipForUser(th1.BasicUser.Id, thread.Id) + require.Nil(t, appErr) + require.NotNil(t, member1) + + member2, appErr := th1.App.GetThreadMembershipForUser(th1.BasicUser2.Id, thread.Id) + require.Nil(t, appErr) + require.NotNil(t, member2) + + var b bytes.Buffer + err := th1.App.BulkExport(th1.Context, &b, "somePath", nil, model.BulkExportOpts{}) + require.Nil(t, err) + assertThreadFollowers(t, &b, thread.CreateAt, []string{th1.BasicUser.Username, th1.BasicUser2.Username}) + }) +} + func TestBulkExport(t *testing.T) { th := Setup(t) testsDir, _ := fileutils.FindDir("tests") diff --git a/server/channels/app/helper_test.go b/server/channels/app/helper_test.go index 5ceecc74be..6c967973bb 100644 --- a/server/channels/app/helper_test.go +++ b/server/channels/app/helper_test.go @@ -477,6 +477,26 @@ func (th *TestHelper) CreateMessagePost(channel *model.Channel, message string) return post } +func (th *TestHelper) CreatePostReply(root *model.Post) *model.Post { + id := model.NewId() + post := &model.Post{ + UserId: th.BasicUser.Id, + ChannelId: root.ChannelId, + RootId: root.Id, + Message: "message_" + id, + CreateAt: model.GetMillis() - 10000, + } + + ch, err := th.App.GetChannel(th.Context, root.ChannelId) + if err != nil { + panic(err) + } + if post, err = th.App.CreatePost(th.Context, post, ch, false, true); err != nil { + panic(err) + } + return post +} + func (th *TestHelper) LinkUserToTeam(user *model.User, team *model.Team) { _, err := th.App.JoinUserToTeam(th.Context, team, user, "") if err != nil { diff --git a/server/channels/app/import_functions.go b/server/channels/app/import_functions.go index 29544ea01e..faa9b79b3e 100644 --- a/server/channels/app/import_functions.go +++ b/server/channels/app/import_functions.go @@ -1567,11 +1567,13 @@ func (a *App) importMultiplePostLines(rctx request.CTX, lines []imports.LineImpo } var ( - postsWithData = []postAndData{} - postsForCreateList = []*model.Post{} - postsForCreateMap = map[string]int{} - postsForOverwriteList = []*model.Post{} - postsForOverwriteMap = map[string]int{} + postsWithData = []postAndData{} + postsForCreateList = []*model.Post{} + postsForCreateMap = map[string]int{} + postsForOverwriteList = []*model.Post{} + postsForOverwriteMap = map[string]int{} + threadMembersToCreateMap = map[string][]*model.ThreadMembership{} + threadMembersToOverwriteList = []*model.ThreadMembership{} ) for _, line := range lines { @@ -1615,6 +1617,18 @@ func (a *App) importMultiplePostLines(rctx request.CTX, lines []imports.LineImpo if line.Post.IsPinned != nil { post.IsPinned = *line.Post.IsPinned } + if line.Post.ThreadFollowers != nil { + threadMemberships, lineNumber, err := a.extractThreadMembers(&line, users, post) + if err != nil { + return lineNumber, err + } + + if post.Id == "" { + threadMembersToCreateMap[getPostStrID(post)] = threadMemberships + } else { + threadMembersToOverwriteList = append(threadMembersToOverwriteList, threadMemberships...) + } + } fileIDs := a.uploadAttachments(rctx, line.Post.Attachments, post, team.Id, extractContent) for _, fileID := range post.FileIds { @@ -1634,11 +1648,13 @@ func (a *App) importMultiplePostLines(rctx request.CTX, lines []imports.LineImpo postsForOverwriteList = append(postsForOverwriteList, post) postsForOverwriteMap[getPostStrID(post)] = line.LineNumber } + // Tip: the post ID is getting populated after the post is saved, if it's a new post. Otherwise, it's already set. postsWithData = append(postsWithData, postAndData{post: post, postData: line.Post, team: team, lineNumber: line.LineNumber}) } if len(postsForCreateList) > 0 { - if _, idx, nErr := a.Srv().Store().Post().SaveMultiple(postsForCreateList); nErr != nil { + _, idx, nErr := a.Srv().Store().Post().SaveMultiple(postsForCreateList) + if nErr != nil { var appErr *model.AppError var invErr *store.ErrInvalidInput var retErr *model.AppError @@ -1659,6 +1675,36 @@ func (a *App) importMultiplePostLines(rctx request.CTX, lines []imports.LineImpo } return 0, retErr } + + var membersToCreate []*model.ThreadMembership + for _, post := range postsForCreateList { + members, ok := threadMembersToCreateMap[getPostStrID(post)] + if !ok { + continue + } + + for _, member := range members { + if post.Id == "" { + appErr := model.NewAppError("importMultiplePostLines", "app.post.save.thread_membership.app_error", nil, "", http.StatusInternalServerError).Wrap(errors.New("post id cannot be empty")) + if lineNumber, ok := postsForCreateMap[getPostStrID(post)]; ok { + return lineNumber, appErr + } + return 0, appErr + } + member.PostId = post.Id + } + + membersToCreate = append(membersToCreate, members...) + } + + // we have an assumption here is that all these memberships should be brand new because the corresponding posts + // do not exist in the target until the import. + if _, err := a.Srv().Store().Thread().SaveMultipleMemberships(membersToCreate); err != nil { + // we don't know the line number of the post that caused the error + // so we return 0. But at this stage, it's unlikely to receive an error + // due to the thread member itself, most likely it's due to the DB connection etc. + return 0, model.NewAppError("importMultiplePostLines", "app.post.save.thread_membership.app_error", nil, "", http.StatusInternalServerError).Wrap(err) + } } if _, idx, err := a.Srv().Store().Post().OverwriteMultiple(postsForOverwriteList); err != nil { @@ -1671,6 +1717,15 @@ func (a *App) importMultiplePostLines(rctx request.CTX, lines []imports.LineImpo return 0, model.NewAppError("importMultiplePostLines", "app.post.overwrite.app_error", nil, "", http.StatusInternalServerError).Wrap(err) } + // Update thread memberships for posts that were overwritten. Here some of the memberships + // can be brand new, needs to be updated or an older membership should not get updated. + // MaintainMembership method has some logic within to handle those decisions. Unfortunately + // some application code leaked to the store layer here, which should be revisited when there + // is resource (eg. time, human or maybe AI). + if _, sErr := a.Srv().Store().Thread().MaintainMultipleFromImport(threadMembersToOverwriteList); sErr != nil { + return 0, model.NewAppError("importMultiplePostLines", "app.post.save.thread_membership.app_error", nil, "", http.StatusInternalServerError).Wrap(sErr) + } + for _, postWithData := range postsWithData { postWithData := postWithData if postWithData.postData.FlaggedBy != nil { @@ -2008,11 +2063,13 @@ func (a *App) importMultipleDirectPostLines(rctx request.CTX, lines []imports.Li } var ( - postsWithData = []postAndData{} - postsForCreateList = []*model.Post{} - postsForCreateMap = map[string]int{} - postsForOverwriteList = []*model.Post{} - postsForOverwriteMap = map[string]int{} + postsWithData = []postAndData{} + postsForCreateList = []*model.Post{} + postsForCreateMap = map[string]int{} + postsForOverwriteList = []*model.Post{} + postsForOverwriteMap = map[string]int{} + threadMembersToCreateMap = map[string][]*model.ThreadMembership{} + threadMembersToOverwriteList = []*model.ThreadMembership{} ) for _, line := range lines { @@ -2077,6 +2134,18 @@ func (a *App) importMultipleDirectPostLines(rctx request.CTX, lines []imports.Li if line.DirectPost.IsPinned != nil { post.IsPinned = *line.DirectPost.IsPinned } + if line.DirectPost.ThreadFollowers != nil { + threadMemberships, lineNumber, err := a.extractThreadMembers(&line, users, post) + if err != nil { + return lineNumber, err + } + + if post.Id == "" { + threadMembersToCreateMap[getPostStrID(post)] = threadMemberships + } else { + threadMembersToOverwriteList = append(threadMembersToOverwriteList, threadMemberships...) + } + } fileIDs := a.uploadAttachments(rctx, line.DirectPost.Attachments, post, "noteam", extractContent) for _, fileID := range post.FileIds { @@ -2121,7 +2190,33 @@ func (a *App) importMultipleDirectPostLines(rctx request.CTX, lines []imports.Li } return 0, retErr } + + var membersToCreate []*model.ThreadMembership + for _, post := range postsForCreateList { + members, ok := threadMembersToCreateMap[getPostStrID(post)] + if !ok { + continue + } + + for _, member := range members { + if post.Id == "" { + appErr := model.NewAppError("importMultiplePostLines", "app.post.save.thread_membership.app_error", nil, "", http.StatusInternalServerError).Wrap(errors.New("post id cannot be empty")) + if lineNumber, ok := postsForCreateMap[getPostStrID(post)]; ok { + return lineNumber, appErr + } + return 0, appErr + } + member.PostId = post.Id + } + + membersToCreate = append(membersToCreate, members...) + } + + if _, err := a.Srv().Store().Thread().SaveMultipleMemberships(membersToCreate); err != nil { + return 0, model.NewAppError("importMultiplePostLines", "app.post.save.thread_membership.app_error", nil, "", http.StatusInternalServerError).Wrap(err) + } } + if _, idx, err := a.Srv().Store().Post().OverwriteMultiple(postsForOverwriteList); err != nil { if idx != -1 && idx < len(postsForOverwriteList) { post := postsForOverwriteList[idx] @@ -2132,6 +2227,10 @@ func (a *App) importMultipleDirectPostLines(rctx request.CTX, lines []imports.Li return 0, model.NewAppError("importMultiplePostLines", "app.post.overwrite.app_error", nil, "", http.StatusInternalServerError).Wrap(err) } + if _, sErr := a.Srv().Store().Thread().MaintainMultipleFromImport(threadMembersToOverwriteList); sErr != nil { + return 0, model.NewAppError("importMultiplePostLines", "app.post.save.thread_membership.app_error", nil, "", http.StatusInternalServerError).Wrap(sErr) + } + for _, postWithData := range postsWithData { if postWithData.directPostData.FlaggedBy != nil { var preferences model.Preferences @@ -2240,3 +2339,48 @@ func (a *App) importEmoji(rctx request.CTX, data *imports.EmojiImportData, dryRu return nil } + +func (a *App) extractThreadMembers(line *imports.LineImportWorkerData, users map[string]*model.User, post *model.Post) ([]*model.ThreadMembership, int, *model.AppError) { + threadMemberships := []*model.ThreadMembership{} + + var importedFollowers []imports.ThreadFollowerImportData + if line.Post != nil { + importedFollowers = *line.Post.ThreadFollowers + } else if line.DirectPost != nil { + importedFollowers = *line.DirectPost.ThreadFollowers + } + participants := make([]*model.User, len(importedFollowers)) + + for i, member := range importedFollowers { + user, ok := users[strings.ToLower(*member.User)] + if !ok { + // maybe it's a user on target instance but not in the import data. + // This is a rare case, but we need to or can to handle it. + // alternatively, we can continue and discard this follower as maybe they + // were deleted. + var uErr error + user, uErr = a.Srv().Store().User().GetByUsername(*member.User) + if uErr != nil { + return nil, line.LineNumber, model.NewAppError("importMultiplePostLines", "app.import.get_users_by_username.some_users_not_found.error", nil, "", http.StatusBadRequest).Wrap(uErr) + } + } + membership := &model.ThreadMembership{ + PostId: post.Id, // empty if it's a new post, will set later while inserting to the DB. + UserId: user.Id, + Following: true, + } + + if member.LastViewed != nil { + membership.LastViewed = *member.LastViewed + } + if member.UnreadMentions != nil { + membership.UnreadMentions = *member.UnreadMentions + } + // We only need the user ID to update the thread. + participants[i] = &model.User{Id: user.Id} + threadMemberships = append(threadMemberships, membership) + } + post.Participants = participants + + return threadMemberships, 0, nil +} diff --git a/server/channels/app/import_functions_test.go b/server/channels/app/import_functions_test.go index 149b1a6a80..94700fa38c 100644 --- a/server/channels/app/import_functions_test.go +++ b/server/channels/app/import_functions_test.go @@ -11,6 +11,7 @@ import ( "path/filepath" "strings" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -2103,7 +2104,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) // Try adding a valid post in apply mode. - time := model.GetMillis() + createAt := model.GetMillis() data = imports.LineImportWorkerData{ LineImportData: imports.LineImportData{ Post: &imports.PostImportData{ @@ -2111,7 +2112,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { Channel: &channelName, User: &username, Message: ptrStr("Message"), - CreateAt: &time, + CreateAt: &createAt, }, }, LineNumber: 1, @@ -2122,7 +2123,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { AssertAllPostsCount(t, th.App, initialPostCount, 1, team.Id) // Check the post values. - posts, nErr := th.App.Srv().Store().Post().GetPostsCreatedAt(channel.Id, time) + posts, nErr := th.App.Srv().Store().Post().GetPostsCreatedAt(channel.Id, createAt) require.NoError(t, nErr) require.Len(t, posts, 1, "Unexpected number of posts found.") @@ -2139,7 +2140,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { Channel: &channelName, User: &username, Message: ptrStr("Message"), - CreateAt: &time, + CreateAt: &createAt, }, }, LineNumber: 1, @@ -2150,7 +2151,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { AssertAllPostsCount(t, th.App, initialPostCount, 1, team.Id) // Check the post values. - posts, nErr = th.App.Srv().Store().Post().GetPostsCreatedAt(channel.Id, time) + posts, nErr = th.App.Srv().Store().Post().GetPostsCreatedAt(channel.Id, createAt) require.NoError(t, nErr) require.Len(t, posts, 1, "Unexpected number of posts found.") @@ -2160,7 +2161,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { require.False(t, postBool, "Post properties not as expected") // Save the post with a different time. - newTime := time + 1 + newTime := createAt + 1 data = imports.LineImportWorkerData{ LineImportData: imports.LineImportData{ Post: &imports.PostImportData{ @@ -2186,7 +2187,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { Channel: &channelName, User: &username, Message: ptrStr("Message 2"), - CreateAt: &time, + CreateAt: &createAt, }, }, LineNumber: 1, @@ -2197,7 +2198,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { AssertAllPostsCount(t, th.App, initialPostCount, 3, team.Id) // Test with hashtags - hashtagTime := time + 2 + hashtagTime := createAt + 2 data = imports.LineImportWorkerData{ LineImportData: imports.LineImportData{ Post: &imports.PostImportData{ @@ -2505,7 +2506,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { Channel: &channelName, User: &username, Message: ptrStr("another message"), - CreateAt: &time, + CreateAt: &createAt, }, }, LineNumber: 1, @@ -2517,7 +2518,7 @@ func TestImportimportMultiplePostLines(t *testing.T) { Channel: &channelName, User: &username, Message: ptrStr("another message"), - CreateAt: &time, + CreateAt: &createAt, }, }, LineNumber: 1, @@ -2554,6 +2555,166 @@ func TestImportimportMultiplePostLines(t *testing.T) { // Posts should be added to the right team AssertAllPostsCount(t, th.App, initialPostCountForTeam2, 1, team2.Id) AssertAllPostsCount(t, th.App, initialPostCount, 15, team.Id) + + t.Run("Importing a post with a thread", func(t *testing.T) { + // Create a thread. + importCreate := time.Now().Add(-1 * time.Minute).UnixMilli() + data = imports.LineImportWorkerData{ + LineImportData: imports.LineImportData{ + Post: &imports.PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &user.Username, + Message: ptrStr("Thread Message"), + CreateAt: ptrInt64(importCreate), + Replies: &[]imports.ReplyImportData{{ + User: &user.Username, + Message: ptrStr("Reply"), + CreateAt: ptrInt64(model.GetMillis()), + }}, + ThreadFollowers: &[]imports.ThreadFollowerImportData{{ + User: &user.Username, + LastViewed: ptrInt64(model.GetMillis()), + }, { + User: &user2.Username, + LastViewed: ptrInt64(model.GetMillis()), + }}, + }, + }, + LineNumber: 1, + } + + errLine, err = th.App.importMultiplePostLines(th.Context, []imports.LineImportWorkerData{data}, false, true) + require.Nil(t, err) + require.Equal(t, 0, errLine) + + resultPosts, nErr = th.App.Srv().Store().Post().GetPostsCreatedAt(channel.Id, importCreate) + require.NoError(t, nErr) + require.Equal(t, 1, len(resultPosts)) + + followers, err := th.App.Srv().Store().Thread().GetThreadFollowers(resultPosts[0].Id, true) + require.NoError(t, err) + + assert.ElementsMatch(t, []string{user.Id, user2.Id}, followers) + }) + + t.Run("Importing a post with a non existent follower", func(t *testing.T) { + // Create a thread. + importCreate := time.Now().Add(-1 * time.Minute).UnixMilli() + data = imports.LineImportWorkerData{ + LineImportData: imports.LineImportData{ + Post: &imports.PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &user.Username, + Message: ptrStr("Thread Message"), + CreateAt: ptrInt64(importCreate), + Replies: &[]imports.ReplyImportData{{ + User: &user.Username, + Message: ptrStr("Reply"), + CreateAt: ptrInt64(model.GetMillis()), + }}, + ThreadFollowers: &[]imports.ThreadFollowerImportData{{ + User: &user.Username, + LastViewed: ptrInt64(model.GetMillis()), + }, { + User: ptrStr("invalid.user"), + }}, + }, + }, + LineNumber: 1, + } + + errLine, err = th.App.importMultiplePostLines(th.Context, []imports.LineImportWorkerData{data}, false, true) + require.NotNil(t, err) + require.Equal(t, 1, errLine) + }) + + t.Run("Importing a post with a non existent follower", func(t *testing.T) { + importCreate := time.Now().Add(-1 * time.Minute).UnixMilli() + data = imports.LineImportWorkerData{ + LineImportData: imports.LineImportData{ + Post: &imports.PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &user.Username, + Message: ptrStr("Thread Message"), + CreateAt: ptrInt64(importCreate), + Replies: &[]imports.ReplyImportData{{ + User: &user.Username, + Message: ptrStr("Reply"), + CreateAt: ptrInt64(model.GetMillis()), + }}, + ThreadFollowers: &[]imports.ThreadFollowerImportData{{ + User: &user.Username, + LastViewed: ptrInt64(model.GetMillis()), + }, { + User: ptrStr("invalid.user"), + }}, + }, + }, + LineNumber: 1, + } + + errLine, err = th.App.importMultiplePostLines(th.Context, []imports.LineImportWorkerData{data}, false, true) + require.NotNil(t, err) + require.Equal(t, 1, errLine) + }) + + t.Run("Importing a post with new followers", func(t *testing.T) { + importCreate := time.Now().Add(-5 * time.Minute).UnixMilli() + data = imports.LineImportWorkerData{ + LineImportData: imports.LineImportData{ + Post: &imports.PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + Message: ptrStr("Hello"), + CreateAt: ptrInt64(importCreate), + }, + }, + LineNumber: 1, + } + + errLine, err = th.App.importMultiplePostLines(th.Context, []imports.LineImportWorkerData{data}, false, true) + require.Nil(t, err) + require.Equal(t, 0, errLine) + + resultPosts, nErr = th.App.Srv().Store().Post().GetPostsCreatedAt(channel.Id, importCreate) + require.NoError(t, nErr) + require.Equal(t, 1, len(resultPosts)) + + data = imports.LineImportWorkerData{ + LineImportData: imports.LineImportData{ + Post: &imports.PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &user.Username, + Message: ptrStr("Hello"), + CreateAt: ptrInt64(importCreate), + Replies: &[]imports.ReplyImportData{{ + User: &user.Username, + Message: ptrStr("Reply"), + CreateAt: ptrInt64(model.GetMillis()), + }}, + ThreadFollowers: &[]imports.ThreadFollowerImportData{{ + User: &user.Username, + LastViewed: ptrInt64(model.GetMillis()), + }}, + }, + }, + LineNumber: 1, + } + + errLine, err = th.App.importMultiplePostLines(th.Context, []imports.LineImportWorkerData{data}, false, true) + require.Nil(t, err) + require.Equal(t, 0, errLine) + + followers, err := th.App.Srv().Store().Thread().GetThreadFollowers(resultPosts[0].Id, true) + require.NoError(t, err) + + assert.ElementsMatch(t, []string{user.Id}, followers) + }) } func TestImportImportPost(t *testing.T) { @@ -3890,6 +4051,109 @@ func TestImportImportDirectPost(t *testing.T) { require.True(t, post.IsPinned) }) + t.Run("Importing a direct post with a thread", func(t *testing.T) { + // Create a thread. + importCreate := time.Now().Add(-1 * time.Minute).UnixMilli() + data := imports.LineImportWorkerData{ + LineImportData: imports.LineImportData{ + DirectPost: &imports.DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Thread Message"), + CreateAt: ptrInt64(importCreate), + Replies: &[]imports.ReplyImportData{{ + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Reply"), + CreateAt: ptrInt64(model.GetMillis()), + }}, + ThreadFollowers: &[]imports.ThreadFollowerImportData{{ + User: ptrStr(th.BasicUser.Username), + LastViewed: ptrInt64(model.GetMillis()), + }, { + User: ptrStr(th.BasicUser2.Username), + LastViewed: ptrInt64(model.GetMillis()), + }}, + }, + }, + LineNumber: 1, + } + + errLine, err := th.App.importMultipleDirectPostLines(th.Context, []imports.LineImportWorkerData{data}, false, true) + require.Nil(t, err) + require.Equal(t, 0, errLine) + + resultPosts, nErr := th.App.Srv().Store().Post().GetPostsCreatedAt(channel.Id, importCreate) + require.NoError(t, nErr) + require.Equal(t, 1, len(resultPosts)) + + followers, nErr := th.App.Srv().Store().Thread().GetThreadFollowers(resultPosts[0].Id, true) + require.NoError(t, nErr) + + assert.ElementsMatch(t, []string{th.BasicUser.Id, th.BasicUser2.Id}, followers) + }) + + t.Run("Importing a direct post with new followers", func(t *testing.T) { + importCreate := time.Now().Add(-5 * time.Minute).UnixMilli() + data := imports.LineImportWorkerData{ + LineImportData: imports.LineImportData{ + DirectPost: &imports.DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Hello"), + CreateAt: ptrInt64(importCreate), + }, + }, + LineNumber: 1, + } + + errLine, err := th.App.importMultipleDirectPostLines(th.Context, []imports.LineImportWorkerData{data}, false, true) + require.Nil(t, err) + require.Equal(t, 0, errLine) + + resultPosts, nErr := th.App.Srv().Store().Post().GetPostsCreatedAt(channel.Id, importCreate) + require.NoError(t, nErr) + require.Equal(t, 1, len(resultPosts)) + + data = imports.LineImportWorkerData{ + LineImportData: imports.LineImportData{ + DirectPost: &imports.DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Hello"), + CreateAt: ptrInt64(importCreate), + Replies: &[]imports.ReplyImportData{{ + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Reply"), + CreateAt: ptrInt64(model.GetMillis()), + }}, + ThreadFollowers: &[]imports.ThreadFollowerImportData{{ + User: ptrStr(th.BasicUser.Username), + LastViewed: ptrInt64(model.GetMillis()), + }}, + }, + }, + LineNumber: 1, + } + + errLine, err = th.App.importMultipleDirectPostLines(th.Context, []imports.LineImportWorkerData{data}, false, true) + require.Nil(t, err) + require.Equal(t, 0, errLine) + + followers, nErr := th.App.Srv().Store().Thread().GetThreadFollowers(resultPosts[0].Id, true) + require.NoError(t, nErr) + + assert.ElementsMatch(t, []string{th.BasicUser.Id}, followers) + }) + // ------------------ Group Channel ------------------------- // Create the GROUP channel. diff --git a/server/channels/app/imports/import_types.go b/server/channels/app/imports/import_types.go index dde9d09324..11e8003e79 100644 --- a/server/channels/app/imports/import_types.go +++ b/server/channels/app/imports/import_types.go @@ -187,6 +187,8 @@ type PostImportData struct { Replies *[]ReplyImportData `json:"replies,omitempty"` Attachments *[]AttachmentImportData `json:"attachments,omitempty"` IsPinned *bool `json:"is_pinned,omitempty"` + + ThreadFollowers *[]ThreadFollowerImportData `json:"thread_followers,omitempty"` } type DirectChannelImportData struct { @@ -213,6 +215,8 @@ type DirectPostImportData struct { Replies *[]ReplyImportData `json:"replies"` Attachments *[]AttachmentImportData `json:"attachments"` IsPinned *bool `json:"is_pinned,omitempty"` + + ThreadFollowers *[]ThreadFollowerImportData `json:"thread_followers,omitempty"` } type SchemeImportData struct { @@ -255,3 +259,11 @@ type ComparablePreference struct { Category string Name string } + +type ThreadFollowerImportData struct { + // User is the username of the follower. It's the general convention + // for import data types to name it as user for the username. + User *string `json:"user"` + LastViewed *int64 `json:"last_viewed,omitempty"` + UnreadMentions *int64 `json:"unread_mentions,omitempty"` +} diff --git a/server/channels/store/opentracinglayer/opentracinglayer.go b/server/channels/store/opentracinglayer/opentracinglayer.go index aeb196dfd7..ab0c5b6c87 100644 --- a/server/channels/store/opentracinglayer/opentracinglayer.go +++ b/server/channels/store/opentracinglayer/opentracinglayer.go @@ -10806,6 +10806,24 @@ func (s *OpenTracingLayerThreadStore) GetThreadForUser(threadMembership *model.T return result, err } +func (s *OpenTracingLayerThreadStore) GetThreadMembershipsForExport(postID string) ([]*model.ThreadMembershipForExport, error) { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.GetThreadMembershipsForExport") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + result, err := s.ThreadStore.GetThreadMembershipsForExport(postID) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return result, err +} + func (s *OpenTracingLayerThreadStore) GetThreadUnreadReplyCount(threadMembership *model.ThreadMembership) (int64, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.GetThreadUnreadReplyCount") @@ -10932,6 +10950,24 @@ func (s *OpenTracingLayerThreadStore) MaintainMembership(userID string, postID s return result, err } +func (s *OpenTracingLayerThreadStore) MaintainMultipleFromImport(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.MaintainMultipleFromImport") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + result, err := s.ThreadStore.MaintainMultipleFromImport(memberships) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return result, err +} + func (s *OpenTracingLayerThreadStore) MarkAllAsRead(userID string, threadIds []string) error { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.MarkAllAsRead") @@ -11040,6 +11076,24 @@ func (s *OpenTracingLayerThreadStore) PermanentDeleteBatchThreadMembershipsForRe return result, resultVar1, err } +func (s *OpenTracingLayerThreadStore) SaveMultipleMemberships(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.SaveMultipleMemberships") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + result, err := s.ThreadStore.SaveMultipleMemberships(memberships) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return result, 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/server/channels/store/retrylayer/retrylayer.go b/server/channels/store/retrylayer/retrylayer.go index 05c276bdfa..2ccedbcc48 100644 --- a/server/channels/store/retrylayer/retrylayer.go +++ b/server/channels/store/retrylayer/retrylayer.go @@ -12365,6 +12365,27 @@ func (s *RetryLayerThreadStore) GetThreadForUser(threadMembership *model.ThreadM } +func (s *RetryLayerThreadStore) GetThreadMembershipsForExport(postID string) ([]*model.ThreadMembershipForExport, error) { + + tries := 0 + for { + result, err := s.ThreadStore.GetThreadMembershipsForExport(postID) + 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) GetThreadUnreadReplyCount(threadMembership *model.ThreadMembership) (int64, error) { tries := 0 @@ -12512,6 +12533,27 @@ func (s *RetryLayerThreadStore) MaintainMembership(userID string, postID string, } +func (s *RetryLayerThreadStore) MaintainMultipleFromImport(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) { + + tries := 0 + for { + result, err := s.ThreadStore.MaintainMultipleFromImport(memberships) + 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) MarkAllAsRead(userID string, threadIds []string) error { tries := 0 @@ -12638,6 +12680,27 @@ func (s *RetryLayerThreadStore) PermanentDeleteBatchThreadMembershipsForRetentio } +func (s *RetryLayerThreadStore) SaveMultipleMemberships(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) { + + tries := 0 + for { + result, err := s.ThreadStore.SaveMultipleMemberships(memberships) + 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) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { tries := 0 diff --git a/server/channels/store/sqlstore/thread_store.go b/server/channels/store/sqlstore/thread_store.go index 5132a96dc4..f579db55e2 100644 --- a/server/channels/store/sqlstore/thread_store.go +++ b/server/channels/store/sqlstore/thread_store.go @@ -505,6 +505,28 @@ func (s *SqlThreadStore) GetThreadFollowers(threadID string, fetchOnlyActive boo return users, nil } +func (s *SqlThreadStore) GetThreadMembershipsForExport(postID string) ([]*model.ThreadMembershipForExport, error) { + members := []*model.ThreadMembershipForExport{} + + fetchConditions := sq.And{ + sq.Eq{"PostId": postID}, + sq.Eq{"Following": true}, + } + + query := s.getQueryBuilder(). + Select("Users.Username, ThreadMemberships.LastViewed, ThreadMemberships.UnreadMentions"). + From("ThreadMemberships"). + InnerJoin("Users ON ThreadMemberships.UserId = Users.Id"). + Where(fetchConditions) + + err := s.GetReplicaX().SelectBuilder(&members, query) + if err != nil { + return nil, errors.Wrapf(err, "failed to get thread members for thread id=%s", postID) + } + + return members, nil +} + func (s *SqlThreadStore) GetThreadForUser(threadMembership *model.ThreadMembership, extended, postPriorityEnabled bool) (*model.ThreadResponse, error) { if !threadMembership.Following { return nil, store.NewErrNotFound("ThreadMembership", "") @@ -807,22 +829,83 @@ func (s *SqlThreadStore) DeleteMembershipForUser(userId string, postId string) e // - post creation (mentions handling) // - channel marked unread // - user explicitly following a thread -func (s *SqlThreadStore) MaintainMembership(userId, postId string, opts store.ThreadMembershipOpts) (_ *model.ThreadMembership, err error) { +func (s *SqlThreadStore) MaintainMembership(userID, postID string, opts store.ThreadMembershipOpts) (_ *model.ThreadMembership, err error) { trx, err := s.GetMasterX().Beginx() if err != nil { return nil, errors.Wrap(err, "begin_transaction") } defer finalizeTransactionX(trx, &err) - membership, err := s.getMembershipForUser(trx, userId, postId) + membership, err := s.maintainMembershipTx(trx, userID, postID, opts) + if err != nil { + return nil, err + } + + if err = trx.Commit(); err != nil { + return nil, errors.Wrap(err, "commit_transaction") + } + + return membership, nil +} + +func (s *SqlThreadStore) MaintainMultipleFromImport(memberships []*model.ThreadMembership) (_ []*model.ThreadMembership, err error) { + trx, err := s.GetMasterX().Beginx() + if err != nil { + return nil, errors.Wrap(err, "begin_transaction") + } + defer finalizeTransactionX(trx, &err) + + for _, member := range memberships { + membership, err2 := s.maintainMembershipTx(trx, member.UserId, member.PostId, store.ThreadMembershipOpts{ + ImportData: &store.ThreadMembershipImportData{ + UnreadMentions: member.UnreadMentions, + LastViewed: member.LastViewed, + }, + }) + if err2 != nil { + return nil, err2 + } + + memberships = append(memberships, membership) + } + + if err = trx.Commit(); err != nil { + return nil, errors.Wrap(err, "commit_transaction") + } + + return memberships, nil +} + +func (s *SqlThreadStore) maintainMembershipTx(trx *sqlxTxWrapper, userID, postID string, opts store.ThreadMembershipOpts) (_ *model.ThreadMembership, err error) { + membership, err := s.getMembershipForUser(trx, userID, postID) now := utils.MillisFromTime(time.Now()) // if membership exists, update it if: // a. user started/stopped following a thread // b. mention count changed // c. user viewed a thread + // d. the membership is imported if err == nil { followingNeedsUpdate := (opts.UpdateFollowing && (membership.Following != opts.Following)) - if followingNeedsUpdate || opts.IncrementMentions || opts.UpdateViewedTimestamp { + if imported := opts.ImportData; imported != nil { + // Only the active followers are getting exported, so we can safely assume + // that the user is following the thread. + if membership.LastUpdated > imported.LastViewed { + // User may have stopped following the thread, + // we need to be smart if we should activate the membership + return membership, nil + } + membership.Following = true + membership.LastUpdated = now + membership.UnreadMentions = imported.UnreadMentions + membership.LastViewed = imported.LastViewed + if _, err = s.updateMembership(trx, membership); err != nil { + return nil, err + } + + if err = s.updateThreadParticipantsForUserTx(trx, postID, userID); err != nil { + return nil, err + } + } else if followingNeedsUpdate || opts.IncrementMentions || opts.UpdateViewedTimestamp { if followingNeedsUpdate { membership.Following = opts.Following } @@ -838,10 +921,6 @@ func (s *SqlThreadStore) MaintainMembership(userId, postId string, opts store.Th } } - if err = trx.Commit(); err != nil { - return nil, errors.Wrap(err, "commit_transaction") - } - return membership, err } @@ -851,16 +930,25 @@ func (s *SqlThreadStore) MaintainMembership(userId, postId string, opts store.Th } membership = &model.ThreadMembership{ - PostId: postId, - UserId: userId, + PostId: postID, + UserId: userID, Following: opts.Following, LastUpdated: now, } - if opts.IncrementMentions { - membership.UnreadMentions = 1 - } - if opts.UpdateViewedTimestamp { - membership.LastViewed = now + if opts.ImportData != nil { + membership.UnreadMentions = opts.ImportData.UnreadMentions + membership.LastViewed = opts.ImportData.LastViewed + membership.Following = true + // If we are importing data, we need to update the thread participants regardless + // of what is given from the options. + opts.UpdateParticipants = true + } else { + if opts.IncrementMentions { + membership.UnreadMentions = 1 + } + if opts.UpdateViewedTimestamp { + membership.LastViewed = now + } } membership, err = s.saveMembership(trx, membership) if err != nil { @@ -868,37 +956,11 @@ func (s *SqlThreadStore) MaintainMembership(userId, postId string, opts store.Th } if opts.UpdateParticipants { - if s.DriverName() == model.DatabaseDriverPostgres { - userIdParam, err2 := jsonArray([]string{userId}).Value() - if err2 != nil { - return nil, err2 - } - if s.IsBinaryParamEnabled() { - userIdParam = AppendBinaryFlag(userIdParam.([]byte)) - } - - if _, err2 := trx.ExecRaw(`UPDATE Threads - SET participants = participants || $1::jsonb - WHERE postid=$2 - AND NOT participants ? $3`, userIdParam, postId, userId); err2 != nil { - return nil, err2 - } - } else { - // CONCAT('$[', JSON_LENGTH(Participants), ']') just generates $[n] - // which is the positional syntax required for appending. - if _, err2 := trx.Exec(`UPDATE Threads - SET Participants = JSON_ARRAY_INSERT(Participants, CONCAT('$[', JSON_LENGTH(Participants), ']'), ?) - WHERE PostId=? - AND NOT JSON_CONTAINS(Participants, ?)`, userId, postId, strconv.Quote(userId)); err2 != nil { - return nil, err2 - } + if err = s.updateThreadParticipantsForUserTx(trx, postID, userID); err != nil { + return nil, err } } - if err = trx.Commit(); err != nil { - return nil, errors.Wrap(err, "commit_transaction") - } - return membership, err } @@ -987,3 +1049,71 @@ func (s *SqlThreadStore) GetThreadUnreadReplyCount(threadMembership *model.Threa return unreadReplies, nil } + +// SaveMultipleMemberships saves multiple NEW thread memberships in a single query and meant to be used only in the import +// process. Unlike MaintainMembership, this method does not update the thread participants (which is handled separately +// in the post creation). +func (s *SqlThreadStore) SaveMultipleMemberships(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) { + if len(memberships) == 0 { + return memberships, nil + } + + query := s.getQueryBuilder(). + Insert("ThreadMemberships"). + Columns("PostId", "UserId", "Following", "LastViewed", "LastUpdated", "UnreadMentions") + + for _, member := range memberships { + if err := member.IsValid(); err != nil { + return memberships, err + } + member.LastUpdated = model.GetMillis() + query = query.Values(member.PostId, member.UserId, member.Following, member.LastViewed, member.LastUpdated, member.UnreadMentions) + } + + tx, err := s.GetMasterX().Beginx() + if err != nil { + return nil, errors.Wrap(err, "begin_transaction") + } + defer finalizeTransactionX(tx, &err) + + _, err = tx.ExecBuilder(query) + if err != nil { + return nil, errors.Wrap(err, "failed to save thread memberships") + } + err = tx.Commit() + if err != nil { + return nil, errors.Wrap(err, "commit_transaction") + } + + return memberships, nil +} + +func (s *SqlThreadStore) updateThreadParticipantsForUserTx(trx *sqlxTxWrapper, postID, userID string) error { + if s.DriverName() == model.DatabaseDriverPostgres { + userIdParam, err := jsonArray([]string{userID}).Value() + if err != nil { + return err + } + if s.IsBinaryParamEnabled() { + userIdParam = AppendBinaryFlag(userIdParam.([]byte)) + } + + if _, err := trx.ExecRaw(`UPDATE Threads + SET participants = participants || $1::jsonb + WHERE postid=$2 + AND NOT participants ? $3`, userIdParam, postID, userID); err != nil { + return err + } + } else { + // CONCAT('$[', JSON_LENGTH(Participants), ']') just generates $[n] + // which is the positional syntax required for appending. + if _, err := trx.Exec(`UPDATE Threads + SET Participants = JSON_ARRAY_INSERT(Participants, CONCAT('$[', JSON_LENGTH(Participants), ']'), ?) + WHERE PostId=? + AND NOT JSON_CONTAINS(Participants, ?)`, userID, postID, strconv.Quote(userID)); err != nil { + return err + } + } + + return nil +} diff --git a/server/channels/store/store.go b/server/channels/store/store.go index 73b70cee4f..d7215b7c50 100644 --- a/server/channels/store/store.go +++ b/server/channels/store/store.go @@ -321,6 +321,7 @@ type ChannelMemberHistoryStore interface { } type ThreadStore interface { GetThreadFollowers(threadID string, fetchOnlyActive bool) ([]string, error) + GetThreadMembershipsForExport(postID string) ([]*model.ThreadMembershipForExport, error) Get(id string) (*model.Thread, error) GetTotalUnreadThreads(userId, teamID string, opts model.GetUserThreadsOpts) (int64, error) @@ -346,6 +347,9 @@ type ThreadStore interface { DeleteOrphanedRows(limit int) (deleted int64, err error) GetThreadUnreadReplyCount(threadMembership *model.ThreadMembership) (int64, error) DeleteMembershipsForChannel(userID, channelID string) error + + SaveMultipleMemberships(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) + MaintainMultipleFromImport(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) } type PostStore interface { @@ -1103,6 +1107,9 @@ type ThreadMembershipOpts struct { // UpdateParticipants indicates whether or not the thread's participants list // should be updated. UpdateParticipants bool + // ImportData contains the data only when the membership is imported. + // and triggers a different workflow. + ImportData *ThreadMembershipImportData } // PostReminderMetadata contains some info needed to send @@ -1121,3 +1128,10 @@ type SidebarCategorySearchOpts struct { ExcludeTeam bool Type model.SidebarCategoryType } + +type ThreadMembershipImportData struct { + // LastViewed is the timestamp to set the LastViewed field to. + LastViewed int64 + // UnreadMentions is the number of unread mentions to set the UnreadMentions field to. + UnreadMentions int64 +} diff --git a/server/channels/store/storetest/mocks/ThreadStore.go b/server/channels/store/storetest/mocks/ThreadStore.go index b8fbd477ac..488b631806 100644 --- a/server/channels/store/storetest/mocks/ThreadStore.go +++ b/server/channels/store/storetest/mocks/ThreadStore.go @@ -259,6 +259,36 @@ func (_m *ThreadStore) GetThreadForUser(threadMembership *model.ThreadMembership return r0, r1 } +// GetThreadMembershipsForExport provides a mock function with given fields: postID +func (_m *ThreadStore) GetThreadMembershipsForExport(postID string) ([]*model.ThreadMembershipForExport, error) { + ret := _m.Called(postID) + + if len(ret) == 0 { + panic("no return value specified for GetThreadMembershipsForExport") + } + + var r0 []*model.ThreadMembershipForExport + var r1 error + if rf, ok := ret.Get(0).(func(string) ([]*model.ThreadMembershipForExport, error)); ok { + return rf(postID) + } + if rf, ok := ret.Get(0).(func(string) []*model.ThreadMembershipForExport); ok { + r0 = rf(postID) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.ThreadMembershipForExport) + } + } + + if rf, ok := ret.Get(1).(func(string) error); ok { + r1 = rf(postID) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // GetThreadUnreadReplyCount provides a mock function with given fields: threadMembership func (_m *ThreadStore) GetThreadUnreadReplyCount(threadMembership *model.ThreadMembership) (int64, error) { ret := _m.Called(threadMembership) @@ -459,6 +489,36 @@ func (_m *ThreadStore) MaintainMembership(userID string, postID string, opts sto return r0, r1 } +// MaintainMultipleFromImport provides a mock function with given fields: memberships +func (_m *ThreadStore) MaintainMultipleFromImport(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) { + ret := _m.Called(memberships) + + if len(ret) == 0 { + panic("no return value specified for MaintainMultipleFromImport") + } + + var r0 []*model.ThreadMembership + var r1 error + if rf, ok := ret.Get(0).(func([]*model.ThreadMembership) ([]*model.ThreadMembership, error)); ok { + return rf(memberships) + } + if rf, ok := ret.Get(0).(func([]*model.ThreadMembership) []*model.ThreadMembership); ok { + r0 = rf(memberships) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.ThreadMembership) + } + } + + if rf, ok := ret.Get(1).(func([]*model.ThreadMembership) error); ok { + r1 = rf(memberships) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // MarkAllAsRead provides a mock function with given fields: userID, threadIds func (_m *ThreadStore) MarkAllAsRead(userID string, threadIds []string) error { ret := _m.Called(userID, threadIds) @@ -601,6 +661,36 @@ func (_m *ThreadStore) PermanentDeleteBatchThreadMembershipsForRetentionPolicies return r0, r1, r2 } +// SaveMultipleMemberships provides a mock function with given fields: memberships +func (_m *ThreadStore) SaveMultipleMemberships(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) { + ret := _m.Called(memberships) + + if len(ret) == 0 { + panic("no return value specified for SaveMultipleMemberships") + } + + var r0 []*model.ThreadMembership + var r1 error + if rf, ok := ret.Get(0).(func([]*model.ThreadMembership) ([]*model.ThreadMembership, error)); ok { + return rf(memberships) + } + if rf, ok := ret.Get(0).(func([]*model.ThreadMembership) []*model.ThreadMembership); ok { + r0 = rf(memberships) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.ThreadMembership) + } + } + + if rf, ok := ret.Get(1).(func([]*model.ThreadMembership) error); ok { + r1 = rf(memberships) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // 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/server/channels/store/storetest/thread_store.go b/server/channels/store/storetest/thread_store.go index 2497a6517c..75ebca667b 100644 --- a/server/channels/store/storetest/thread_store.go +++ b/server/channels/store/storetest/thread_store.go @@ -30,6 +30,8 @@ func TestThreadStore(t *testing.T, rctx request.CTX, ss store.Store, s SqlStore) t.Run("MarkAllAsReadByChannels", func(t *testing.T) { testMarkAllAsReadByChannels(t, rctx, ss) }) t.Run("MarkAllAsReadByTeam", func(t *testing.T) { testMarkAllAsReadByTeam(t, rctx, ss) }) t.Run("DeleteMembershipsForChannel", func(t *testing.T) { testDeleteMembershipsForChannel(t, rctx, ss) }) + t.Run("SaveMultipleMemberships", func(t *testing.T) { testSaveMultipleMemberships(t, ss) }) + t.Run("MaintainMultipleFromImport", func(t *testing.T) { testMaintainMultipleFromImport(t, rctx, ss) }) } func testThreadStorePopulation(t *testing.T, rctx request.CTX, ss store.Store) { @@ -1201,6 +1203,75 @@ func testVarious(t *testing.T, rctx request.CTX, ss store.Store) { }) } }) + + t.Run(("GetThreadMembershipsForExport"), func(t *testing.T) { + t.Run("Get members for thread, ensure usernames", func(t *testing.T) { + members, err := ss.Thread().GetThreadMembershipsForExport(team1channel1post1.Id) + require.NoError(t, err) + + // team1channel1post1 has 1 member + assert.Len(t, members, 1) + + userIDs, err := ss.Thread().GetThreadFollowers(team1channel1post1.Id, true) + require.NoError(t, err) + require.Len(t, userIDs, 1) + + u, err := ss.User().Get(context.Background(), userIDs[0]) + require.NoError(t, err) + + assert.Equal(t, u.Username, members[0].Username) + + members, err = ss.Thread().GetThreadMembershipsForExport(team1channel1post2.Id) + require.NoError(t, err) + + // team1channel1post2 has 2 members + assert.Len(t, members, 2) + + userIDs, err = ss.Thread().GetThreadFollowers(team1channel1post2.Id, true) + require.NoError(t, err) + require.Len(t, userIDs, 2) + + for i := range userIDs { + u, err := ss.User().Get(context.Background(), userIDs[i]) + require.NoError(t, err) + + assert.Equal(t, u.Username, members[i].Username) + } + }) + + t.Run("Get members for a thread, ensure only following members are exported", func(t *testing.T) { + createThreadMembership(user2ID, team1channel1post1.Id, false) + + members, err := ss.Thread().GetThreadMembershipsForExport(team1channel1post1.Id) + require.NoError(t, err) + + // team1channel1post1 should have 2 members + assert.Len(t, members, 2) + + _, err = ss.Thread().MaintainMembership(user2ID, team1channel1post1.Id, store.ThreadMembershipOpts{ + Following: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: true, + }) + require.NoError(t, err) + + members, err = ss.Thread().GetThreadMembershipsForExport(team1channel1post1.Id) + require.NoError(t, err) + + // team1channel1post1 should have 1 following member + assert.Len(t, members, 1) + + userIDs, err := ss.Thread().GetThreadFollowers(team1channel1post1.Id, true) + require.NoError(t, err) + require.Len(t, userIDs, 1) + + u, err := ss.User().Get(context.Background(), userIDs[0]) + require.NoError(t, err) + + assert.Equal(t, u.Username, members[0].Username) + }) + }) } func testMarkAllAsReadByChannels(t *testing.T, rctx request.CTX, ss store.Store) { @@ -1690,3 +1761,241 @@ func testDeleteMembershipsForChannel(t *testing.T, rctx request.CTX, ss store.St require.ElementsMatch(t, []*model.ThreadMembership{memB1}, membershipsB) }) } + +func testSaveMultipleMemberships(t *testing.T, ss store.Store) { + t.Run("should save multiple memberships", func(t *testing.T) { + memberships := []*model.ThreadMembership{ + { + PostId: model.NewId(), + UserId: model.NewId(), + Following: true, + }, + { + PostId: model.NewId(), + UserId: model.NewId(), + Following: true, + }, + } + + _, err := ss.Thread().SaveMultipleMemberships(memberships) + require.NoError(t, err) + }) + + t.Run("should return error if any of the memberships is invalid", func(t *testing.T) { + memberships := []*model.ThreadMembership{ + { + PostId: model.NewId(), + UserId: "invalid", + Following: true, + }, + { + PostId: model.NewId(), + UserId: model.NewId(), + Following: true, + }, + } + + _, err := ss.Thread().SaveMultipleMemberships(memberships) + require.Error(t, err) + }) + + t.Run("should not fail if the list is empty", func(t *testing.T) { + _, err := ss.Thread().SaveMultipleMemberships([]*model.ThreadMembership{}) + require.NoError(t, err) + }) + + t.Run("should fail if there is a conflict", func(t *testing.T) { + postID := model.NewId() + userID := model.NewId() + + memberships := []*model.ThreadMembership{ + { + PostId: postID, + UserId: userID, + Following: true, + }, + { + PostId: postID, + UserId: userID, + Following: true, + }, + } + + _, err := ss.Thread().SaveMultipleMemberships(memberships) + require.Error(t, err) + }) +} + +func testMaintainMultipleFromImport(t *testing.T, rctx request.CTX, ss store.Store) { + createThreadMembership := func(userID, postID string, following bool) (*model.ThreadMembership, func()) { + t.Helper() + opts := store.ThreadMembershipOpts{ + Following: following, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + mem, err := ss.Thread().MaintainMembership(userID, postID, opts) + require.NoError(t, err) + + return mem, func() { + err := ss.Thread().DeleteMembershipForUser(userID, postID) + require.NoError(t, err) + } + } + + cleanMembers := func(userIDs []string, postID string) error { + // clean the thread memberships + for _, id := range userIDs { + err := ss.Thread().DeleteMembershipForUser(id, postID) + if err != nil { + return err + } + } + return nil + } + + postingUserID := model.NewId() + + team, err := ss.Team().Save(&model.Team{ + DisplayName: "DisplayName", + Name: "team" + model.NewId(), + Email: MakeEmail(), + Type: model.TeamOpen, + }) + require.NoError(t, err) + + channel1, err := ss.Channel().Save(rctx, &model.Channel{ + TeamId: team.Id, + DisplayName: "DisplayName", + Name: "channel1" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + + rootPost1, err := ss.Post().Save(rctx, &model.Post{ + ChannelId: channel1.Id, + UserId: postingUserID, + Message: model.NewRandomString(10), + }) + require.NoError(t, err) + + _, err = ss.Post().Save(rctx, &model.Post{ + ChannelId: channel1.Id, + UserId: postingUserID, + Message: model.NewRandomString(10), + RootId: rootPost1.Id, + }) + require.NoError(t, err) + + t.Run("Should create new memberships from new list", func(t *testing.T) { + userAID := model.NewId() + userBID := model.NewId() + + _, err := ss.Thread().MaintainMultipleFromImport([]*model.ThreadMembership{ + { + UserId: userAID, + PostId: rootPost1.Id, + Following: true, + }, + { + UserId: userBID, + PostId: rootPost1.Id, + Following: true, + }, + }) + require.NoError(t, err) + + followers, err := ss.Thread().GetThreadFollowers(rootPost1.Id, true) + require.NoError(t, err) + require.ElementsMatch(t, followers, []string{userAID, userBID}) + + // clean the thread memberships + err = cleanMembers(followers, rootPost1.Id) + require.NoError(t, err) + }) + + t.Run("Should add incoming memberships from the list", func(t *testing.T) { + userAID := model.NewId() + userBID := model.NewId() + + _, clean := createThreadMembership(userAID, rootPost1.Id, true) + defer clean() + + _, err := ss.Thread().MaintainMultipleFromImport([]*model.ThreadMembership{ + { + UserId: userBID, + PostId: rootPost1.Id, + Following: true, + }, + }) + require.NoError(t, err) + + followers, err := ss.Thread().GetThreadFollowers(rootPost1.Id, true) + require.NoError(t, err) + require.ElementsMatch(t, followers, []string{userAID, userBID}) + + // clean the thread memberships + err = cleanMembers(followers, rootPost1.Id) + require.NoError(t, err) + }) + + t.Run("Should update memberships if they are newer", func(t *testing.T) { + userAID := model.NewId() + + old, clean := createThreadMembership(userAID, rootPost1.Id, true) + defer clean() + + _, err := ss.Thread().MaintainMultipleFromImport([]*model.ThreadMembership{ + { + UserId: userAID, + PostId: rootPost1.Id, + Following: true, + LastViewed: time.Now().Add(time.Minute).UnixMilli(), + }, + }) + require.NoError(t, err) + + followers, err := ss.Thread().GetThreadFollowers(rootPost1.Id, true) + require.NoError(t, err) + require.ElementsMatch(t, followers, []string{userAID}) + + updated, err := ss.Thread().GetMembershipForUser(userAID, rootPost1.Id) + require.NoError(t, err) + require.Greater(t, updated.LastViewed, old.LastViewed) + + // clean the thread memberships + err = cleanMembers(followers, rootPost1.Id) + require.NoError(t, err) + }) + + t.Run("Should not update membership if incoming is not newer", func(t *testing.T) { + userAID := model.NewId() + + _, clean := createThreadMembership(userAID, rootPost1.Id, false) + defer clean() + + _, err := ss.Thread().MaintainMultipleFromImport([]*model.ThreadMembership{ + { + UserId: userAID, + PostId: rootPost1.Id, + Following: true, + LastViewed: time.Now().Add(-1 * time.Hour).UnixMilli(), + }, + }) + require.NoError(t, err) + + followers, err := ss.Thread().GetThreadFollowers(rootPost1.Id, true) + require.NoError(t, err) + require.Empty(t, followers) + + m, err := ss.Thread().GetMembershipForUser(userAID, rootPost1.Id) + require.NoError(t, err) + require.False(t, m.Following) + + // clean the thread memberships + err = cleanMembers(followers, rootPost1.Id) + require.NoError(t, err) + }) +} diff --git a/server/channels/store/timerlayer/timerlayer.go b/server/channels/store/timerlayer/timerlayer.go index c9d8bc7a7d..3ec99d88f7 100644 --- a/server/channels/store/timerlayer/timerlayer.go +++ b/server/channels/store/timerlayer/timerlayer.go @@ -9719,6 +9719,22 @@ func (s *TimerLayerThreadStore) GetThreadForUser(threadMembership *model.ThreadM return result, err } +func (s *TimerLayerThreadStore) GetThreadMembershipsForExport(postID string) ([]*model.ThreadMembershipForExport, error) { + start := time.Now() + + result, err := s.ThreadStore.GetThreadMembershipsForExport(postID) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.GetThreadMembershipsForExport", success, elapsed) + } + return result, err +} + func (s *TimerLayerThreadStore) GetThreadUnreadReplyCount(threadMembership *model.ThreadMembership) (int64, error) { start := time.Now() @@ -9831,6 +9847,22 @@ func (s *TimerLayerThreadStore) MaintainMembership(userID string, postID string, return result, err } +func (s *TimerLayerThreadStore) MaintainMultipleFromImport(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) { + start := time.Now() + + result, err := s.ThreadStore.MaintainMultipleFromImport(memberships) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.MaintainMultipleFromImport", success, elapsed) + } + return result, err +} + func (s *TimerLayerThreadStore) MarkAllAsRead(userID string, threadIds []string) error { start := time.Now() @@ -9927,6 +9959,22 @@ func (s *TimerLayerThreadStore) PermanentDeleteBatchThreadMembershipsForRetentio return result, resultVar1, err } +func (s *TimerLayerThreadStore) SaveMultipleMemberships(memberships []*model.ThreadMembership) ([]*model.ThreadMembership, error) { + start := time.Now() + + result, err := s.ThreadStore.SaveMultipleMemberships(memberships) + + elapsed := float64(time.Since(start)) / float64(time.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.SaveMultipleMemberships", success, elapsed) + } + return result, err +} + func (s *TimerLayerThreadStore) UpdateMembership(membership *model.ThreadMembership) (*model.ThreadMembership, error) { start := time.Now() diff --git a/server/i18n/en.json b/server/i18n/en.json index 26227fdf89..3f385a63d8 100644 --- a/server/i18n/en.json +++ b/server/i18n/en.json @@ -6222,6 +6222,10 @@ "id": "app.post.save.existing.app_error", "translation": "You cannot update an existing Post." }, + { + "id": "app.post.save.thread_membership.app_error", + "translation": "Unable to save thread membership for post." + }, { "id": "app.post.search.app_error", "translation": "Error searching posts" @@ -6718,6 +6722,10 @@ "id": "app.terms_of_service.get.no_rows.app_error", "translation": "No terms of service found." }, + { + "id": "app.thread.get_threadmembers_for_export.app_error", + "translation": "Unable to get thread members for export." + }, { "id": "app.thread.mark_all_as_read_by_channels.app_error", "translation": "Unable to mark all threads as read by channel" @@ -9658,6 +9666,14 @@ "id": "model.team_member.is_valid.user_id.app_error", "translation": "Invalid user id." }, + { + "id": "model.thread.is_valid.post_id.app_error", + "translation": "Invalid post ID." + }, + { + "id": "model.thread.is_valid.user_id.app_error", + "translation": "Invalid user ID." + }, { "id": "model.token.is_valid.expiry", "translation": "Invalid token expiry" diff --git a/server/public/model/thread.go b/server/public/model/thread.go index ce8ebcca3c..b510d1196a 100644 --- a/server/public/model/thread.go +++ b/server/public/model/thread.go @@ -3,6 +3,8 @@ package model +import "net/http" + // Thread tracks the metadata associated with a root post and its reply posts. // // Note that Thread metadata does not exist until the first reply to a root post. @@ -126,3 +128,21 @@ type ThreadMembership struct { // threads with the mention count. UnreadMentions int64 `json:"unread_mentions"` } + +func (o *ThreadMembership) IsValid() *AppError { + if !IsValidId(o.PostId) { + return NewAppError("ThreadMembership.IsValid", "model.thread.is_valid.post_id.app_error", nil, "", http.StatusBadRequest) + } + + if !IsValidId(o.UserId) { + return NewAppError("ThreadMembership.IsValid", "model.thread.is_valid.user_id.app_error", nil, "", http.StatusBadRequest) + } + + return nil +} + +type ThreadMembershipForExport struct { + Username string `json:"user_name"` + LastViewed int64 `json:"last_viewed"` + UnreadMentions int64 `json:"unread_mentions"` +} diff --git a/server/public/model/thread_test.go b/server/public/model/thread_test.go new file mode 100644 index 0000000000..42c139cd2f --- /dev/null +++ b/server/public/model/thread_test.go @@ -0,0 +1,39 @@ +// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved. +// See LICENSE.txt for license information. + +package model + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestThreadMembershipIsValid(t *testing.T) { + cases := map[string]struct { + Member *ThreadMembership + ShouldBeValid bool + }{ + "valid member": { + Member: &ThreadMembership{PostId: NewId(), UserId: NewId()}, + ShouldBeValid: true, + }, + "empty post id": { + Member: &ThreadMembership{PostId: "", UserId: NewId()}, + ShouldBeValid: false, + }, + "empty user id": { + Member: &ThreadMembership{PostId: NewId(), UserId: ""}, + ShouldBeValid: false, + }, + "invalid post id": { + Member: &ThreadMembership{PostId: "invalid", UserId: NewId()}, + ShouldBeValid: false, + }, + } + for name, tc := range cases { + t.Run(name, func(t *testing.T) { + require.Equal(t, tc.ShouldBeValid, (tc.Member.IsValid() == nil)) + }) + } +}