[MM-59367] export: enable exporting thread followers for CRT (#27623)

Этот коммит содержится в:
Ibrahim Serdar Acikgoz
2024-08-29 14:06:41 +02:00
коммит произвёл GitHub
родитель 3dc0e63c03
Коммит d5cc2eb2f6
17 изменённых файлов: 1430 добавлений и 63 удалений

Просмотреть файл

@@ -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()