MM-35396 Refactor GetThreadForUser store func to take membership as argument to prevent replica lag issues and reduce joins in query (#17754)

* Update store function GetThreadForUser to use master DB to fix replica lag

* Refactor GetThreadForUser store func to take membership as argument to prevent replica lag issues and reduce joins in query

* Add translation

* Fix test

* Updates per feedback

* Minor clean-up per feedback
Этот коммит содержится в:
Joram Wilander
2021-06-14 12:33:08 -04:00
коммит произвёл GitHub
родитель 24fb0033f4
Коммит d0778486ad
16 изменённых файлов: 119 добавлений и 46 удалений

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

@@ -8866,7 +8866,7 @@ func (s *OpenTracingLayerThreadStore) GetThreadFollowers(threadID string) ([]str
return result, err
}
func (s *OpenTracingLayerThreadStore) GetThreadForUser(userID string, teamID string, threadId string, extended bool) (*model.ThreadResponse, error) {
func (s *OpenTracingLayerThreadStore) GetThreadForUser(teamID string, threadMembership *model.ThreadMembership, extended bool) (*model.ThreadResponse, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.GetThreadForUser")
s.Root.Store.SetContext(newCtx)
@@ -8875,7 +8875,7 @@ func (s *OpenTracingLayerThreadStore) GetThreadForUser(userID string, teamID str
}()
defer span.Finish()
result, err := s.ThreadStore.GetThreadForUser(userID, teamID, threadId, extended)
result, err := s.ThreadStore.GetThreadForUser(teamID, threadMembership, extended)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)

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

@@ -9648,11 +9648,11 @@ func (s *RetryLayerThreadStore) GetThreadFollowers(threadID string) ([]string, e
}
func (s *RetryLayerThreadStore) GetThreadForUser(userID string, teamID string, threadId string, extended bool) (*model.ThreadResponse, error) {
func (s *RetryLayerThreadStore) GetThreadForUser(teamID string, threadMembership *model.ThreadMembership, extended bool) (*model.ThreadResponse, error) {
tries := 0
for {
result, err := s.ThreadStore.GetThreadForUser(userID, teamID, threadId, extended)
result, err := s.ThreadStore.GetThreadForUser(teamID, threadMembership, extended)
if err == nil {
return result, nil
}

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

@@ -321,7 +321,11 @@ func (s *SqlThreadStore) GetThreadFollowers(threadID string) ([]string, error) {
return users, nil
}
func (s *SqlThreadStore) GetThreadForUser(userId, teamId, threadId string, extended bool) (*model.ThreadResponse, error) {
func (s *SqlThreadStore) GetThreadForUser(teamId string, threadMembership *model.ThreadMembership, extended bool) (*model.ThreadResponse, error) {
if !threadMembership.Following {
return nil, nil // in case the thread is not followed anymore - return nil error to be interpreted as 404
}
type JoinedThread struct {
PostId string
Following bool
@@ -334,38 +338,45 @@ func (s *SqlThreadStore) GetThreadForUser(userId, teamId, threadId string, exten
model.Post
}
unreadRepliesQuery := "SELECT COUNT(Posts.Id) From Posts Where Posts.RootId=ThreadMemberships.PostId AND Posts.CreateAt >= ThreadMemberships.LastViewed AND Posts.DeleteAt=0"
unreadRepliesQuery, unreadRepliesArgs := sq.
Select("COUNT(Posts.Id)").
From("Posts").
Where(sq.And{
sq.Eq{"Posts.RootId": threadMembership.PostId},
sq.GtOrEq{"Posts.CreateAt": threadMembership.LastViewed},
sq.Eq{"Posts.DeleteAt": 0},
}).MustSql()
fetchConditions := sq.And{
sq.Or{sq.Eq{"Channels.TeamId": teamId}, sq.Eq{"Channels.TeamId": ""}},
sq.Eq{"ThreadMemberships.UserId": userId},
sq.Eq{"Threads.PostId": threadId},
sq.Eq{"Threads.PostId": threadMembership.PostId},
}
var thread JoinedThread
query, args, _ := s.getQueryBuilder().
Select("Threads.*, Posts.*, ThreadMemberships.LastViewed as LastViewedAt, ThreadMemberships.UnreadMentions as UnreadMentions, ThreadMemberships.Following").
query, threadArgs, _ := s.getQueryBuilder().
Select("Threads.*, Posts.*").
From("Threads").
Column(sq.Alias(sq.Expr(unreadRepliesQuery), "UnreadReplies")).
LeftJoin("Posts ON Posts.Id = Threads.PostId").
LeftJoin("Channels ON Posts.ChannelId = Channels.Id").
LeftJoin("ThreadMemberships ON ThreadMemberships.PostId = Threads.PostId").
Where(fetchConditions).ToSql()
err := s.GetReplica().SelectOne(&thread, query, args...)
args := append(unreadRepliesArgs, threadArgs...)
err := s.GetReplica().SelectOne(&thread, query, args...)
if err != nil {
return nil, err
}
if !thread.Following {
return nil, nil // in case the thread is not followed anymore - return nil error to be interpreted as 404
}
thread.LastViewedAt = threadMembership.LastViewed
thread.UnreadMentions = threadMembership.UnreadMentions
var users []*model.User
if extended {
var err error
users, err = s.User().GetProfileByIds(context.Background(), thread.Participants, &store.UserGetByIdsOpts{}, true)
if err != nil {
return nil, errors.Wrapf(err, "failed to get threads for user id=%s", userId)
return nil, errors.Wrapf(err, "failed to get thread for user id=%s", threadMembership.UserId)
}
} else {
for _, userId := range thread.Participants {

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

@@ -283,7 +283,7 @@ type ThreadStore interface {
Update(thread *model.Thread) (*model.Thread, error)
Get(id string) (*model.Thread, error)
GetThreadsForUser(userId, teamID string, opts model.GetUserThreadsOpts) (*model.Threads, error)
GetThreadForUser(userID, teamID, threadId string, extended bool) (*model.ThreadResponse, error)
GetThreadForUser(teamID string, threadMembership *model.ThreadMembership, extended bool) (*model.ThreadResponse, error)
Delete(postID string) error
GetPosts(threadID string, since int64) ([]*model.Post, error)

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

@@ -180,13 +180,13 @@ func (_m *ThreadStore) GetThreadFollowers(threadID string) ([]string, error) {
return r0, r1
}
// GetThreadForUser provides a mock function with given fields: userID, teamID, threadId, extended
func (_m *ThreadStore) GetThreadForUser(userID string, teamID string, threadId string, extended bool) (*model.ThreadResponse, error) {
ret := _m.Called(userID, teamID, threadId, extended)
// GetThreadForUser provides a mock function with given fields: teamID, threadMembership, extended
func (_m *ThreadStore) GetThreadForUser(teamID string, threadMembership *model.ThreadMembership, extended bool) (*model.ThreadResponse, error) {
ret := _m.Called(teamID, threadMembership, extended)
var r0 *model.ThreadResponse
if rf, ok := ret.Get(0).(func(string, string, string, bool) *model.ThreadResponse); ok {
r0 = rf(userID, teamID, threadId, extended)
if rf, ok := ret.Get(0).(func(string, *model.ThreadMembership, bool) *model.ThreadResponse); ok {
r0 = rf(teamID, threadMembership, extended)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.ThreadResponse)
@@ -194,8 +194,8 @@ func (_m *ThreadStore) GetThreadForUser(userID string, teamID string, threadId s
}
var r1 error
if rf, ok := ret.Get(1).(func(string, string, string, bool) error); ok {
r1 = rf(userID, teamID, threadId, extended)
if rf, ok := ret.Get(1).(func(string, *model.ThreadMembership, bool) error); ok {
r1 = rf(teamID, threadMembership, extended)
} else {
r1 = ret.Error(1)
}

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

@@ -339,14 +339,14 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) {
newPosts := makeSomePosts()
m, err := ss.Thread().MaintainMembership(newPosts[0].UserId, newPosts[0].Id, true, false, true, false, false)
require.NoError(t, err)
th, err := ss.Thread().GetThreadForUser(newPosts[0].UserId, "", newPosts[0].Id, false)
th, err := ss.Thread().GetThreadForUser("", m, false)
require.NoError(t, err)
require.Equal(t, int64(2), th.UnreadReplies)
m.LastViewed = newPosts[2].UpdateAt + 1
_, err = ss.Thread().UpdateMembership(m)
require.NoError(t, err)
th, err = ss.Thread().GetThreadForUser(newPosts[0].UserId, "", newPosts[0].Id, false)
th, err = ss.Thread().GetThreadForUser("", m, false)
require.NoError(t, err)
require.Equal(t, int64(0), th.UnreadReplies)
@@ -355,7 +355,7 @@ func testThreadStorePopulation(t *testing.T, ss store.Store) {
_, err = ss.Post().Update(editedPost, newPosts[2])
require.NoError(t, err)
th, err = ss.Thread().GetThreadForUser(newPosts[0].UserId, "", newPosts[0].Id, false)
th, err = ss.Thread().GetThreadForUser("", m, false)
require.NoError(t, err)
require.Equal(t, int64(0), th.UnreadReplies)
})

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

@@ -7990,10 +7990,10 @@ func (s *TimerLayerThreadStore) GetThreadFollowers(threadID string) ([]string, e
return result, err
}
func (s *TimerLayerThreadStore) GetThreadForUser(userID string, teamID string, threadId string, extended bool) (*model.ThreadResponse, error) {
func (s *TimerLayerThreadStore) GetThreadForUser(teamID string, threadMembership *model.ThreadMembership, extended bool) (*model.ThreadResponse, error) {
start := timemodule.Now()
result, err := s.ThreadStore.GetThreadForUser(userID, teamID, threadId, extended)
result, err := s.ThreadStore.GetThreadForUser(teamID, threadMembership, extended)
elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second)
if s.Root.Metrics != nil {