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