From 27d536b212c0b52f4b553c96df99af7667e66f6a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jes=C3=BAs=20Espino?= Date: Wed, 11 Mar 2020 14:29:32 +0100 Subject: [PATCH] MM-21552: Adding SaveMultiple to posts (#13766) * Adding SaveMultiple to posts * Improving tests * fixing i18n * Fixing tests * Improving testing on top of Save and SaveMultiple * Fixing shadow variables * Addressing some PR comments * More clear update post test * Addressing some PR comments * Addressing some PR comments and simplifying the code * Improting replies in bulk too * Fixing reply count and processing last imported replies * Adding OverwriteMultiple to posts aggregating everything in the same transaction * Adding 2 pending tests to implement * Adding tests for overwrite multiple posts * Adding tests for TeamStore.GetByNames method * Fixing shadow variables * Addressing PR comments * Extracting i18n strings * Fixing tests * Fixing tests * Adding more test cases * Using a variable instead of a fake timestamp --- app/import.go | 49 +- app/import_functions.go | 658 +++++++----- app/import_functions_test.go | 1352 ++++++++++++++++++------- app/import_test.go | 4 +- cmd/mattermost/commands/sampledata.go | 16 + i18n/en.json | 36 +- store/sqlstore/post_store.go | 182 +++- store/sqlstore/team_store.go | 28 + store/store.go | 3 + store/storetest/mocks/PostStore.go | 50 + store/storetest/mocks/TeamStore.go | 25 + store/storetest/post_store.go | 482 +++++++-- store/storetest/team_store.go | 74 +- store/timer_layer.go | 16 + 14 files changed, 2208 insertions(+), 767 deletions(-) diff --git a/app/import.go b/app/import.go index 39e379252f..155d942bdd 100644 --- a/app/import.go +++ b/app/import.go @@ -16,7 +16,10 @@ import ( "github.com/mattermost/mattermost-server/v5/model" ) -const maxScanTokenSize = 16 * 1024 * 1024 // Need to set a higher limit than default because some customers cross the limit. See MM-22314 +const ( + importMultiplePostsThreshold = 1000 + maxScanTokenSize = 16 * 1024 * 1024 // Need to set a higher limit than default because some customers cross the limit. See MM-22314 +) func stopOnError(err LineImportWorkerError) bool { if err.Error.Id == "api.file.upload_file.large_image.app_error" { @@ -27,11 +30,41 @@ func stopOnError(err LineImportWorkerError) bool { } func (a *App) bulkImportWorker(dryRun bool, wg *sync.WaitGroup, lines <-chan LineImportWorkerData, errors chan<- LineImportWorkerError) { + posts := []*PostImportData{} + directPosts := []*DirectPostImportData{} for line := range lines { - if err := a.importLine(line.LineImportData, dryRun); err != nil { - errors <- LineImportWorkerError{err, line.LineNumber} + switch { + case line.LineImportData.Type == "post": + posts = append(posts, line.Post) + if line.Post == nil { + errors <- LineImportWorkerError{model.NewAppError("BulkImport", "app.import.import_line.null_post.error", nil, "", http.StatusBadRequest), line.LineNumber} + } + if len(posts) >= importMultiplePostsThreshold { + a.importMultiplePosts(posts, dryRun) + posts = []*PostImportData{} + } + case line.LineImportData.Type == "direct_post": + directPosts = append(directPosts, line.DirectPost) + if line.DirectPost == nil { + errors <- LineImportWorkerError{model.NewAppError("BulkImport", "app.import.import_line.null_direct_post.error", nil, "", http.StatusBadRequest), line.LineNumber} + } + if len(directPosts) >= importMultiplePostsThreshold { + a.importMultipleDirectPosts(directPosts, dryRun) + directPosts = []*DirectPostImportData{} + } + default: + if err := a.importLine(line.LineImportData, dryRun); err != nil { + errors <- LineImportWorkerError{err, line.LineNumber} + } } } + + if len(posts) > 0 { + a.importMultiplePosts(posts, dryRun) + } + if len(directPosts) > 0 { + a.importMultipleDirectPosts(directPosts, dryRun) + } wg.Done() } @@ -159,21 +192,11 @@ func (a *App) importLine(line LineImportData, dryRun bool) *model.AppError { return model.NewAppError("BulkImport", "app.import.import_line.null_user.error", nil, "", http.StatusBadRequest) } return a.importUser(line.User, dryRun) - case line.Type == "post": - if line.Post == nil { - return model.NewAppError("BulkImport", "app.import.import_line.null_post.error", nil, "", http.StatusBadRequest) - } - return a.importPost(line.Post, dryRun) case line.Type == "direct_channel": if line.DirectChannel == nil { return model.NewAppError("BulkImport", "app.import.import_line.null_direct_channel.error", nil, "", http.StatusBadRequest) } return a.importDirectChannel(line.DirectChannel, dryRun) - case line.Type == "direct_post": - if line.DirectPost == nil { - return model.NewAppError("BulkImport", "app.import.import_line.null_direct_post.error", nil, "", http.StatusBadRequest) - } - return a.importDirectPost(line.DirectPost, dryRun) case line.Type == "emoji": if line.Emoji == nil { return model.NewAppError("BulkImport", "app.import.import_line.null_emoji.error", nil, "", http.StatusBadRequest) diff --git a/app/import_functions.go b/app/import_functions.go index 7bf26f4a70..e67185c7aa 100644 --- a/app/import_functions.go +++ b/app/import_functions.go @@ -851,67 +851,87 @@ func (a *App) importReaction(data *ReactionImportData, post *model.Post, dryRun return nil } -func (a *App) importReply(data *ReplyImportData, post *model.Post, teamId string, dryRun bool) *model.AppError { +func (a *App) importReplies(data []ReplyImportData, post *model.Post, teamId string, dryRun bool) *model.AppError { var err *model.AppError - if err = validateReplyImportData(data, post.CreateAt, a.MaxPostSize()); err != nil { - return err - } - - var user *model.User - user, err = a.Srv().Store.User().GetByUsername(*data.User) - if err != nil { - return model.NewAppError("BulkImport", "app.import.import_post.user_not_found.error", map[string]interface{}{"Username": data.User}, err.Error(), http.StatusBadRequest) - } - - // Check if this post already exists. - replies, err := a.Srv().Store.Post().GetPostsCreatedAt(post.ChannelId, *data.CreateAt) - if err != nil { - return err - } - - var reply *model.Post - for _, r := range replies { - if r.Message == *data.Message && r.RootId == post.Id { - reply = r - break - } - } - - if reply == nil { - reply = &model.Post{} - } - reply.UserId = user.Id - reply.ChannelId = post.ChannelId - reply.ParentId = post.Id - reply.RootId = post.Id - reply.Message = *data.Message - reply.CreateAt = *data.CreateAt - - fileIds, err := a.uploadAttachments(data.Attachments, reply, teamId, dryRun) - if err != nil { - return err - } - for _, fileID := range reply.FileIds { - if _, ok := fileIds[fileID]; !ok { - a.Srv().Store.FileInfo().PermanentDelete(fileID) - } - } - reply.FileIds = make([]string, 0) - for fileID := range fileIds { - reply.FileIds = append(reply.FileIds, fileID) - } - - if reply.Id == "" { - if _, err := a.Srv().Store.Post().Save(reply); err != nil { + usernames := []string{} + for _, replyData := range data { + if err = validateReplyImportData(&replyData, post.CreateAt, a.MaxPostSize()); err != nil { return err } - } else { - if _, err := a.Srv().Store.Post().Overwrite(reply); err != nil { + usernames = append(usernames, *replyData.User) + } + + users, err := a.getUsersByUsernames(usernames) + if err != nil { + return err + } + + postsWithData := []postAndData{} + postsForCreateList := []*model.Post{} + postsForOverwriteList := []*model.Post{} + + for _, replyData := range data { + user := users[*replyData.User] + + // Check if this post already exists. + replies, err := a.Srv().Store.Post().GetPostsCreatedAt(post.ChannelId, *replyData.CreateAt) + if err != nil { + return err + } + + var reply *model.Post + for _, r := range replies { + if r.Message == *replyData.Message && r.RootId == post.Id { + reply = r + break + } + } + + if reply == nil { + reply = &model.Post{} + } + reply.UserId = user.Id + reply.ChannelId = post.ChannelId + reply.ParentId = post.Id + reply.RootId = post.Id + reply.Message = *replyData.Message + reply.CreateAt = *replyData.CreateAt + + fileIds, err := a.uploadAttachments(replyData.Attachments, reply, teamId, dryRun) + if err != nil { + return err + } + for _, fileID := range reply.FileIds { + if _, ok := fileIds[fileID]; !ok { + a.Srv().Store.FileInfo().PermanentDelete(fileID) + } + } + reply.FileIds = make([]string, 0) + for fileID := range fileIds { + reply.FileIds = append(reply.FileIds, fileID) + } + + if len(reply.Id) == 0 { + postsForCreateList = append(postsForCreateList, reply) + } else { + postsForOverwriteList = append(postsForOverwriteList, reply) + } + postsWithData = append(postsWithData, postAndData{post: reply, replyData: &replyData}) + } + + if len(postsForCreateList) > 0 { + if _, err := a.Srv().Store.Post().SaveMultiple(postsForCreateList); err != nil { return err } } - a.updateFileInfoWithPostId(reply) + if _, err := a.Srv().Store.Post().OverwriteMultiple(postsForOverwriteList); err != nil { + return err + } + + for _, postWithData := range postsWithData { + a.updateFileInfoWithPostId(postWithData.post) + } return nil } @@ -961,9 +981,70 @@ func (a *App) importAttachment(data *AttachmentImportData, post *model.Post, tea return fileInfo, nil } -func (a *App) importPost(data *PostImportData, dryRun bool) *model.AppError { - if err := validatePostImportData(data, a.MaxPostSize()); err != nil { - return err +type postAndData struct { + post *model.Post + postData *PostImportData + directPostData *DirectPostImportData + replyData *ReplyImportData + team *model.Team +} + +func (a *App) getUsersByUsernames(usernames []string) (map[string]*model.User, *model.AppError) { + uniqueUsernames := utils.RemoveDuplicatesFromStringArray(usernames) + allUsers, err := a.Srv().Store.User().GetProfilesByUsernames(uniqueUsernames, nil) + if err != nil { + return nil, model.NewAppError("BulkImport", "app.import.get_users_by_username.some_users_not_found.error", nil, err.Error(), http.StatusBadRequest) + } + + if len(allUsers) != len(uniqueUsernames) { + return nil, model.NewAppError("BulkImport", "app.import.get_users_by_username.some_users_not_found.error", nil, "", http.StatusBadRequest) + } + + users := make(map[string]*model.User) + for _, user := range allUsers { + users[user.Username] = user + } + return users, nil +} + +func (a *App) getTeamsByNames(names []string) (map[string]*model.Team, *model.AppError) { + allTeams, err := a.Srv().Store.Team().GetByNames(names) + if err != nil { + return nil, model.NewAppError("BulkImport", "app.import.get_teams_by_names.some_teams_not_found.error", nil, err.Error(), http.StatusBadRequest) + } + + teams := make(map[string]*model.Team) + for _, team := range allTeams { + teams[team.Name] = team + } + return teams, nil +} + +func (a *App) getChannelsForPosts(teams map[string]*model.Team, data []*PostImportData) (map[string]*model.Channel, *model.AppError) { + channels := make(map[string]*model.Channel) + for _, postData := range data { + team := teams[*postData.Team] + if channel, ok := channels[*postData.Channel]; !ok || channel == nil { + var err *model.AppError + channel, err = a.Srv().Store.Channel().GetByName(team.Id, *postData.Channel, true) + if err != nil { + return nil, model.NewAppError("BulkImport", "app.import.import_post.channel_not_found.error", map[string]interface{}{"ChannelName": *postData.Channel}, err.Error(), http.StatusBadRequest) + } + channels[*postData.Channel] = channel + } + } + return channels, nil +} + +func (a *App) importMultiplePosts(data []*PostImportData, dryRun bool) *model.AppError { + if len(data) == 0 { + return nil + } + + for _, postData := range data { + if err := validatePostImportData(postData, a.MaxPostSize()); err != nil { + return err + } } // If this is a Dry Run, do not continue any further. @@ -971,113 +1052,148 @@ func (a *App) importPost(data *PostImportData, dryRun bool) *model.AppError { return nil } - team, err := a.Srv().Store.Team().GetByName(*data.Team) - if err != nil { - return model.NewAppError("BulkImport", "app.import.import_post.team_not_found.error", map[string]interface{}{"TeamName": *data.Team}, err.Error(), http.StatusBadRequest) + usernames := []string{} + teamNames := []string{} + for _, postData := range data { + usernames = append(usernames, *postData.User) + if postData.FlaggedBy != nil { + usernames = append(usernames, *postData.FlaggedBy...) + } + teamNames = append(teamNames, *postData.Team) } - channel, err := a.Srv().Store.Channel().GetByName(team.Id, *data.Channel, false) - if err != nil { - return model.NewAppError("BulkImport", "app.import.import_post.channel_not_found.error", map[string]interface{}{"ChannelName": *data.Channel}, err.Error(), http.StatusBadRequest) - } - - var user *model.User - user, err = a.Srv().Store.User().GetByUsername(*data.User) - if err != nil { - return model.NewAppError("BulkImport", "app.import.import_post.user_not_found.error", map[string]interface{}{"Username": *data.User}, err.Error(), http.StatusBadRequest) - } - - // Check if this post already exists. - posts, err := a.Srv().Store.Post().GetPostsCreatedAt(channel.Id, *data.CreateAt) + users, err := a.getUsersByUsernames(usernames) if err != nil { return err } - var post *model.Post - for _, p := range posts { - if p.Message == *data.Message { - post = p - break - } - } - - if post == nil { - post = &model.Post{} - } - - post.ChannelId = channel.Id - post.Message = *data.Message - post.UserId = user.Id - post.CreateAt = *data.CreateAt - - post.Hashtags, _ = model.ParseHashtags(post.Message) - - fileIds, err := a.uploadAttachments(data.Attachments, post, team.Id, dryRun) + teams, err := a.getTeamsByNames(teamNames) if err != nil { return err } - for _, fileID := range post.FileIds { - if _, ok := fileIds[fileID]; !ok { - a.Srv().Store.FileInfo().PermanentDelete(fileID) - } - } - post.FileIds = make([]string, 0) - for fileID := range fileIds { - post.FileIds = append(post.FileIds, fileID) - } - if post.Id == "" { - if _, err = a.Srv().Store.Post().Save(post); err != nil { + channels, err := a.getChannelsForPosts(teams, data) + if err != nil { + return err + } + postsWithData := []postAndData{} + postsForCreateList := []*model.Post{} + postsForOverwriteList := []*model.Post{} + + for _, postData := range data { + team := teams[*postData.Team] + channel := channels[*postData.Channel] + user := users[*postData.User] + + // Check if this post already exists. + posts, err := a.Srv().Store.Post().GetPostsCreatedAt(channel.Id, *postData.CreateAt) + if err != nil { return err } - } else { - if _, err = a.Srv().Store.Post().Overwrite(post); err != nil { + + var post *model.Post + for _, p := range posts { + if p.Message == *postData.Message { + post = p + break + } + } + + if post == nil { + post = &model.Post{} + } + + post.ChannelId = channel.Id + post.Message = *postData.Message + post.UserId = user.Id + post.CreateAt = *postData.CreateAt + post.Hashtags, _ = model.ParseHashtags(post.Message) + + fileIds, err := a.uploadAttachments(postData.Attachments, post, team.Id, dryRun) + if err != nil { + return err + } + for _, fileID := range post.FileIds { + if _, ok := fileIds[fileID]; !ok { + a.Srv().Store.FileInfo().PermanentDelete(fileID) + } + } + post.FileIds = make([]string, 0) + for fileID := range fileIds { + post.FileIds = append(post.FileIds, fileID) + } + + if len(post.Id) == 0 { + postsForCreateList = append(postsForCreateList, post) + } else { + postsForOverwriteList = append(postsForOverwriteList, post) + } + postsWithData = append(postsWithData, postAndData{post: post, postData: postData, team: team}) + } + + if len(postsForCreateList) > 0 { + if _, err := a.Srv().Store.Post().SaveMultiple(postsForCreateList); err != nil { return err } } - if data.FlaggedBy != nil { - var preferences model.Preferences + if _, err := a.Srv().Store.Post().OverwriteMultiple(postsForOverwriteList); err != nil { + return err + } - for _, username := range *data.FlaggedBy { - var user *model.User - user, err = a.Srv().Store.User().GetByUsername(username) - if err != nil { - return model.NewAppError("BulkImport", "app.import.import_post.user_not_found.error", map[string]interface{}{"Username": username}, err.Error(), http.StatusBadRequest) + var lastPostWithData *postAndData + repliesBulk := []ReplyImportData{} + for _, postWithData := range postsWithData { + if postWithData.postData.FlaggedBy != nil { + var preferences model.Preferences + + for _, username := range *postWithData.postData.FlaggedBy { + user := users[username] + + preferences = append(preferences, model.Preference{ + UserId: user.Id, + Category: model.PREFERENCE_CATEGORY_FLAGGED_POST, + Name: postWithData.post.Id, + Value: "true", + }) } - preferences = append(preferences, model.Preference{ - UserId: user.Id, - Category: model.PREFERENCE_CATEGORY_FLAGGED_POST, - Name: post.Id, - Value: "true", - }) + if len(preferences) > 0 { + if err := a.Srv().Store.Preference().Save(&preferences); err != nil { + return model.NewAppError("BulkImport", "app.import.import_post.save_preferences.error", nil, err.Error(), http.StatusInternalServerError) + } + } } - if len(preferences) > 0 { - if err := a.Srv().Store.Preference().Save(&preferences); err != nil { - return model.NewAppError("BulkImport", "app.import.import_post.save_preferences.error", nil, err.Error(), http.StatusInternalServerError) + if postWithData.postData.Reactions != nil { + for _, reaction := range *postWithData.postData.Reactions { + if err := a.importReaction(&reaction, postWithData.post, dryRun); err != nil { + return err + } } } + + if postWithData.postData.Replies != nil { + repliesBulk = append(repliesBulk, *postWithData.postData.Replies...) + if len(repliesBulk) >= importMultiplePostsThreshold { + err := a.importReplies(repliesBulk, postWithData.post, postWithData.team.Id, dryRun) + if err != nil { + return err + } + repliesBulk = []ReplyImportData{} + } + } + a.updateFileInfoWithPostId(postWithData.post) + lastPostWithData = &postWithData + } + + if len(repliesBulk) >= 0 && lastPostWithData != nil { + err := a.importReplies(repliesBulk, lastPostWithData.post, lastPostWithData.team.Id, dryRun) + if err != nil { + return err + } } - if data.Reactions != nil { - for _, reaction := range *data.Reactions { - if err := a.importReaction(&reaction, post, dryRun); err != nil { - return err - } - } - } - - if data.Replies != nil { - for _, reply := range *data.Replies { - if err := a.importReply(&reply, post, team.Id, dryRun); err != nil { - return err - } - } - } - - a.updateFileInfoWithPostId(post) return nil } @@ -1116,15 +1232,12 @@ func (a *App) importDirectChannel(data *DirectChannelImportData, dryRun bool) *m } var userIds []string - userMap := make(map[string]string) - for _, username := range *data.Members { - var user *model.User - user, err = a.Srv().Store.User().GetByUsername(username) - if err != nil { - return model.NewAppError("BulkImport", "app.import.import_direct_channel.member_not_found.error", nil, err.Error(), http.StatusBadRequest) - } - userIds = append(userIds, user.Id) - userMap[username] = user.Id + userMap, err := a.getUsersByUsernames(*data.Members) + if err != nil { + return err + } + for _, user := range *data.Members { + userIds = append(userIds, userMap[user].Id) } var channel *model.Channel @@ -1157,7 +1270,7 @@ func (a *App) importDirectChannel(data *DirectChannelImportData, dryRun bool) *m if data.FavoritedBy != nil { for _, favoriter := range *data.FavoritedBy { preferences = append(preferences, model.Preference{ - UserId: userMap[favoriter], + UserId: userMap[favoriter].Id, Category: model.PREFERENCE_CATEGORY_FAVORITE_CHANNEL, Name: channel.Id, Value: "true", @@ -1180,10 +1293,15 @@ func (a *App) importDirectChannel(data *DirectChannelImportData, dryRun bool) *m return nil } -func (a *App) importDirectPost(data *DirectPostImportData, dryRun bool) *model.AppError { - var err *model.AppError - if err = validateDirectPostImportData(data, a.MaxPostSize()); err != nil { - return err +func (a *App) importMultipleDirectPosts(data []*DirectPostImportData, dryRun bool) *model.AppError { + if len(data) == 0 { + return nil + } + + for _, postData := range data { + if err := validateDirectPostImportData(postData, a.MaxPostSize()); err != nil { + return err + } } // If this is a Dry Run, do not continue any further. @@ -1191,129 +1309,143 @@ func (a *App) importDirectPost(data *DirectPostImportData, dryRun bool) *model.A return nil } - var userIds []string - for _, username := range *data.ChannelMembers { - var user *model.User - user, err = a.Srv().Store.User().GetByUsername(username) + usernames := []string{} + for _, postData := range data { + usernames = append(usernames, *postData.User) + if postData.FlaggedBy != nil { + usernames = append(usernames, *postData.FlaggedBy...) + } + usernames = append(usernames, *postData.ChannelMembers...) + } + + users, err := a.getUsersByUsernames(usernames) + if err != nil { + return err + } + + postsWithData := []postAndData{} + postsForCreateList := []*model.Post{} + postsForOverwriteList := []*model.Post{} + + for _, postData := range data { + var userIds []string + var err *model.AppError + for _, username := range *postData.ChannelMembers { + user := users[username] + userIds = append(userIds, user.Id) + } + + var channel *model.Channel + var ch *model.Channel + if len(userIds) == 2 { + ch, err = a.GetOrCreateDirectChannel(userIds[0], userIds[1]) + if err != nil && err.Id != store.CHANNEL_EXISTS_ERROR { + return model.NewAppError("BulkImport", "app.import.import_direct_post.create_direct_channel.error", nil, err.Error(), http.StatusBadRequest) + } + channel = ch + } else { + ch, err = a.createGroupChannel(userIds, userIds[0]) + if err != nil && err.Id != store.CHANNEL_EXISTS_ERROR { + return model.NewAppError("BulkImport", "app.import.import_direct_post.create_group_channel.error", nil, err.Error(), http.StatusBadRequest) + } + channel = ch + } + + user := users[*postData.User] + + // Check if this post already exists. + posts, err := a.Srv().Store.Post().GetPostsCreatedAt(channel.Id, *postData.CreateAt) if err != nil { - return model.NewAppError("BulkImport", "app.import.import_direct_post.channel_member_not_found.error", nil, err.Error(), http.StatusBadRequest) - } - userIds = append(userIds, user.Id) - } - - var channel *model.Channel - var ch *model.Channel - if len(userIds) == 2 { - ch, err = a.createDirectChannel(userIds[0], userIds[1]) - if err != nil && err.Id != store.CHANNEL_EXISTS_ERROR { - return model.NewAppError("BulkImport", "app.import.import_direct_post.create_direct_channel.error", nil, err.Error(), http.StatusBadRequest) - } - channel = ch - } else { - ch, err = a.createGroupChannel(userIds, userIds[0]) - if err != nil && err.Id != store.CHANNEL_EXISTS_ERROR { - return model.NewAppError("BulkImport", "app.import.import_direct_post.create_group_channel.error", nil, err.Error(), http.StatusBadRequest) - } - channel = ch - } - - var user *model.User - user, err = a.Srv().Store.User().GetByUsername(*data.User) - if err != nil { - return model.NewAppError("BulkImport", "app.import.import_direct_post.user_not_found.error", map[string]interface{}{"Username": *data.User}, "", http.StatusBadRequest) - } - - // Check if this post already exists. - posts, err := a.Srv().Store.Post().GetPostsCreatedAt(channel.Id, *data.CreateAt) - if err != nil { - return err - } - - var post *model.Post - for _, p := range posts { - if p.Message == *data.Message { - post = p - break - } - } - - if post == nil { - post = &model.Post{} - } - - post.ChannelId = channel.Id - post.Message = *data.Message - post.UserId = user.Id - post.CreateAt = *data.CreateAt - - post.Hashtags, _ = model.ParseHashtags(post.Message) - - fileIds, err := a.uploadAttachments(data.Attachments, post, "noteam", dryRun) - if err != nil { - return err - } - for _, fileID := range post.FileIds { - if _, ok := fileIds[fileID]; !ok { - a.Srv().Store.FileInfo().PermanentDelete(fileID) - } - } - post.FileIds = make([]string, 0) - for fileID := range fileIds { - post.FileIds = append(post.FileIds, fileID) - } - - if post.Id == "" { - if _, err = a.Srv().Store.Post().Save(post); err != nil { return err } - } else { - if _, err = a.Srv().Store.Post().Overwrite(post); err != nil { + + var post *model.Post + for _, p := range posts { + if p.Message == *postData.Message { + post = p + break + } + } + + if post == nil { + post = &model.Post{} + } + + post.ChannelId = channel.Id + post.Message = *postData.Message + post.UserId = user.Id + post.CreateAt = *postData.CreateAt + post.Hashtags, _ = model.ParseHashtags(post.Message) + + fileIds, err := a.uploadAttachments(postData.Attachments, post, "noteam", dryRun) + if err != nil { + return err + } + for _, fileID := range post.FileIds { + if _, ok := fileIds[fileID]; !ok { + a.Srv().Store.FileInfo().PermanentDelete(fileID) + } + } + post.FileIds = make([]string, 0) + for fileID := range fileIds { + post.FileIds = append(post.FileIds, fileID) + } + + if len(post.Id) == 0 { + postsForCreateList = append(postsForCreateList, post) + } else { + postsForOverwriteList = append(postsForOverwriteList, post) + } + postsWithData = append(postsWithData, postAndData{post: post, directPostData: postData}) + } + + if len(postsForCreateList) > 0 { + if _, err := a.Srv().Store.Post().SaveMultiple(postsForCreateList); err != nil { return err } } - - if data.FlaggedBy != nil { - var preferences model.Preferences - - for _, username := range *data.FlaggedBy { - var user *model.User - user, err = a.Srv().Store.User().GetByUsername(username) - if err != nil { - return model.NewAppError("BulkImport", "app.import.import_direct_post.user_not_found.error", map[string]interface{}{"Username": username}, "", http.StatusBadRequest) - } - - preferences = append(preferences, model.Preference{ - UserId: user.Id, - Category: model.PREFERENCE_CATEGORY_FLAGGED_POST, - Name: post.Id, - Value: "true", - }) - } - - if len(preferences) > 0 { - if err := a.Srv().Store.Preference().Save(&preferences); err != nil { - return model.NewAppError("BulkImport", "app.import.import_direct_post.save_preferences.error", nil, err.Error(), http.StatusInternalServerError) - } - } + if _, err := a.Srv().Store.Post().OverwriteMultiple(postsForOverwriteList); err != nil { + return err } - if data.Reactions != nil { - for _, reaction := range *data.Reactions { - if err := a.importReaction(&reaction, post, dryRun); err != nil { + for _, postWithData := range postsWithData { + if postWithData.directPostData.FlaggedBy != nil { + var preferences model.Preferences + + for _, username := range *postWithData.directPostData.FlaggedBy { + user := users[username] + + preferences = append(preferences, model.Preference{ + UserId: user.Id, + Category: model.PREFERENCE_CATEGORY_FLAGGED_POST, + Name: postWithData.post.Id, + Value: "true", + }) + } + + if len(preferences) > 0 { + if err := a.Srv().Store.Preference().Save(&preferences); err != nil { + return model.NewAppError("BulkImport", "app.import.import_post.save_preferences.error", nil, err.Error(), http.StatusInternalServerError) + } + } + } + + if postWithData.directPostData.Reactions != nil { + for _, reaction := range *postWithData.directPostData.Reactions { + if err := a.importReaction(&reaction, postWithData.post, dryRun); err != nil { + return err + } + } + } + + if postWithData.directPostData.Replies != nil { + if err := a.importReplies(*postWithData.directPostData.Replies, postWithData.post, "noteam", dryRun); err != nil { return err } } - } - if data.Replies != nil { - for _, reply := range *data.Replies { - if err := a.importReply(&reply, post, "noteam", dryRun); err != nil { - return err - } - } + a.updateFileInfoWithPostId(postWithData.post) } - - a.updateFileInfoWithPostId(post) return nil } diff --git a/app/import_functions_test.go b/app/import_functions_test.go index fc18c32667..ef57c39b6b 100644 --- a/app/import_functions_test.go +++ b/app/import_functions_test.go @@ -1565,7 +1565,7 @@ func TestImportUserDefaultNotifyProps(t *testing.T) { } } -func TestImportImportPost(t *testing.T) { +func TestImportimportMultiplePosts(t *testing.T) { th := Setup(t) defer th.TearDown() @@ -1609,7 +1609,7 @@ func TestImportImportPost(t *testing.T) { Channel: &channelName, User: &username, } - err = th.App.importPost(data, true) + err = th.App.importMultiplePosts([]*PostImportData{data}, true) assert.NotNil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) @@ -1621,7 +1621,7 @@ func TestImportImportPost(t *testing.T) { Message: ptrStr("Hello"), CreateAt: ptrInt64(model.GetMillis()), } - err = th.App.importPost(data, true) + err = th.App.importMultiplePosts([]*PostImportData{data}, true) assert.Nil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) @@ -1632,7 +1632,7 @@ func TestImportImportPost(t *testing.T) { User: &username, CreateAt: ptrInt64(model.GetMillis()), } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) assert.NotNil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) @@ -1644,7 +1644,7 @@ func TestImportImportPost(t *testing.T) { Message: ptrStr("Message"), CreateAt: ptrInt64(model.GetMillis()), } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) assert.NotNil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) @@ -1656,7 +1656,7 @@ func TestImportImportPost(t *testing.T) { Message: ptrStr("Message"), CreateAt: ptrInt64(model.GetMillis()), } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) assert.NotNil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) @@ -1668,7 +1668,7 @@ func TestImportImportPost(t *testing.T) { Message: ptrStr("Message"), CreateAt: ptrInt64(model.GetMillis()), } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) assert.NotNil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) @@ -1681,7 +1681,7 @@ func TestImportImportPost(t *testing.T) { Message: ptrStr("Message"), CreateAt: &time, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) assert.Nil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 1, team.Id) @@ -1703,7 +1703,7 @@ func TestImportImportPost(t *testing.T) { Message: ptrStr("Message"), CreateAt: &time, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) assert.Nil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 1, team.Id) @@ -1726,7 +1726,7 @@ func TestImportImportPost(t *testing.T) { Message: ptrStr("Message"), CreateAt: &newTime, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) assert.Nil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 2, team.Id) @@ -1738,7 +1738,7 @@ func TestImportImportPost(t *testing.T) { Message: ptrStr("Message 2"), CreateAt: &time, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) assert.Nil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 3, team.Id) @@ -1751,7 +1751,7 @@ func TestImportImportPost(t *testing.T) { Message: ptrStr("Message 2 #hashtagmashupcity"), CreateAt: &hashtagTime, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) assert.Nil(t, err) AssertAllPostsCount(t, th.App, initialPostCount, 4, team.Id) @@ -1788,7 +1788,7 @@ func TestImportImportPost(t *testing.T) { }, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) require.Nil(t, err, "Expected success.") AssertAllPostsCount(t, th.App, initialPostCount, 5, team.Id) @@ -1821,7 +1821,7 @@ func TestImportImportPost(t *testing.T) { CreateAt: &reactionTime, }}, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) require.Nil(t, err, "Expected success.") AssertAllPostsCount(t, th.App, initialPostCount, 6, team.Id) @@ -1856,7 +1856,7 @@ func TestImportImportPost(t *testing.T) { CreateAt: &replyTime, }}, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) require.Nil(t, err, "Expected success.") AssertAllPostsCount(t, th.App, initialPostCount, 8, team.Id) @@ -1896,7 +1896,7 @@ func TestImportImportPost(t *testing.T) { CreateAt: &replyTime, }}, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) require.Nil(t, err, "Expected success.") AssertAllPostsCount(t, th.App, initialPostCount, 8, team.Id) @@ -1914,7 +1914,7 @@ func TestImportImportPost(t *testing.T) { CreateAt: &replyTime, }}, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) require.Nil(t, err, "Expected success.") AssertAllPostsCount(t, th.App, initialPostCount, 10, team.Id) @@ -1932,12 +1932,403 @@ func TestImportImportPost(t *testing.T) { CreateAt: &replyTime, }}, } - err = th.App.importPost(data, false) + err = th.App.importMultiplePosts([]*PostImportData{data}, false) require.Nil(t, err, "Expected success.") AssertAllPostsCount(t, th.App, initialPostCount, 11, team.Id) } +func TestImportImportPost(t *testing.T) { + th := Setup(t) + defer th.TearDown() + + // Create a Team. + teamName := model.NewRandomTeamName() + th.App.importTeam(&TeamImportData{ + Name: &teamName, + DisplayName: ptrStr("Display Name"), + Type: ptrStr("O"), + }, false) + team, appErr := th.App.GetTeamByName(teamName) + require.Nil(t, appErr, "Failed to get team from database.") + + // Create a Channel. + channelName := model.NewId() + th.App.importChannel(&ChannelImportData{ + Team: &teamName, + Name: &channelName, + DisplayName: ptrStr("Display Name"), + Type: ptrStr("O"), + }, false) + channel, appErr := th.App.GetChannelByName(channelName, team.Id, false) + require.Nil(t, appErr, "Failed to get channel from database.") + + // Create a user. + username := model.NewId() + th.App.importUser(&UserImportData{ + Username: &username, + Email: ptrStr(model.NewId() + "@example.com"), + }, false) + user, appErr := th.App.GetUserByUsername(username) + require.Nil(t, appErr, "Failed to get user from database.") + + username2 := model.NewId() + th.App.importUser(&UserImportData{ + Username: &username2, + Email: ptrStr(model.NewId() + "@example.com"), + }, false) + user2, appErr := th.App.GetUserByUsername(username2) + require.Nil(t, appErr, "Failed to get user from database.") + + // Count the number of posts in the testing team. + initialPostCount, appErr := th.App.Srv().Store.Post().AnalyticsPostCount(team.Id, false, false) + require.Nil(t, appErr) + + time := model.GetMillis() + hashtagTime := time + 2 + replyPostTime := hashtagTime + 4 + replyTime := hashtagTime + 5 + + t.Run("Try adding an invalid post in dry run mode", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, true) + assert.NotNil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) + }) + + t.Run("Try adding a valid post in dry run mode", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + Message: ptrStr("Hello"), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultiplePosts([]*PostImportData{data}, true) + assert.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) + }) + + t.Run("Try adding an invalid post in apply mode", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + assert.NotNil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) + }) + + t.Run("Try adding a valid post with invalid team in apply mode", func(t *testing.T) { + data := &PostImportData{ + Team: ptrStr(model.NewId()), + Channel: &channelName, + User: &username, + Message: ptrStr("Message"), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + assert.NotNil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) + }) + + t.Run("Try adding a valid post with invalid channel in apply mode", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: ptrStr(model.NewId()), + User: &username, + Message: ptrStr("Message"), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + assert.NotNil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) + }) + + t.Run("Try adding a valid post with invalid user in apply mode", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: ptrStr(model.NewId()), + Message: ptrStr("Message"), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + assert.NotNil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, team.Id) + }) + + t.Run("Try adding a valid post in apply mode", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + Message: ptrStr("Message"), + CreateAt: &time, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + assert.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 1, team.Id) + + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(channel.Id, time) + require.Nil(t, err) + + require.Len(t, posts, 1, "Unexpected number of posts found.") + + post := posts[0] + postBool := post.Message != *data.Message || post.CreateAt != *data.CreateAt || post.UserId != user.Id + require.False(t, postBool, "Post properties not as expected") + }) + + t.Run("Update the post", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username2, + Message: ptrStr("Message"), + CreateAt: &time, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + assert.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 1, team.Id) + + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(channel.Id, time) + require.Nil(t, err) + + require.Len(t, posts, 1, "Unexpected number of posts found.") + + post := posts[0] + postBool := post.Message != *data.Message || post.CreateAt != *data.CreateAt || post.UserId != user2.Id + require.False(t, postBool, "Post properties not as expected") + }) + + t.Run("Save the post with a different time", func(t *testing.T) { + newTime := time + 1 + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + Message: ptrStr("Message"), + CreateAt: &newTime, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + assert.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 2, team.Id) + }) + + t.Run("Save the post with a different message", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + Message: ptrStr("Message 2"), + CreateAt: &time, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + assert.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 3, team.Id) + }) + + t.Run("Test with hashtag", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + Message: ptrStr("Message 2 #hashtagmashupcity"), + CreateAt: &hashtagTime, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + assert.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 4, team.Id) + + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(channel.Id, hashtagTime) + require.Nil(t, err) + + require.Len(t, posts, 1, "Unexpected number of posts found.") + + post := posts[0] + postBool := post.Message != *data.Message || post.CreateAt != *data.CreateAt || post.UserId != user.Id + require.False(t, postBool, "Post properties not as expected") + + require.Equal(t, "#hashtagmashupcity", post.Hashtags, "Hashtags not as expected: %s", post.Hashtags) + }) + + t.Run("Post with flags", func(t *testing.T) { + flagsTime := hashtagTime + 1 + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + Message: ptrStr("Message with Favorites"), + CreateAt: &flagsTime, + FlaggedBy: &[]string{ + username, + username2, + }, + } + + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 5, team.Id) + + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(channel.Id, flagsTime) + require.Nil(t, err) + + require.Len(t, posts, 1, "Unexpected number of posts found.") + + post := posts[0] + postBool := post.Message != *data.Message || post.CreateAt != *data.CreateAt || post.UserId != user.Id + require.False(t, postBool, "Post properties not as expected") + + checkPreference(t, th.App, user.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") + checkPreference(t, th.App, user2.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") + }) + + t.Run("Post with reaction", func(t *testing.T) { + reactionPostTime := hashtagTime + 2 + reactionTime := hashtagTime + 3 + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + Message: ptrStr("Message with reaction"), + CreateAt: &reactionPostTime, + Reactions: &[]ReactionImportData{{ + User: &user2.Username, + EmojiName: ptrStr("+1"), + CreateAt: &reactionTime, + }}, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 6, team.Id) + + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(channel.Id, reactionPostTime) + require.Nil(t, err) + + require.Len(t, posts, 1, "Unexpected number of posts found.") + + post := posts[0] + postBool := post.Message != *data.Message || post.CreateAt != *data.CreateAt || post.UserId != user.Id || !post.HasReactions + require.False(t, postBool, "Post properties not as expected") + + reactions, err := th.App.Srv().Store.Reaction().GetForPost(post.Id, false) + require.Nil(t, err, "Can't get reaction") + + require.Len(t, reactions, 1, "Invalid number of reactions") + }) + + t.Run("Post with reply", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &username, + Message: ptrStr("Message with reply"), + CreateAt: &replyPostTime, + Replies: &[]ReplyImportData{{ + User: &user2.Username, + Message: ptrStr("Message reply"), + CreateAt: &replyTime, + }}, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 8, team.Id) + + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(channel.Id, replyPostTime) + require.Nil(t, err) + + require.Len(t, posts, 1, "Unexpected number of posts found.") + + post := posts[0] + postBool := post.Message != *data.Message || post.CreateAt != *data.CreateAt || post.UserId != user.Id + require.False(t, postBool, "Post properties not as expected") + + // Check the reply values. + replies, err := th.App.Srv().Store.Post().GetPostsCreatedAt(channel.Id, replyTime) + require.Nil(t, err) + + require.Len(t, replies, 1, "Unexpected number of posts found.") + + reply := replies[0] + replyBool := reply.Message != *(*data.Replies)[0].Message || reply.CreateAt != *(*data.Replies)[0].CreateAt || reply.UserId != user2.Id + require.False(t, replyBool, "Post properties not as expected") + + require.Equal(t, post.Id, reply.RootId, "Unexpected reply RootId") + }) + + t.Run("Update post with replies", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &user2.Username, + Message: ptrStr("Message with reply"), + CreateAt: &replyPostTime, + Replies: &[]ReplyImportData{{ + User: &username, + Message: ptrStr("Message reply"), + CreateAt: &replyTime, + }}, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 8, team.Id) + }) + + t.Run("Create new post with replies based on the previous one", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &user2.Username, + Message: ptrStr("Message with reply 2"), + CreateAt: &replyPostTime, + Replies: &[]ReplyImportData{{ + User: &username, + Message: ptrStr("Message reply"), + CreateAt: &replyTime, + }}, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 10, team.Id) + }) + + t.Run("Create new reply for existing post with replies", func(t *testing.T) { + data := &PostImportData{ + Team: &teamName, + Channel: &channelName, + User: &user2.Username, + Message: ptrStr("Message with reply"), + CreateAt: &replyPostTime, + Replies: &[]ReplyImportData{{ + User: &username, + Message: ptrStr("Message reply 2"), + CreateAt: &replyTime, + }}, + } + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 11, team.Id) + }) +} + func TestImportImportDirectChannel(t *testing.T) { th := Setup(t).InitBasic() defer th.TearDown() @@ -2117,156 +2508,198 @@ func TestImportImportDirectPost(t *testing.T) { th.BasicUser2.Username, }, } - err := th.App.importDirectChannel(&channelData, false) - require.Nil(t, err) + appErr := th.App.importDirectChannel(&channelData, false) + require.Nil(t, appErr) // Get the channel. var directChannel *model.Channel - channel, err := th.App.GetOrCreateDirectChannel(th.BasicUser.Id, th.BasicUser2.Id) - require.Nil(t, err) + channel, appErr := th.App.GetOrCreateDirectChannel(th.BasicUser.Id, th.BasicUser2.Id) + require.Nil(t, appErr) require.NotEmpty(t, channel) directChannel = channel // Get the number of posts in the system. - result, err := th.App.Srv().Store.Post().AnalyticsPostCount("", false, false) - require.Nil(t, err) + result, appErr := th.App.Srv().Store.Post().AnalyticsPostCount("", false, false) + require.Nil(t, appErr) initialPostCount := result + initialDate := model.GetMillis() - // Try adding an invalid post in dry run mode. - data := &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - }, - User: ptrStr(th.BasicUser.Username), - CreateAt: ptrInt64(model.GetMillis()), - } - err = th.App.importDirectPost(data, true) - require.NotNil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + t.Run("Try adding an invalid post in dry run mode", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, true) + require.NotNil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + }) - // Try adding a valid post in dry run mode. - data = &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - }, - User: ptrStr(th.BasicUser.Username), - Message: ptrStr("Message"), - CreateAt: ptrInt64(model.GetMillis()), - } - err = th.App.importDirectPost(data, true) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + t.Run("Try adding a valid post in dry run mode", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, true) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + }) - // Try adding an invalid post in apply mode. - data = &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - model.NewId(), - }, - User: ptrStr(th.BasicUser.Username), - Message: ptrStr("Message"), - CreateAt: ptrInt64(model.GetMillis()), - } - err = th.App.importDirectPost(data, false) - require.NotNil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + t.Run("Try adding an invalid post in apply mode", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + model.NewId(), + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.NotNil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + }) - // Try adding a valid post in apply mode. - data = &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - }, - User: ptrStr(th.BasicUser.Username), - Message: ptrStr("Message"), - CreateAt: ptrInt64(model.GetMillis()), - } - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 1, "") + t.Run("Try adding a valid post in apply mode", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(initialDate), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 1, "") - // Check the post values. - posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(directChannel.Id, *data.CreateAt) - require.Nil(t, err) - require.Len(t, posts, 1) + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(directChannel.Id, *data.CreateAt) + require.Nil(t, err) + require.Len(t, posts, 1) - post := posts[0] - require.Equal(t, post.Message, *data.Message) - require.Equal(t, post.CreateAt, *data.CreateAt) - require.Equal(t, post.UserId, th.BasicUser.Id) + post := posts[0] + require.Equal(t, post.Message, *data.Message) + require.Equal(t, post.CreateAt, *data.CreateAt) + require.Equal(t, post.UserId, th.BasicUser.Id) + }) - // Import the post again. - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 1, "") + t.Run("Import the post again", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(initialDate), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 1, "") - // Check the post values. - posts, err = th.App.Srv().Store.Post().GetPostsCreatedAt(directChannel.Id, *data.CreateAt) - require.Nil(t, err) - require.Len(t, posts, 1) + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(directChannel.Id, *data.CreateAt) + require.Nil(t, err) + require.Len(t, posts, 1) - post = posts[0] - require.Equal(t, post.Message, *data.Message) - require.Equal(t, post.CreateAt, *data.CreateAt) - require.Equal(t, post.UserId, th.BasicUser.Id) + post := posts[0] + require.Equal(t, post.Message, *data.Message) + require.Equal(t, post.CreateAt, *data.CreateAt) + require.Equal(t, post.UserId, th.BasicUser.Id) + }) - // Save the post with a different time. - data.CreateAt = ptrInt64(*data.CreateAt + 1) - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 2, "") + t.Run("Save the post with a different time", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(initialDate + 1), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 2, "") + }) - // Save the post with a different message. - data.Message = ptrStr("Message 2") - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 3, "") + t.Run("Save the post with a different message", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message 2"), + CreateAt: ptrInt64(initialDate + 1), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 3, "") + }) - // Test with hashtags - data.Message = ptrStr("Message 2 #hashtagmashupcity") - data.CreateAt = ptrInt64(*data.CreateAt + 1) - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 4, "") + t.Run("Test with hashtag", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message 2 #hashtagmashupcity"), + CreateAt: ptrInt64(initialDate + 2), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 4, "") - posts, err = th.App.Srv().Store.Post().GetPostsCreatedAt(directChannel.Id, *data.CreateAt) - require.Nil(t, err) - require.Len(t, posts, 1) + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(directChannel.Id, *data.CreateAt) + require.Nil(t, err) + require.Len(t, posts, 1) - post = posts[0] - require.Equal(t, post.Message, *data.Message) - require.Equal(t, post.CreateAt, *data.CreateAt) - require.Equal(t, post.UserId, th.BasicUser.Id) - require.Equal(t, post.Hashtags, "#hashtagmashupcity") + post := posts[0] + require.Equal(t, post.Message, *data.Message) + require.Equal(t, post.CreateAt, *data.CreateAt) + require.Equal(t, post.UserId, th.BasicUser.Id) + require.Equal(t, post.Hashtags, "#hashtagmashupcity") + }) - // Test with some flags. - data = &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - }, - FlaggedBy: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - }, - User: ptrStr(th.BasicUser.Username), - Message: ptrStr("Message"), - CreateAt: ptrInt64(model.GetMillis()), - } + t.Run("Test with some flags", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + FlaggedBy: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(model.GetMillis()), + } - err = th.App.importDirectPost(data, false) - require.Nil(t, err) + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) - // Check the post values. - posts, err = th.App.Srv().Store.Post().GetPostsCreatedAt(directChannel.Id, *data.CreateAt) - require.Nil(t, err) - require.Len(t, posts, 1) + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(directChannel.Id, *data.CreateAt) + require.Nil(t, err) + require.Len(t, posts, 1) - post = posts[0] - checkPreference(t, th.App, th.BasicUser.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") - checkPreference(t, th.App, th.BasicUser2.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") + post := posts[0] + checkPreference(t, th.App, th.BasicUser.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") + checkPreference(t, th.App, th.BasicUser2.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") + }) // ------------------ Group Channel ------------------------- @@ -2279,8 +2712,8 @@ func TestImportImportDirectPost(t *testing.T) { user3.Username, }, } - err = th.App.importDirectChannel(&channelData, false) - require.Nil(t, err) + appErr = th.App.importDirectChannel(&channelData, false) + require.Nil(t, appErr) // Get the channel. var groupChannel *model.Channel @@ -2289,157 +2722,360 @@ func TestImportImportDirectPost(t *testing.T) { th.BasicUser2.Id, user3.Id, } - channel, err = th.App.createGroupChannel(userIds, th.BasicUser.Id) - require.Equal(t, err.Id, store.CHANNEL_EXISTS_ERROR) + channel, appErr = th.App.createGroupChannel(userIds, th.BasicUser.Id) + require.Equal(t, appErr.Id, store.CHANNEL_EXISTS_ERROR) groupChannel = channel // Get the number of posts in the system. - result, err = th.App.Srv().Store.Post().AnalyticsPostCount("", false, false) - require.Nil(t, err) + result, appErr = th.App.Srv().Store.Post().AnalyticsPostCount("", false, false) + require.Nil(t, appErr) initialPostCount = result - // Try adding an invalid post in dry run mode. - data = &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - user3.Username, - }, - User: ptrStr(th.BasicUser.Username), - CreateAt: ptrInt64(model.GetMillis()), - } - err = th.App.importDirectPost(data, true) - require.NotNil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + t.Run("Try adding an invalid post in dry run mode", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser.Username), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, true) + require.NotNil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + }) - // Try adding a valid post in dry run mode. - data = &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - user3.Username, - }, - User: ptrStr(th.BasicUser.Username), - Message: ptrStr("Message"), - CreateAt: ptrInt64(model.GetMillis()), - } - err = th.App.importDirectPost(data, true) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + t.Run("Try adding a valid post in dry run mode", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, true) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + }) - // Try adding an invalid post in apply mode. - data = &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - user3.Username, - model.NewId(), - }, - User: ptrStr(th.BasicUser.Username), - Message: ptrStr("Message"), - CreateAt: ptrInt64(model.GetMillis()), - } - err = th.App.importDirectPost(data, false) - require.NotNil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + t.Run("Try adding an invalid post in apply mode", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + model.NewId(), + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(model.GetMillis()), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.NotNil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 0, "") + }) - // Try adding a valid post in apply mode. - data = &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - user3.Username, - }, - User: ptrStr(th.BasicUser.Username), - Message: ptrStr("Message"), - CreateAt: ptrInt64(model.GetMillis()), - } - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 1, "") + t.Run("Try adding a valid post in apply mode", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(initialDate + 10), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 1, "") - // Check the post values. - posts, err = th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) - require.Nil(t, err) - require.Len(t, posts, 1) + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) + require.Nil(t, err) + require.Len(t, posts, 1) - post = posts[0] - require.Equal(t, post.Message, *data.Message) - require.Equal(t, post.CreateAt, *data.CreateAt) - require.Equal(t, post.UserId, th.BasicUser.Id) + post := posts[0] + require.Equal(t, post.Message, *data.Message) + require.Equal(t, post.CreateAt, *data.CreateAt) + require.Equal(t, post.UserId, th.BasicUser.Id) + }) - // Import the post again. - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 1, "") + t.Run("Import the post again", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(initialDate + 10), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 1, "") - // Check the post values. - posts, err = th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) - require.Nil(t, err) - require.Len(t, posts, 1) + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) + require.Nil(t, err) + require.Len(t, posts, 1) - post = posts[0] - require.Equal(t, post.Message, *data.Message) - require.Equal(t, post.CreateAt, *data.CreateAt) - require.Equal(t, post.UserId, th.BasicUser.Id) + post := posts[0] + require.Equal(t, post.Message, *data.Message) + require.Equal(t, post.CreateAt, *data.CreateAt) + require.Equal(t, post.UserId, th.BasicUser.Id) + }) - // Save the post with a different time. - data.CreateAt = ptrInt64(*data.CreateAt + 1) - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 2, "") + t.Run("Save the post with a different time", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(initialDate + 11), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 2, "") + }) - // Save the post with a different message. - data.Message = ptrStr("Message 2") - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 3, "") + t.Run("Save the post with a different message", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message 2"), + CreateAt: ptrInt64(initialDate + 11), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 3, "") + }) - // Test with hashtags - data.Message = ptrStr("Message 2 #hashtagmashupcity") - data.CreateAt = ptrInt64(*data.CreateAt + 1) - err = th.App.importDirectPost(data, false) - require.Nil(t, err) - AssertAllPostsCount(t, th.App, initialPostCount, 4, "") + t.Run("Test with hashtag", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message 2 #hashtagmashupcity"), + CreateAt: ptrInt64(initialDate + 12), + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) + AssertAllPostsCount(t, th.App, initialPostCount, 4, "") - posts, err = th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) - require.Nil(t, err) - require.Len(t, posts, 1) + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) + require.Nil(t, err) + require.Len(t, posts, 1) - post = posts[0] - require.Equal(t, post.Message, *data.Message) - require.Equal(t, post.CreateAt, *data.CreateAt) - require.Equal(t, post.UserId, th.BasicUser.Id) - require.Equal(t, post.Hashtags, "#hashtagmashupcity") + post := posts[0] + require.Equal(t, post.Message, *data.Message) + require.Equal(t, post.CreateAt, *data.CreateAt) + require.Equal(t, post.UserId, th.BasicUser.Id) + require.Equal(t, post.Hashtags, "#hashtagmashupcity") + }) - // Test with some flags. - data = &DirectPostImportData{ - ChannelMembers: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - user3.Username, - }, - FlaggedBy: &[]string{ - th.BasicUser.Username, - th.BasicUser2.Username, - }, - User: ptrStr(th.BasicUser.Username), - Message: ptrStr("Message"), - CreateAt: ptrInt64(model.GetMillis()), - } + t.Run("Test with some flags", func(t *testing.T) { + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + FlaggedBy: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message"), + CreateAt: ptrInt64(model.GetMillis()), + } - err = th.App.importDirectPost(data, false) - require.Nil(t, err) + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err) - // Check the post values. - posts, err = th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) - require.Nil(t, err) - require.Len(t, posts, 1) + AssertAllPostsCount(t, th.App, initialPostCount, 5, "") - post = posts[0] - checkPreference(t, th.App, th.BasicUser.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") - checkPreference(t, th.App, th.BasicUser2.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) + require.Nil(t, err) + require.Len(t, posts, 1) + post := posts[0] + checkPreference(t, th.App, th.BasicUser.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") + checkPreference(t, th.App, th.BasicUser2.Id, model.PREFERENCE_CATEGORY_FLAGGED_POST, post.Id, "true") + }) + + t.Run("Post with reaction", func(t *testing.T) { + reactionPostTime := ptrInt64(initialDate + 22) + reactionTime := ptrInt64(initialDate + 23) + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message with reaction"), + CreateAt: reactionPostTime, + Reactions: &[]ReactionImportData{{ + User: ptrStr(th.BasicUser2.Username), + EmojiName: ptrStr("+1"), + CreateAt: reactionTime, + }}, + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 6, "") + + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) + require.Nil(t, err) + + require.Len(t, posts, 1, "Unexpected number of posts found.") + + post := posts[0] + postBool := post.Message != *data.Message || post.CreateAt != *data.CreateAt || post.UserId != th.BasicUser.Id || !post.HasReactions + require.False(t, postBool, "Post properties not as expected") + + reactions, err := th.App.Srv().Store.Reaction().GetForPost(post.Id, false) + require.Nil(t, err, "Can't get reaction") + + require.Len(t, reactions, 1, "Invalid number of reactions") + }) + + t.Run("Post with reply", func(t *testing.T) { + replyPostTime := ptrInt64(initialDate + 25) + replyTime := ptrInt64(initialDate + 26) + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message with reply"), + CreateAt: replyPostTime, + Replies: &[]ReplyImportData{{ + User: ptrStr(th.BasicUser2.Username), + Message: ptrStr("Message reply"), + CreateAt: replyTime, + }}, + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 8, "") + + // Check the post values. + posts, err := th.App.Srv().Store.Post().GetPostsCreatedAt(groupChannel.Id, *data.CreateAt) + require.Nil(t, err) + + require.Len(t, posts, 1, "Unexpected number of posts found.") + + post := posts[0] + postBool := post.Message != *data.Message || post.CreateAt != *data.CreateAt || post.UserId != th.BasicUser.Id + require.False(t, postBool, "Post properties not as expected") + + // Check the reply values. + replies, err := th.App.Srv().Store.Post().GetPostsCreatedAt(channel.Id, *replyTime) + require.Nil(t, err) + + require.Len(t, replies, 1, "Unexpected number of posts found.") + + reply := replies[0] + replyBool := reply.Message != *(*data.Replies)[0].Message || reply.CreateAt != *(*data.Replies)[0].CreateAt || reply.UserId != th.BasicUser2.Id + require.False(t, replyBool, "Post properties not as expected") + + require.Equal(t, post.Id, reply.RootId, "Unexpected reply RootId") + }) + + t.Run("Update post with replies", func(t *testing.T) { + replyPostTime := ptrInt64(initialDate + 25) + replyTime := ptrInt64(initialDate + 26) + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser2.Username), + Message: ptrStr("Message with reply"), + CreateAt: replyPostTime, + Replies: &[]ReplyImportData{{ + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message reply"), + CreateAt: replyTime, + }}, + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 8, "") + }) + + t.Run("Create new post with replies based on the previous one", func(t *testing.T) { + replyPostTime := ptrInt64(initialDate + 27) + replyTime := ptrInt64(initialDate + 28) + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser2.Username), + Message: ptrStr("Message with reply 2"), + CreateAt: replyPostTime, + Replies: &[]ReplyImportData{{ + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message reply"), + CreateAt: replyTime, + }}, + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 10, "") + }) + + t.Run("Create new reply for existing post with replies", func(t *testing.T) { + replyPostTime := ptrInt64(initialDate + 25) + replyTime := ptrInt64(initialDate + 29) + data := &DirectPostImportData{ + ChannelMembers: &[]string{ + th.BasicUser.Username, + th.BasicUser2.Username, + user3.Username, + }, + User: ptrStr(th.BasicUser2.Username), + Message: ptrStr("Message with reply"), + CreateAt: replyPostTime, + Replies: &[]ReplyImportData{{ + User: ptrStr(th.BasicUser.Username), + Message: ptrStr("Message reply 2"), + CreateAt: replyTime, + }}, + } + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{data}, false) + require.Nil(t, err, "Expected success.") + + AssertAllPostsCount(t, th.App, initialPostCount, 11, "") + }) } func TestImportImportEmoji(t *testing.T) { @@ -2516,8 +3152,8 @@ func TestImportPostAndRepliesWithAttachments(t *testing.T) { DisplayName: ptrStr("Display Name"), Type: ptrStr("O"), }, false) - team, err := th.App.GetTeamByName(teamName) - require.Nil(t, err, "Failed to get team from database.") + team, appErr := th.App.GetTeamByName(teamName) + require.Nil(t, appErr, "Failed to get team from database.") // Create a Channel. channelName := model.NewId() @@ -2527,8 +3163,8 @@ func TestImportPostAndRepliesWithAttachments(t *testing.T) { DisplayName: ptrStr("Display Name"), Type: ptrStr("O"), }, false) - _, err = th.App.GetChannelByName(channelName, team.Id, false) - require.Nil(t, err, "Failed to get channel from database.") + _, appErr = th.App.GetChannelByName(channelName, team.Id, false) + require.Nil(t, appErr, "Failed to get channel from database.") // Create a user3. username := model.NewId() @@ -2536,18 +3172,36 @@ func TestImportPostAndRepliesWithAttachments(t *testing.T) { Username: &username, Email: ptrStr(model.NewId() + "@example.com"), }, false) - user3, err := th.App.GetUserByUsername(username) - require.Nil(t, err, "Failed to get user3 from database.") + user3, appErr := th.App.GetUserByUsername(username) + require.Nil(t, appErr, "Failed to get user3 from database.") username2 := model.NewId() th.App.importUser(&UserImportData{ Username: &username2, Email: ptrStr(model.NewId() + "@example.com"), }, false) - user4, err := th.App.GetUserByUsername(username2) - require.Nil(t, err, "Failed to get user3 from database.") + user2, appErr := th.App.GetUserByUsername(username2) + require.Nil(t, appErr, "Failed to get user3 from database.") - // Post with attachments. + // Create direct post users. + username3 := model.NewId() + th.App.importUser(&UserImportData{ + Username: &username3, + Email: ptrStr(model.NewId() + "@example.com"), + }, false) + user3, appErr = th.App.GetUserByUsername(username3) + require.Nil(t, appErr, "Failed to get user3 from database.") + + username4 := model.NewId() + th.App.importUser(&UserImportData{ + Username: &username4, + Email: ptrStr(model.NewId() + "@example.com"), + }, false) + + user4, appErr := th.App.GetUserByUsername(username4) + require.Nil(t, appErr, "Failed to get user3 from database.") + + // Post with attachments time := model.GetMillis() attachmentsPostTime := time attachmentsReplyTime := time + 1 @@ -2557,7 +3211,7 @@ func TestImportPostAndRepliesWithAttachments(t *testing.T) { data := &PostImportData{ Team: &teamName, Channel: &channelName, - User: &username, + User: &username3, Message: ptrStr("Message with reply"), CreateAt: &attachmentsPostTime, Attachments: &[]AttachmentImportData{{Path: &testImage}, {Path: &testMarkDown}}, @@ -2569,75 +3223,63 @@ func TestImportPostAndRepliesWithAttachments(t *testing.T) { }}, } - // import with attachments - err = th.App.importPost(data, false) - assert.Nil(t, err) + t.Run("import with attachment", func(t *testing.T) { + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + require.Nil(t, err) - attachments := GetAttachments(user3.Id, th, t) - assert.Len(t, attachments, 2) - assert.Contains(t, attachments[0].Path, team.Id) - assert.Contains(t, attachments[1].Path, team.Id) - AssertFileIdsInPost(attachments, th, t) + attachments := GetAttachments(user3.Id, th, t) + require.Len(t, attachments, 2) + assert.Contains(t, attachments[0].Path, team.Id) + assert.Contains(t, attachments[1].Path, team.Id) + AssertFileIdsInPost(attachments, th, t) - // import existing post with new attachments - data.Attachments = &[]AttachmentImportData{{Path: &testImage}} - err = th.App.importPost(data, false) - assert.Nil(t, err) + attachments = GetAttachments(user4.Id, th, t) + require.Len(t, attachments, 1) + assert.Contains(t, attachments[0].Path, team.Id) + AssertFileIdsInPost(attachments, th, t) + }) - attachments = GetAttachments(user3.Id, th, t) - assert.Len(t, attachments, 1) - assert.Contains(t, attachments[0].Path, team.Id) - AssertFileIdsInPost(attachments, th, t) + t.Run("import existing post with new attachment", func(t *testing.T) { + data.Attachments = &[]AttachmentImportData{{Path: &testImage}} + err := th.App.importMultiplePosts([]*PostImportData{data}, false) + require.Nil(t, err) - attachments = GetAttachments(user4.Id, th, t) - assert.Len(t, attachments, 1) - assert.Contains(t, attachments[0].Path, team.Id) - AssertFileIdsInPost(attachments, th, t) + attachments := GetAttachments(user3.Id, th, t) + require.Len(t, attachments, 1) + assert.Contains(t, attachments[0].Path, team.Id) + AssertFileIdsInPost(attachments, th, t) - // Reply with Attachments in Direct Post + attachments = GetAttachments(user4.Id, th, t) + require.Len(t, attachments, 1) + assert.Contains(t, attachments[0].Path, team.Id) + AssertFileIdsInPost(attachments, th, t) + }) - // Create direct post users. + t.Run("Reply with Attachments in Direct Pos", func(t *testing.T) { + directImportData := &DirectPostImportData{ + ChannelMembers: &[]string{ + user3.Username, + user2.Username, + }, + User: &user3.Username, + Message: ptrStr("Message with Replies"), + CreateAt: ptrInt64(model.GetMillis()), + Replies: &[]ReplyImportData{{ + User: &user2.Username, + Message: ptrStr("Message reply with attachment"), + CreateAt: ptrInt64(model.GetMillis()), + Attachments: &[]AttachmentImportData{{Path: &testImage}}, + }}, + } - username3 := model.NewId() - th.App.importUser(&UserImportData{ - Username: &username3, - Email: ptrStr(model.NewId() + "@example.com"), - }, false) - user3, err = th.App.GetUserByUsername(username3) - require.Nil(t, err, "Failed to get user3 from database.") + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{directImportData}, false) + require.Nil(t, err, "Expected success.") - username4 := model.NewId() - th.App.importUser(&UserImportData{ - Username: &username4, - Email: ptrStr(model.NewId() + "@example.com"), - }, false) - - user4, err = th.App.GetUserByUsername(username4) - require.Nil(t, err, "Failed to get user3 from database.") - - directImportData := &DirectPostImportData{ - ChannelMembers: &[]string{ - user3.Username, - user4.Username, - }, - User: &user3.Username, - Message: ptrStr("Message with Replies"), - CreateAt: ptrInt64(model.GetMillis()), - Replies: &[]ReplyImportData{{ - User: &user4.Username, - Message: ptrStr("Message reply with attachment"), - CreateAt: ptrInt64(model.GetMillis()), - Attachments: &[]AttachmentImportData{{Path: &testImage}}, - }}, - } - - err = th.App.importDirectPost(directImportData, false) - require.Nil(t, err, "Expected success.") - - attachments = GetAttachments(user4.Id, th, t) - assert.Len(t, attachments, 1) - assert.Contains(t, attachments[0].Path, "noteam") - AssertFileIdsInPost(attachments, th, t) + attachments := GetAttachments(user2.Id, th, t) + require.Len(t, attachments, 1) + assert.Contains(t, attachments[0].Path, "noteam") + AssertFileIdsInPost(attachments, th, t) + }) } func TestImportDirectPostWithAttachments(t *testing.T) { @@ -2662,8 +3304,8 @@ func TestImportDirectPostWithAttachments(t *testing.T) { Username: &username, Email: ptrStr(model.NewId() + "@example.com"), }, false) - user1, err := th.App.GetUserByUsername(username) - require.Nil(t, err, "Failed to get user1 from database.") + user1, appErr := th.App.GetUserByUsername(username) + require.Nil(t, appErr, "Failed to get user1 from database.") username2 := model.NewId() th.App.importUser(&UserImportData{ @@ -2671,8 +3313,8 @@ func TestImportDirectPostWithAttachments(t *testing.T) { Email: ptrStr(model.NewId() + "@example.com"), }, false) - user2, err := th.App.GetUserByUsername(username2) - require.Nil(t, err, "Failed to get user2 from database.") + user2, appErr := th.App.GetUserByUsername(username2) + require.Nil(t, appErr, "Failed to get user2 from database.") directImportData := &DirectPostImportData{ ChannelMembers: &[]string{ @@ -2686,21 +3328,21 @@ func TestImportDirectPostWithAttachments(t *testing.T) { } t.Run("Regular import of attachment", func(t *testing.T) { - err := th.App.importDirectPost(directImportData, false) + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{directImportData}, false) require.Nil(t, err, "Expected success.") attachments := GetAttachments(user1.Id, th, t) - assert.Len(t, attachments, 1) + require.Len(t, attachments, 1) assert.Contains(t, attachments[0].Path, "noteam") AssertFileIdsInPost(attachments, th, t) }) t.Run("Attempt to import again with same file entirely, should NOT add an attachment", func(t *testing.T) { - err := th.App.importDirectPost(directImportData, false) + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{directImportData}, false) require.Nil(t, err, "Expected success.") attachments := GetAttachments(user1.Id, th, t) - assert.Len(t, attachments, 1) + require.Len(t, attachments, 1) }) t.Run("Attempt to import again with same name and size but different content, SHOULD add an attachment", func(t *testing.T) { @@ -2715,11 +3357,11 @@ func TestImportDirectPostWithAttachments(t *testing.T) { Attachments: &[]AttachmentImportData{{Path: &testImageFake}}, } - err := th.App.importDirectPost(directImportDataFake, false) + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{directImportDataFake}, false) require.Nil(t, err, "Expected success.") attachments := GetAttachments(user1.Id, th, t) - assert.Len(t, attachments, 2) + require.Len(t, attachments, 2) }) t.Run("Attempt to import again with same data, SHOULD add an attachment, since it's different name", func(t *testing.T) { @@ -2734,10 +3376,10 @@ func TestImportDirectPostWithAttachments(t *testing.T) { Attachments: &[]AttachmentImportData{{Path: &testImage2}}, } - err := th.App.importDirectPost(directImportData2, false) + err := th.App.importMultipleDirectPosts([]*DirectPostImportData{directImportData2}, false) require.Nil(t, err, "Expected success.") attachments := GetAttachments(user1.Id, th, t) - assert.Len(t, attachments, 3) + require.Len(t, attachments, 3) }) } diff --git a/app/import_test.go b/app/import_test.go index 517d455b77..e9efc571ec 100644 --- a/app/import_test.go +++ b/app/import_test.go @@ -239,12 +239,12 @@ func GetAttachments(userId string, th *TestHelper, t *testing.T) []*model.FileIn func AssertFileIdsInPost(files []*model.FileInfo, th *TestHelper, t *testing.T) { postId := files[0].PostId - assert.NotNil(t, postId) + require.NotNil(t, postId) posts, err := th.App.Srv().Store.Post().GetPostsByIds([]string{postId}) require.Nil(t, err) - assert.Equal(t, len(posts), 1) + require.Len(t, posts, 1) for _, file := range files { assert.Contains(t, posts[0].FileIds, file.Id) } diff --git a/cmd/mattermost/commands/sampledata.go b/cmd/mattermost/commands/sampledata.go index 20423a87f8..725eccc2e5 100644 --- a/cmd/mattermost/commands/sampledata.go +++ b/cmd/mattermost/commands/sampledata.go @@ -325,6 +325,11 @@ func sampleDataCmdF(command *cobra.Command, args []string) error { user2 := allUsers[rand.Intn(len(allUsers))] channelLine := createDirectChannel([]string{user1, user2}) encoder.Encode(channelLine) + } + + for i := 0; i < directChannels; i++ { + user1 := allUsers[rand.Intn(len(allUsers))] + user2 := allUsers[rand.Intn(len(allUsers))] dates := sortedRandomDates(postsPerDirectChannel) for j := 0; j < postsPerDirectChannel; j++ { @@ -344,6 +349,17 @@ func sampleDataCmdF(command *cobra.Command, args []string) error { } channelLine := createDirectChannel(users) encoder.Encode(channelLine) + } + + for i := 0; i < groupChannels; i++ { + users := []string{} + totalUsers := 3 + rand.Intn(3) + for len(users) < totalUsers { + user := allUsers[rand.Intn(len(allUsers))] + if !sliceIncludes(users, user) { + users = append(users, user) + } + } dates := sortedRandomDates(postsPerGroupChannel) for j := 0; j < postsPerGroupChannel; j++ { diff --git a/i18n/en.json b/i18n/en.json index 63d6bfe240..2f1bf7f2ab 100644 --- a/i18n/en.json +++ b/i18n/en.json @@ -2958,6 +2958,14 @@ "id": "app.import.emoji.bad_file.error", "translation": "Error reading import emoji image file. Emoji with name: \"{{.EmojiName}}\"" }, + { + "id": "app.import.get_teams_by_names.some_teams_not_found.error", + "translation": "Some teams not found" + }, + { + "id": "app.import.get_users_by_username.some_users_not_found.error", + "translation": "Some users not found" + }, { "id": "app.import.import_channel.scheme_deleted.error", "translation": "Unable to set a channel to use a deleted scheme." @@ -2978,18 +2986,10 @@ "id": "app.import.import_direct_channel.create_group_channel.error", "translation": "Failed to create group channel" }, - { - "id": "app.import.import_direct_channel.member_not_found.error", - "translation": "Could not find channel member when importing direct channel" - }, { "id": "app.import.import_direct_channel.update_header_failed.error", "translation": "Failed to update direct channel header" }, - { - "id": "app.import.import_direct_post.channel_member_not_found.error", - "translation": "Could not find channel member when importing direct channel post" - }, { "id": "app.import.import_direct_post.create_direct_channel.error", "translation": "Failed to get direct channel" @@ -2998,14 +2998,6 @@ "id": "app.import.import_direct_post.create_group_channel.error", "translation": "Failed to get group channel" }, - { - "id": "app.import.import_direct_post.save_preferences.error", - "translation": "Error importing direct post. Failed to save preferences." - }, - { - "id": "app.import.import_direct_post.user_not_found.error", - "translation": "Post user does not exist" - }, { "id": "app.import.import_line.null_channel.error", "translation": "Import data line has type \"channel\" but the channel object is null." @@ -3050,10 +3042,6 @@ "id": "app.import.import_post.save_preferences.error", "translation": "Error importing post. Failed to save preferences." }, - { - "id": "app.import.import_post.team_not_found.error", - "translation": "Error importing post. Team with name \"{{.TeamName}}\" could not be found." - }, { "id": "app.import.import_post.user_not_found.error", "translation": "Error importing post. User with username \"{{.Username}}\" could not be found." @@ -6982,6 +6970,14 @@ "id": "store.sql_team.get_by_name.missing.app_error", "translation": "Unable to find the existing team." }, + { + "id": "store.sql_team.get_by_names.app_error", + "translation": "Unable to get the teams by names" + }, + { + "id": "store.sql_team.get_by_names.missing.app_error", + "translation": "Unable to find some of the requested teams" + }, { "id": "store.sql_team.get_by_scheme.app_error", "translation": "Unable to get the channels for the provided scheme." diff --git a/store/sqlstore/post_store.go b/store/sqlstore/post_store.go index 797651bb9e..78c6d82172 100644 --- a/store/sqlstore/post_store.go +++ b/store/sqlstore/post_store.go @@ -30,6 +30,33 @@ type SqlPostStore struct { func (s *SqlPostStore) ClearCaches() { } +func postSliceColumns() []string { + return []string{"Id", "CreateAt", "UpdateAt", "EditAt", "DeleteAt", "IsPinned", "UserId", "ChannelId", "RootId", "ParentId", "OriginalId", "Message", "Type", "Props", "Hashtags", "Filenames", "FileIds", "HasReactions"} +} + +func postToSlice(post *model.Post) []interface{} { + return []interface{}{ + post.Id, + post.CreateAt, + post.UpdateAt, + post.EditAt, + post.DeleteAt, + post.IsPinned, + post.UserId, + post.ChannelId, + post.RootId, + post.ParentId, + post.OriginalId, + post.Message, + post.Type, + model.StringInterfaceToJson(post.Props), + post.Hashtags, + model.ArrayToJson(post.Filenames), + model.ArrayToJson(post.FileIds), + post.HasReactions, + } +} + func newSqlPostStore(sqlStore SqlStore, metrics einterfaces.MetricsInterface) store.PostStore { s := &SqlPostStore{ SqlStore: sqlStore, @@ -72,48 +99,97 @@ func (s *SqlPostStore) createIndexesIfNotExists() { s.CreateFullTextIndexIfNotExists("idx_posts_hashtags_txt", "Posts", "Hashtags") } -func (s *SqlPostStore) Save(post *model.Post) (*model.Post, *model.AppError) { - if len(post.Id) > 0 { - return nil, model.NewAppError("SqlPostStore.Save", "store.sql_post.save.existing.app_error", nil, "id="+post.Id, http.StatusBadRequest) - } - - maxPostSize := s.GetMaxPostSize() - - post.PreSave() - if err := post.IsValid(maxPostSize); err != nil { - return nil, err - } - - if err := s.GetMaster().Insert(post); err != nil { - return nil, model.NewAppError("SqlPostStore.Save", "store.sql_post.save.app_error", nil, "id="+post.Id+", "+err.Error(), http.StatusInternalServerError) - } - - time := post.UpdateAt - - if !post.IsJoinLeaveMessage() { - if _, err := s.GetMaster().Exec("UPDATE Channels SET LastPostAt = GREATEST(:LastPostAt, LastPostAt), TotalMsgCount = TotalMsgCount + 1 WHERE Id = :ChannelId", map[string]interface{}{"LastPostAt": time, "ChannelId": post.ChannelId}); err != nil { - mlog.Error("Error updating Channel LastPostAt.", mlog.Err(err)) +func (s *SqlPostStore) SaveMultiple(posts []*model.Post) ([]*model.Post, *model.AppError) { + channelNewPosts := make(map[string]int) + maxDateNewPosts := make(map[string]int64) + rootIds := make(map[string]int) + maxDateRootIds := make(map[string]int64) + for _, post := range posts { + if len(post.Id) > 0 { + return nil, model.NewAppError("SqlPostStore.Save", "store.sql_post.save.existing.app_error", nil, "id="+post.Id, http.StatusBadRequest) } - } else { - // don't update TotalMsgCount for unimportant messages so that the channel isn't marked as unread - if _, err := s.GetMaster().Exec("UPDATE Channels SET LastPostAt = :LastPostAt WHERE Id = :ChannelId AND LastPostAt < :LastPostAt", map[string]interface{}{"LastPostAt": time, "ChannelId": post.ChannelId}); err != nil { + post.PreSave() + maxPostSize := s.GetMaxPostSize() + if err := post.IsValid(maxPostSize); err != nil { + return nil, err + } + + currentChannelCount, ok := channelNewPosts[post.ChannelId] + if !ok { + if post.IsJoinLeaveMessage() { + channelNewPosts[post.ChannelId] = 0 + } else { + channelNewPosts[post.ChannelId] = 1 + } + maxDateNewPosts[post.ChannelId] = post.CreateAt + } else { + if !post.IsJoinLeaveMessage() { + channelNewPosts[post.ChannelId] = currentChannelCount + 1 + } + if post.CreateAt > maxDateNewPosts[post.ChannelId] { + maxDateNewPosts[post.ChannelId] = post.CreateAt + } + } + + if len(post.RootId) == 0 { + continue + } + + currentRootCount, ok := rootIds[post.RootId] + if !ok { + rootIds[post.RootId] = 1 + maxDateRootIds[post.RootId] = post.CreateAt + } else { + rootIds[post.RootId] = currentRootCount + 1 + if post.CreateAt > maxDateRootIds[post.RootId] { + maxDateRootIds[post.RootId] = post.CreateAt + } + } + } + + query := s.getQueryBuilder().Insert("Posts").Columns(postSliceColumns()...) + for _, post := range posts { + query = query.Values(postToSlice(post)...) + } + sql, args, err := query.ToSql() + if err != nil { + return nil, model.NewAppError("SqlPostStore.Save", "store.sql_post.save.app_error", nil, err.Error(), http.StatusInternalServerError) + } + + if _, err := s.GetMaster().Exec(sql, args...); err != nil { + return nil, model.NewAppError("SqlPostStore.Save", "store.sql_post.save.app_error", nil, err.Error(), http.StatusInternalServerError) + } + + for channelId, count := range channelNewPosts { + if _, err := s.GetMaster().Exec("UPDATE Channels SET LastPostAt = GREATEST(:LastPostAt, LastPostAt), TotalMsgCount = TotalMsgCount + :Count WHERE Id = :ChannelId", map[string]interface{}{"LastPostAt": maxDateNewPosts[channelId], "ChannelId": channelId, "Count": count}); err != nil { mlog.Error("Error updating Channel LastPostAt.", mlog.Err(err)) } } - if len(post.RootId) > 0 { - if _, err := s.GetMaster().Exec("UPDATE Posts SET UpdateAt = :UpdateAt WHERE Id = :RootId", map[string]interface{}{"UpdateAt": time, "RootId": post.RootId}); err != nil { + for rootId := range rootIds { + if _, err := s.GetMaster().Exec("UPDATE Posts SET UpdateAt = :UpdateAt WHERE Id = :RootId", map[string]interface{}{"UpdateAt": maxDateRootIds[rootId], "RootId": rootId}); err != nil { mlog.Error("Error updating Post UpdateAt.", mlog.Err(err)) } - } else { - if count, err := s.GetMaster().SelectInt("SELECT COUNT(*) FROM Posts WHERE RootId = :Id", map[string]interface{}{"Id": post.Id}); err != nil { - mlog.Error("Error fetching post's thread.", mlog.Err(err)) - } else { - post.ReplyCount = count + } + + for _, post := range posts { + if len(post.RootId) == 0 { + count, ok := rootIds[post.Id] + if ok { + post.ReplyCount += int64(count) + } } } - return post, nil + return posts, nil +} + +func (s *SqlPostStore) Save(post *model.Post) (*model.Post, *model.AppError) { + posts, err := s.SaveMultiple([]*model.Post{post}) + if err != nil { + return nil, err + } + return posts[0], nil } func (s *SqlPostStore) Update(newPost *model.Post, oldPost *model.Post) (*model.Post, *model.AppError) { @@ -149,19 +225,45 @@ func (s *SqlPostStore) Update(newPost *model.Post, oldPost *model.Post) (*model. return newPost, nil } -func (s *SqlPostStore) Overwrite(post *model.Post) (*model.Post, *model.AppError) { - post.UpdateAt = model.GetMillis() - +func (s *SqlPostStore) OverwriteMultiple(posts []*model.Post) ([]*model.Post, *model.AppError) { + updateAt := model.GetMillis() maxPostSize := s.GetMaxPostSize() - if appErr := post.IsValid(maxPostSize); appErr != nil { - return nil, appErr + for _, post := range posts { + post.UpdateAt = updateAt + if appErr := post.IsValid(maxPostSize); appErr != nil { + return nil, appErr + } } - if _, err := s.GetMaster().Update(post); err != nil { - return nil, model.NewAppError("SqlPostStore.Overwrite", "store.sql_post.overwrite.app_error", nil, "id="+post.Id+", "+err.Error(), http.StatusInternalServerError) + tx, err := s.GetMaster().Begin() + if err != nil { + return nil, model.NewAppError("SqlPostStore.Overwrite", "store.sql_post.overwrite.app_error", nil, err.Error(), http.StatusInternalServerError) + } + for _, post := range posts { + if _, err = tx.Update(post); err != nil { + txErr := tx.Rollback() + if txErr != nil { + return nil, model.NewAppError("SqlPostStore.Overwrite", "store.sql_post.overwrite.app_error", nil, txErr.Error(), http.StatusInternalServerError) + } + + return nil, model.NewAppError("SqlPostStore.Overwrite", "store.sql_post.overwrite.app_error", nil, "id="+post.Id+", "+err.Error(), http.StatusInternalServerError) + } + } + err = tx.Commit() + if err != nil { + return nil, model.NewAppError("SqlPostStore.Overwrite", "store.sql_post.overwrite.app_error", nil, err.Error(), http.StatusInternalServerError) } - return post, nil + return posts, nil +} + +func (s *SqlPostStore) Overwrite(post *model.Post) (*model.Post, *model.AppError) { + posts, err := s.OverwriteMultiple([]*model.Post{post}) + if err != nil { + return nil, err + } + + return posts[0], nil } func (s *SqlPostStore) GetFlaggedPosts(userId string, offset int, limit int) (*model.PostList, *model.AppError) { diff --git a/store/sqlstore/team_store.go b/store/sqlstore/team_store.go index 9adc9917f4..563c94d917 100644 --- a/store/sqlstore/team_store.go +++ b/store/sqlstore/team_store.go @@ -13,6 +13,7 @@ import ( "github.com/mattermost/gorp" "github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/store" + "github.com/mattermost/mattermost-server/v5/utils" ) const ( @@ -287,6 +288,33 @@ func (s SqlTeamStore) GetByName(name string) (*model.Team, *model.AppError) { return &team, nil } +func (s SqlTeamStore) GetByNames(names []string) ([]*model.Team, *model.AppError) { + uniqueNames := utils.RemoveDuplicatesFromStringArray(names) + + query := s.getQueryBuilder(). + Select("*"). + From("Teams"). + Where(sq.Eq{"Name": uniqueNames}) + + queryString, args, err := query.ToSql() + if err != nil { + return nil, model.NewAppError("SqlTeamStore.GetByNames", "store.sql_team.get_by_names.app_error", nil, err.Error(), http.StatusInternalServerError) + } + + teams := []*model.Team{} + _, err = s.GetReplica().Select(&teams, queryString, args...) + if err != nil { + if err == sql.ErrNoRows { + return nil, model.NewAppError("SqlTeamStore.GetByNames", "store.sql_team.get_by_names.missing.app_error", nil, err.Error(), http.StatusNotFound) + } + return nil, model.NewAppError("SqlTeamStore.GetByNames", "store.sql_team.get_by_names.app_error", nil, err.Error(), http.StatusInternalServerError) + } + if len(teams) != len(uniqueNames) { + return nil, model.NewAppError("SqlTeamStore.GetByNames", "store.sql_team.get_by_names.missing.app_error", nil, "", http.StatusNotFound) + } + return teams, nil +} + func (s SqlTeamStore) SearchAll(term string) ([]*model.Team, *model.AppError) { var teams []*model.Team diff --git a/store/store.go b/store/store.go index 29ccdafceb..c4612d7008 100644 --- a/store/store.go +++ b/store/store.go @@ -67,6 +67,7 @@ type TeamStore interface { Update(team *model.Team) (*model.Team, *model.AppError) Get(id string) (*model.Team, *model.AppError) GetByName(name string) (*model.Team, *model.AppError) + GetByNames(name []string) ([]*model.Team, *model.AppError) SearchAll(term string) ([]*model.Team, *model.AppError) SearchAllPaged(term string, page int, perPage int) ([]*model.Team, int64, *model.AppError) SearchOpen(term string) ([]*model.Team, *model.AppError) @@ -218,6 +219,7 @@ type ChannelMemberHistoryStore interface { } type PostStore interface { + SaveMultiple(posts []*model.Post) ([]*model.Post, *model.AppError) Save(post *model.Post) (*model.Post, *model.AppError) Update(newPost *model.Post, oldPost *model.Post) (*model.Post, *model.AppError) Get(id string, skipFetchThreads bool) (*model.PostList, *model.AppError) @@ -245,6 +247,7 @@ type PostStore interface { InvalidateLastPostTimeCache(channelId string) GetPostsCreatedAt(channelId string, time int64) ([]*model.Post, *model.AppError) Overwrite(post *model.Post) (*model.Post, *model.AppError) + OverwriteMultiple(posts []*model.Post) ([]*model.Post, *model.AppError) GetPostsByIds(postIds []string) ([]*model.Post, *model.AppError) GetPostsBatchForIndexing(startTime int64, endTime int64, limit int) ([]*model.PostForIndexing, *model.AppError) PermanentDeleteBatch(endTime int64, limit int64) (int64, *model.AppError) diff --git a/store/storetest/mocks/PostStore.go b/store/storetest/mocks/PostStore.go index bd879d60d0..d8442ab6fa 100644 --- a/store/storetest/mocks/PostStore.go +++ b/store/storetest/mocks/PostStore.go @@ -637,6 +637,31 @@ func (_m *PostStore) Overwrite(post *model.Post) (*model.Post, *model.AppError) return r0, r1 } +// OverwriteMultiple provides a mock function with given fields: posts +func (_m *PostStore) OverwriteMultiple(posts []*model.Post) ([]*model.Post, *model.AppError) { + ret := _m.Called(posts) + + var r0 []*model.Post + if rf, ok := ret.Get(0).(func([]*model.Post) []*model.Post); ok { + r0 = rf(posts) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.Post) + } + } + + var r1 *model.AppError + if rf, ok := ret.Get(1).(func([]*model.Post) *model.AppError); ok { + r1 = rf(posts) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 +} + // PermanentDeleteBatch provides a mock function with given fields: endTime, limit func (_m *PostStore) PermanentDeleteBatch(endTime int64, limit int64) (int64, *model.AppError) { ret := _m.Called(endTime, limit) @@ -717,6 +742,31 @@ func (_m *PostStore) Save(post *model.Post) (*model.Post, *model.AppError) { return r0, r1 } +// SaveMultiple provides a mock function with given fields: posts +func (_m *PostStore) SaveMultiple(posts []*model.Post) ([]*model.Post, *model.AppError) { + ret := _m.Called(posts) + + var r0 []*model.Post + if rf, ok := ret.Get(0).(func([]*model.Post) []*model.Post); ok { + r0 = rf(posts) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.Post) + } + } + + var r1 *model.AppError + if rf, ok := ret.Get(1).(func([]*model.Post) *model.AppError); ok { + r1 = rf(posts) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 +} + // Search provides a mock function with given fields: teamId, userId, params func (_m *PostStore) Search(teamId string, userId string, params *model.SearchParams) (*model.PostList, *model.AppError) { ret := _m.Called(teamId, userId, params) diff --git a/store/storetest/mocks/TeamStore.go b/store/storetest/mocks/TeamStore.go index 9807d391ff..33f4b476b8 100644 --- a/store/storetest/mocks/TeamStore.go +++ b/store/storetest/mocks/TeamStore.go @@ -425,6 +425,31 @@ func (_m *TeamStore) GetByName(name string) (*model.Team, *model.AppError) { return r0, r1 } +// GetByNames provides a mock function with given fields: name +func (_m *TeamStore) GetByNames(name []string) ([]*model.Team, *model.AppError) { + ret := _m.Called(name) + + var r0 []*model.Team + if rf, ok := ret.Get(0).(func([]string) []*model.Team); ok { + r0 = rf(name) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).([]*model.Team) + } + } + + var r1 *model.AppError + if rf, ok := ret.Get(1).(func([]string) *model.AppError); ok { + r1 = rf(name) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 +} + // GetChannelUnreadsForAllTeams provides a mock function with given fields: excludeTeamId, userId func (_m *TeamStore) GetChannelUnreadsForAllTeams(excludeTeamId string, userId string) ([]*model.ChannelUnread, *model.AppError) { ret := _m.Called(excludeTeamId, userId) diff --git a/store/storetest/post_store.go b/store/storetest/post_store.go index 13fe753f05..8d1f49be2b 100644 --- a/store/storetest/post_store.go +++ b/store/storetest/post_store.go @@ -18,6 +18,7 @@ import ( ) func TestPostStore(t *testing.T, ss store.Store, s SqlSupplier) { + t.Run("SaveMultiple", func(t *testing.T) { testPostStoreSaveMultiple(t, ss) }) t.Run("Save", func(t *testing.T) { testPostStoreSave(t, ss) }) t.Run("SaveAndUpdateChannelMsgCounts", func(t *testing.T) { testPostStoreSaveChannelMsgCounts(t, ss) }) t.Run("Get", func(t *testing.T) { testPostStoreGet(t, ss) }) @@ -41,6 +42,7 @@ func TestPostStore(t *testing.T, ss store.Store, s SqlSupplier) { t.Run("GetFlaggedPostsForChannel", func(t *testing.T) { testPostStoreGetFlaggedPostsForChannel(t, ss) }) t.Run("GetPostsCreatedAt", func(t *testing.T) { testPostStoreGetPostsCreatedAt(t, ss) }) t.Run("Overwrite", func(t *testing.T) { testPostStoreOverwrite(t, ss) }) + t.Run("OverwriteMultiple", func(t *testing.T) { testPostStoreOverwriteMultiple(t, ss) }) t.Run("GetPostsByIds", func(t *testing.T) { testPostStoreGetPostsByIds(t, ss) }) t.Run("GetPostsBatchForIndexing", func(t *testing.T) { testPostStoreGetPostsBatchForIndexing(t, ss) }) t.Run("PermanentDeleteBatch", func(t *testing.T) { testPostStorePermanentDeleteBatch(t, ss) }) @@ -54,16 +56,223 @@ func TestPostStore(t *testing.T, ss store.Store, s SqlSupplier) { } func testPostStoreSave(t *testing.T, ss store.Store) { - o1 := model.Post{} - o1.ChannelId = model.NewId() - o1.UserId = model.NewId() - o1.Message = "zz" + model.NewId() + "b" + t.Run("Save post", func(t *testing.T) { + o1 := model.Post{} + o1.ChannelId = model.NewId() + o1.UserId = model.NewId() + o1.Message = "zz" + model.NewId() + "b" - _, err := ss.Post().Save(&o1) - require.Nil(t, err, "couldn't save item") + _, err := ss.Post().Save(&o1) + require.Nil(t, err, "couldn't save item") + }) - _, err = ss.Post().Save(&o1) - require.NotNil(t, err, "shouldn't be able to update from save") + t.Run("Try to save existing post", func(t *testing.T) { + o1 := model.Post{} + o1.ChannelId = model.NewId() + o1.UserId = model.NewId() + o1.Message = "zz" + model.NewId() + "b" + + _, err := ss.Post().Save(&o1) + require.Nil(t, err, "couldn't save item") + + _, err = ss.Post().Save(&o1) + require.NotNil(t, err, "shouldn't be able to update from save") + }) + + t.Run("Update reply should update the UpdateAt of the root post", func(t *testing.T) { + rootPost := model.Post{} + rootPost.ChannelId = model.NewId() + rootPost.UserId = model.NewId() + rootPost.Message = "zz" + model.NewId() + "b" + + _, err := ss.Post().Save(&rootPost) + require.Nil(t, err) + + replyPost := model.Post{} + replyPost.ChannelId = rootPost.ChannelId + replyPost.UserId = model.NewId() + replyPost.Message = "zz" + model.NewId() + "b" + replyPost.RootId = rootPost.Id + + _, err = ss.Post().Save(&replyPost) + require.Nil(t, err) + + rrootPost, err := ss.Post().GetSingle(rootPost.Id) + require.Nil(t, err) + assert.Greater(t, rrootPost.UpdateAt, rootPost.UpdateAt) + }) + + t.Run("Create a post should update the channel LastPostAt and the total messages count by one", func(t *testing.T) { + channel := model.Channel{} + channel.Name = "zz" + model.NewId() + "b" + channel.DisplayName = "zz" + model.NewId() + "b" + channel.Type = model.CHANNEL_OPEN + + _, err := ss.Channel().Save(&channel, 100) + require.Nil(t, err) + + post := model.Post{} + post.ChannelId = channel.Id + post.UserId = model.NewId() + post.Message = "zz" + model.NewId() + "b" + + _, err = ss.Post().Save(&post) + require.Nil(t, err) + + rchannel, err := ss.Channel().Get(channel.Id, false) + require.Nil(t, err) + assert.Greater(t, rchannel.LastPostAt, channel.LastPostAt) + assert.Equal(t, int64(1), rchannel.TotalMsgCount) + + post = model.Post{} + post.ChannelId = channel.Id + post.UserId = model.NewId() + post.Message = "zz" + model.NewId() + "b" + post.CreateAt = 5 + + _, err = ss.Post().Save(&post) + require.Nil(t, err) + + rchannel2, err := ss.Channel().Get(channel.Id, false) + require.Nil(t, err) + assert.Equal(t, rchannel.LastPostAt, rchannel2.LastPostAt) + assert.Equal(t, int64(2), rchannel2.TotalMsgCount) + + post = model.Post{} + post.ChannelId = channel.Id + post.UserId = model.NewId() + post.Message = "zz" + model.NewId() + "b" + + _, err = ss.Post().Save(&post) + require.Nil(t, err) + + rchannel3, err := ss.Channel().Get(channel.Id, false) + require.Nil(t, err) + assert.Greater(t, rchannel3.LastPostAt, rchannel2.LastPostAt) + assert.Equal(t, int64(3), rchannel3.TotalMsgCount) + }) +} + +func testPostStoreSaveMultiple(t *testing.T, ss store.Store) { + p1 := model.Post{} + p1.ChannelId = model.NewId() + p1.UserId = model.NewId() + p1.Message = "zz" + model.NewId() + "b" + + p2 := model.Post{} + p2.ChannelId = model.NewId() + p2.UserId = model.NewId() + p2.Message = "zz" + model.NewId() + "b" + + p3 := model.Post{} + p3.ChannelId = model.NewId() + p3.UserId = model.NewId() + p3.Message = "zz" + model.NewId() + "b" + + p4 := model.Post{} + p4.ChannelId = model.NewId() + p4.UserId = model.NewId() + p4.Message = "zz" + model.NewId() + "b" + + t.Run("Save correctly a new set of posts", func(t *testing.T) { + newPosts, err := ss.Post().SaveMultiple([]*model.Post{&p1, &p2, &p3}) + require.Nil(t, err) + for _, post := range newPosts { + storedPost, err := ss.Post().GetSingle(post.Id) + assert.Nil(t, err) + assert.Equal(t, post.ChannelId, storedPost.ChannelId) + assert.Equal(t, post.Message, storedPost.Message) + assert.Equal(t, post.UserId, storedPost.UserId) + } + }) + + t.Run("Try to save mixed, already saved and not saved posts", func(t *testing.T) { + newPosts, err := ss.Post().SaveMultiple([]*model.Post{&p4, &p3}) + require.NotNil(t, err) + require.Nil(t, newPosts) + storedPost, err := ss.Post().GetSingle(p3.Id) + assert.Nil(t, err) + assert.Equal(t, p3.ChannelId, storedPost.ChannelId) + assert.Equal(t, p3.Message, storedPost.Message) + assert.Equal(t, p3.UserId, storedPost.UserId) + + storedPost, err = ss.Post().GetSingle(p4.Id) + assert.NotNil(t, err) + assert.Nil(t, storedPost) + }) + + t.Run("Update reply should update the UpdateAt of the root post", func(t *testing.T) { + rootPost := model.Post{} + rootPost.ChannelId = model.NewId() + rootPost.UserId = model.NewId() + rootPost.Message = "zz" + model.NewId() + "b" + + replyPost := model.Post{} + replyPost.ChannelId = rootPost.ChannelId + replyPost.UserId = model.NewId() + replyPost.Message = "zz" + model.NewId() + "b" + replyPost.RootId = rootPost.Id + + _, err := ss.Post().SaveMultiple([]*model.Post{&rootPost, &replyPost}) + require.Nil(t, err) + + rrootPost, err := ss.Post().GetSingle(rootPost.Id) + require.Nil(t, err) + assert.Equal(t, rrootPost.UpdateAt, rootPost.UpdateAt) + + replyPost2 := model.Post{} + replyPost2.ChannelId = rootPost.ChannelId + replyPost2.UserId = model.NewId() + replyPost2.Message = "zz" + model.NewId() + "b" + replyPost2.RootId = rootPost.Id + + replyPost3 := model.Post{} + replyPost3.ChannelId = rootPost.ChannelId + replyPost3.UserId = model.NewId() + replyPost3.Message = "zz" + model.NewId() + "b" + replyPost3.RootId = rootPost.Id + + _, err = ss.Post().SaveMultiple([]*model.Post{&replyPost2, &replyPost3}) + require.Nil(t, err) + + rrootPost2, err := ss.Post().GetSingle(rootPost.Id) + require.Nil(t, err) + assert.Greater(t, rrootPost2.UpdateAt, rrootPost.UpdateAt) + }) + + t.Run("Create a post should update the channel LastPostAt and the total messages count by one", func(t *testing.T) { + channel := model.Channel{} + channel.Name = "zz" + model.NewId() + "b" + channel.DisplayName = "zz" + model.NewId() + "b" + channel.Type = model.CHANNEL_OPEN + + _, err := ss.Channel().Save(&channel, 100) + require.Nil(t, err) + + post1 := model.Post{} + post1.ChannelId = channel.Id + post1.UserId = model.NewId() + post1.Message = "zz" + model.NewId() + "b" + + post2 := model.Post{} + post2.ChannelId = channel.Id + post2.UserId = model.NewId() + post2.Message = "zz" + model.NewId() + "b" + post2.CreateAt = 5 + + post3 := model.Post{} + post3.ChannelId = channel.Id + post3.UserId = model.NewId() + post3.Message = "zz" + model.NewId() + "b" + + _, err = ss.Post().SaveMultiple([]*model.Post{&post1, &post2, &post3}) + require.Nil(t, err) + + rchannel, err := ss.Channel().Get(channel.Id, false) + require.Nil(t, err) + assert.Greater(t, rchannel.LastPostAt, channel.LastPostAt) + assert.Equal(t, int64(3), rchannel.TotalMsgCount) + }) } func testPostStoreSaveChannelMsgCounts(t *testing.T, ss store.Store) { @@ -2005,6 +2214,136 @@ func testPostStoreGetPostsCreatedAt(t *testing.T, ss store.Store) { assert.Equal(t, 2, len(r1)) } +func testPostStoreOverwriteMultiple(t *testing.T, ss store.Store) { + o1 := &model.Post{} + o1.ChannelId = model.NewId() + o1.UserId = model.NewId() + o1.Message = "zz" + model.NewId() + "AAAAAAAAAAA" + o1, err := ss.Post().Save(o1) + require.Nil(t, err) + + o2 := &model.Post{} + o2.ChannelId = o1.ChannelId + o2.UserId = model.NewId() + o2.Message = "zz" + model.NewId() + "CCCCCCCCC" + o2.ParentId = o1.Id + o2.RootId = o1.Id + o2, err = ss.Post().Save(o2) + require.Nil(t, err) + + o3 := &model.Post{} + o3.ChannelId = o1.ChannelId + o3.UserId = model.NewId() + o3.Message = "zz" + model.NewId() + "QQQQQQQQQQ" + o3, err = ss.Post().Save(o3) + require.Nil(t, err) + + o4, err := ss.Post().Save(&model.Post{ + ChannelId: model.NewId(), + UserId: model.NewId(), + Message: model.NewId(), + Filenames: []string{"test"}, + }) + require.Nil(t, err) + + o5, err := ss.Post().Save(&model.Post{ + ChannelId: model.NewId(), + UserId: model.NewId(), + Message: model.NewId(), + Filenames: []string{"test2", "test3"}, + }) + require.Nil(t, err) + + r1, err := ss.Post().Get(o1.Id, false) + require.Nil(t, err) + ro1 := r1.Posts[o1.Id] + + r2, err := ss.Post().Get(o2.Id, false) + require.Nil(t, err) + ro2 := r2.Posts[o2.Id] + + r3, err := ss.Post().Get(o3.Id, false) + require.Nil(t, err) + ro3 := r3.Posts[o3.Id] + + r4, err := ss.Post().Get(o4.Id, false) + require.Nil(t, err) + ro4 := r4.Posts[o4.Id] + + r5, err := ss.Post().Get(o5.Id, false) + require.Nil(t, err) + ro5 := r5.Posts[o5.Id] + + require.Equal(t, ro1.Message, o1.Message, "Failed to save/get") + require.Equal(t, ro2.Message, o2.Message, "Failed to save/get") + require.Equal(t, ro3.Message, o3.Message, "Failed to save/get") + require.Equal(t, ro4.Message, o4.Message, "Failed to save/get") + require.Equal(t, ro4.Filenames, o4.Filenames, "Failed to save/get") + require.Equal(t, ro5.Message, o5.Message, "Failed to save/get") + require.Equal(t, ro5.Filenames, o5.Filenames, "Failed to save/get") + + t.Run("overwrite changing message", func(t *testing.T) { + o1a := &model.Post{} + *o1a = *ro1 + o1a.Message = ro1.Message + "BBBBBBBBBB" + + o2a := &model.Post{} + *o2a = *ro2 + o2a.Message = ro2.Message + "DDDDDDD" + + o3a := &model.Post{} + *o3a = *ro3 + o3a.Message = ro3.Message + "WWWWWWW" + + _, err = ss.Post().OverwriteMultiple([]*model.Post{o1a, o2a, o3a}) + require.Nil(t, err) + + r1, err = ss.Post().Get(o1.Id, false) + require.Nil(t, err) + ro1a := r1.Posts[o1.Id] + + r2, err = ss.Post().Get(o1.Id, false) + require.Nil(t, err) + ro2a := r2.Posts[o2.Id] + + r3, err = ss.Post().Get(o3.Id, false) + require.Nil(t, err) + ro3a := r3.Posts[o3.Id] + + assert.Equal(t, ro1a.Message, o1a.Message, "Failed to overwrite/get") + assert.Equal(t, ro2a.Message, o2a.Message, "Failed to overwrite/get") + assert.Equal(t, ro3a.Message, o3a.Message, "Failed to overwrite/get") + }) + + t.Run("overwrite clearing filenames", func(t *testing.T) { + o4a := &model.Post{} + *o4a = *ro4 + o4a.Filenames = []string{} + o4a.FileIds = []string{model.NewId()} + + o5a := &model.Post{} + *o5a = *ro5 + o5a.Filenames = []string{} + o5a.FileIds = []string{} + + _, err = ss.Post().OverwriteMultiple([]*model.Post{o4a, o5a}) + require.Nil(t, err) + + r4, err = ss.Post().Get(o4.Id, false) + require.Nil(t, err) + ro4a := r4.Posts[o4.Id] + + r5, err = ss.Post().Get(o5.Id, false) + require.Nil(t, err) + ro5a := r5.Posts[o5.Id] + + require.Empty(t, ro4a.Filenames, "Failed to clear Filenames") + require.Len(t, ro4a.FileIds, 1, "Failed to set FileIds") + require.Empty(t, ro5a.Filenames, "Failed to clear Filenames") + require.Empty(t, ro5a.FileIds, "Failed to set FileIds") + }) +} + func testPostStoreOverwrite(t *testing.T, ss store.Store) { o1 := &model.Post{} o1.ChannelId = model.NewId() @@ -2029,56 +2368,6 @@ func testPostStoreOverwrite(t *testing.T, ss store.Store) { o3, err = ss.Post().Save(o3) require.Nil(t, err) - r1, err := ss.Post().Get(o1.Id, false) - require.Nil(t, err) - ro1 := r1.Posts[o1.Id] - - r2, err := ss.Post().Get(o1.Id, false) - require.Nil(t, err) - ro2 := r2.Posts[o2.Id] - - r3, err := ss.Post().Get(o3.Id, false) - require.Nil(t, err) - ro3 := r3.Posts[o3.Id] - - require.Equal(t, ro1.Message, o1.Message, "Failed to save/get") - - o1a := &model.Post{} - *o1a = *ro1 - o1a.Message = ro1.Message + "BBBBBBBBBB" - _, err = ss.Post().Overwrite(o1a) - require.Nil(t, err) - - r1, err = ss.Post().Get(o1.Id, false) - require.Nil(t, err) - ro1a := r1.Posts[o1.Id] - - require.Equal(t, ro1a.Message, o1a.Message, "Failed to overwrite/get") - - o2a := &model.Post{} - *o2a = *ro2 - o2a.Message = ro2.Message + "DDDDDDD" - _, err = ss.Post().Overwrite(o2a) - require.Nil(t, err) - - r2, err = ss.Post().Get(o1.Id, false) - require.Nil(t, err) - ro2a := r2.Posts[o2.Id] - - require.Equal(t, ro2a.Message, o2a.Message, "Failed to overwrite/get") - - o3a := &model.Post{} - *o3a = *ro3 - o3a.Message = ro3.Message + "WWWWWWW" - _, err = ss.Post().Overwrite(o3a) - require.Nil(t, err) - - r3, err = ss.Post().Get(o3.Id, false) - require.Nil(t, err) - ro3a := r3.Posts[o3.Id] - - require.Equal(t, ro3a.Message, o3a.Message, "Failed to overwrite/get") - o4, err := ss.Post().Save(&model.Post{ ChannelId: model.NewId(), UserId: model.NewId(), @@ -2087,23 +2376,78 @@ func testPostStoreOverwrite(t *testing.T, ss store.Store) { }) require.Nil(t, err) + r1, err := ss.Post().Get(o1.Id, false) + require.Nil(t, err) + ro1 := r1.Posts[o1.Id] + + r2, err := ss.Post().Get(o2.Id, false) + require.Nil(t, err) + ro2 := r2.Posts[o2.Id] + + r3, err := ss.Post().Get(o3.Id, false) + require.Nil(t, err) + ro3 := r3.Posts[o3.Id] + r4, err := ss.Post().Get(o4.Id, false) require.Nil(t, err) ro4 := r4.Posts[o4.Id] - o4a := &model.Post{} - *o4a = *ro4 - o4a.Filenames = []string{} - o4a.FileIds = []string{model.NewId()} - _, err = ss.Post().Overwrite(o4a) - require.Nil(t, err) + require.Equal(t, ro1.Message, o1.Message, "Failed to save/get") + require.Equal(t, ro2.Message, o2.Message, "Failed to save/get") + require.Equal(t, ro3.Message, o3.Message, "Failed to save/get") + require.Equal(t, ro4.Message, o4.Message, "Failed to save/get") - r4, err = ss.Post().Get(o4.Id, false) - require.Nil(t, err) + t.Run("overwrite changing message", func(t *testing.T) { + o1a := &model.Post{} + *o1a = *ro1 + o1a.Message = ro1.Message + "BBBBBBBBBB" + _, err = ss.Post().Overwrite(o1a) + require.Nil(t, err) - ro4a := r4.Posts[o4.Id] - require.Empty(t, ro4a.Filenames, "Failed to clear Filenames") - require.Len(t, ro4a.FileIds, 1, "Failed to set FileIds") + o2a := &model.Post{} + *o2a = *ro2 + o2a.Message = ro2.Message + "DDDDDDD" + _, err = ss.Post().Overwrite(o2a) + require.Nil(t, err) + + o3a := &model.Post{} + *o3a = *ro3 + o3a.Message = ro3.Message + "WWWWWWW" + _, err = ss.Post().Overwrite(o3a) + require.Nil(t, err) + + r1, err = ss.Post().Get(o1.Id, false) + require.Nil(t, err) + ro1a := r1.Posts[o1.Id] + + r2, err = ss.Post().Get(o1.Id, false) + require.Nil(t, err) + ro2a := r2.Posts[o2.Id] + + r3, err = ss.Post().Get(o3.Id, false) + require.Nil(t, err) + ro3a := r3.Posts[o3.Id] + + assert.Equal(t, ro1a.Message, o1a.Message, "Failed to overwrite/get") + assert.Equal(t, ro2a.Message, o2a.Message, "Failed to overwrite/get") + assert.Equal(t, ro3a.Message, o3a.Message, "Failed to overwrite/get") + }) + + t.Run("overwrite clearing filenames", func(t *testing.T) { + o4a := &model.Post{} + *o4a = *ro4 + o4a.Filenames = []string{} + o4a.FileIds = []string{model.NewId()} + _, err = ss.Post().Overwrite(o4a) + require.Nil(t, err) + + r4, err = ss.Post().Get(o4.Id, false) + require.Nil(t, err) + + ro4a := r4.Posts[o4.Id] + require.Empty(t, ro4a.Filenames, "Failed to clear Filenames") + require.Len(t, ro4a.FileIds, 1, "Failed to set FileIds") + }) } func testPostStoreGetPostsByIds(t *testing.T, ss store.Store) { diff --git a/store/storetest/team_store.go b/store/storetest/team_store.go index 662adc0321..c269fdfe17 100644 --- a/store/storetest/team_store.go +++ b/store/storetest/team_store.go @@ -30,6 +30,7 @@ func TestTeamStore(t *testing.T, ss store.Store) { t.Run("Update", func(t *testing.T) { testTeamStoreUpdate(t, ss) }) t.Run("Get", func(t *testing.T) { testTeamStoreGet(t, ss) }) t.Run("GetByName", func(t *testing.T) { testTeamStoreGetByName(t, ss) }) + t.Run("GetByNames", func(t *testing.T) { testTeamStoreGetByNames(t, ss) }) t.Run("SearchAll", func(t *testing.T) { testTeamStoreSearchAll(t, ss) }) t.Run("SearchOpen", func(t *testing.T) { testTeamStoreSearchOpen(t, ss) }) t.Run("SearchPrivate", func(t *testing.T) { testTeamStoreSearchPrivate(t, ss) }) @@ -122,6 +123,59 @@ func testTeamStoreGet(t *testing.T, ss store.Store) { require.NotNil(t, err, "Missing id should have failed") } +func testTeamStoreGetByNames(t *testing.T, ss store.Store) { + o1 := model.Team{} + o1.DisplayName = "DisplayName" + o1.Name = "z-z-z" + model.NewId() + "b" + o1.Email = MakeEmail() + o1.Type = model.TEAM_OPEN + + _, err := ss.Team().Save(&o1) + require.Nil(t, err) + + o2 := model.Team{} + o2.DisplayName = "DisplayName2" + o2.Name = "z-z-z" + model.NewId() + "b" + o2.Email = MakeEmail() + o2.Type = model.TEAM_OPEN + + _, err = ss.Team().Save(&o2) + require.Nil(t, err) + + t.Run("Get empty list", func(t *testing.T) { + var teams []*model.Team + teams, err = ss.Team().GetByNames([]string{}) + require.Nil(t, err) + require.Empty(t, teams) + }) + + t.Run("Get existing teams", func(t *testing.T) { + var teams []*model.Team + teams, err = ss.Team().GetByNames([]string{o1.Name, o2.Name}) + require.Nil(t, err) + teamsIds := []string{} + for _, team := range teams { + teamsIds = append(teamsIds, team.Id) + } + assert.Contains(t, teamsIds, o1.Id, "invalid returned team") + assert.Contains(t, teamsIds, o2.Id, "invalid returned team") + }) + + t.Run("Get existing team and one invalid team name", func(t *testing.T) { + _, err = ss.Team().GetByNames([]string{o1.Name, ""}) + require.NotNil(t, err) + }) + + t.Run("Get existing team and not existing team", func(t *testing.T) { + _, err = ss.Team().GetByNames([]string{o1.Name, "not-existing-team-name"}) + require.NotNil(t, err) + }) + t.Run("Get not existing teams", func(t *testing.T) { + _, err = ss.Team().GetByNames([]string{"not-existing-team-name", "not-existing-team-name-2"}) + require.NotNil(t, err) + }) +} + func testTeamStoreGetByName(t *testing.T, ss store.Store) { o1 := model.Team{} o1.DisplayName = "DisplayName" @@ -132,12 +186,22 @@ func testTeamStoreGetByName(t *testing.T, ss store.Store) { _, err := ss.Team().Save(&o1) require.Nil(t, err) - team, err := ss.Team().GetByName(o1.Name) - require.Nil(t, err) - require.Equal(t, *team, o1, "invalid returned team") + t.Run("Get existing team", func(t *testing.T) { + var team *model.Team + team, err = ss.Team().GetByName(o1.Name) + require.Nil(t, err) + require.Equal(t, *team, o1, "invalid returned team") + }) - _, err = ss.Team().GetByName("") - require.NotNil(t, err, "Missing id should have failed") + t.Run("Get invalid team name", func(t *testing.T) { + _, err = ss.Team().GetByName("") + require.NotNil(t, err, "Missing id should have failed") + }) + + t.Run("Get not existing team", func(t *testing.T) { + _, err = ss.Team().GetByName("not-existing-team-name") + require.NotNil(t, err, "Missing id should have failed") + }) } func testTeamStoreSearchAll(t *testing.T, ss store.Store) { diff --git a/store/timer_layer.go b/store/timer_layer.go index 620f9479a5..cda63f05c0 100644 --- a/store/timer_layer.go +++ b/store/timer_layer.go @@ -4412,6 +4412,22 @@ func (s *TimerLayerPostStore) Save(post *model.Post) (*model.Post, *model.AppErr return resultVar0, resultVar1 } +func (s *TimerLayerPostStore) SaveMultiple(posts []*model.Post) ([]*model.Post, *model.AppError) { + start := timemodule.Now() + + resultVar0, resultVar1 := s.PostStore.SaveMultiple(posts) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if resultVar1 == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("PostStore.SaveMultiple", success, elapsed) + } + return resultVar0, resultVar1 +} + func (s *TimerLayerPostStore) Search(teamId string, userId string, params *model.SearchParams) (*model.PostList, *model.AppError) { start := timemodule.Now()