From 5a0a3e6d13e63a03f22facd84446609a15732dce Mon Sep 17 00:00:00 2001 From: Jesse Hallam Date: Wed, 9 Nov 2022 14:00:07 -0400 Subject: [PATCH] Inline ThreadStore.MarkAllAsUnreadByTeam (#20958) * Add teamId to Threads table * Get rid of multiple teamId reads * Fix failed test * Inline ThreadStore.MarkAllAsUnreadByTeam The query to `MarkAllAsUnreadByTeam` first fetched all thread memberships, then fed just the ids back to a second query to ensure all are marked as unread. Optimize this by simply doing a single `UPDATE` query with the necessary joins. Co-authored-by: iomodo --- store/sqlstore/post_store.go | 6 +- store/sqlstore/thread_store.go | 29 ++-- store/storetest/thread_store.go | 228 +++++++++++++++++++++++++++++++- 3 files changed, 245 insertions(+), 18 deletions(-) diff --git a/store/sqlstore/post_store.go b/store/sqlstore/post_store.go index 141565c816..a11a3ee347 100644 --- a/store/sqlstore/post_store.go +++ b/store/sqlstore/post_store.go @@ -1670,7 +1670,7 @@ func (s *SqlPostStore) getParentsPosts(channelId string, offset int, limit int, FROM Posts WHERE - ChannelId = ? ` + deleteAtCondition + ` + ChannelId = ? ` + deleteAtCondition + ` ORDER BY CreateAt DESC LIMIT ? OFFSET ?) q WHERE q.RootId != ''` @@ -1757,13 +1757,13 @@ func (s *SqlPostStore) getParentsPostsPostgreSQL(channelId string, offset int, l FROM Posts WHERE - Posts.ChannelId = ? `+deleteAtSubQueryCondition+` + Posts.ChannelId = ? `+deleteAtSubQueryCondition+` ORDER BY Posts.CreateAt DESC LIMIT ? OFFSET ?) q3 WHERE q3.RootId != '') q1 ON `+onStatement+` WHERE - q2.ChannelId = ? `+deleteAtQueryCondition+` + q2.ChannelId = ? `+deleteAtQueryCondition+` ORDER BY q2.CreateAt`, channelId, limit, offset, channelId) if err != nil { return nil, errors.Wrapf(err, "failed to find Posts with channelId=%s", channelId) diff --git a/store/sqlstore/thread_store.go b/store/sqlstore/thread_store.go index 195e1b535a..9e4e47835f 100644 --- a/store/sqlstore/thread_store.go +++ b/store/sqlstore/thread_store.go @@ -549,27 +549,28 @@ func (s *SqlThreadStore) MarkAllAsRead(userId string, threadIds []string) error // MarkAllAsReadByTeam marks all threads for the given user in the given team as read from the // current time. func (s *SqlThreadStore) MarkAllAsReadByTeam(userId, teamId string) error { - memberships, err := s.GetMembershipsForUser(userId, teamId) - if err != nil { - return err - } - membershipIds := []string{} - for _, m := range memberships { - membershipIds = append(membershipIds, m.PostId) - } timestamp := model.GetMillis() - query := s.getQueryBuilder(). - Update("ThreadMemberships"). - Where(sq.Eq{"PostId": membershipIds}). - Where(sq.Eq{"UserId": userId}). + + var query sq.UpdateBuilder + if s.DriverName() == model.DatabaseDriverPostgres { + query = s.getQueryBuilder().Update("ThreadMemberships").From("Threads") + } else { + query = s.getQueryBuilder().Update("ThreadMemberships", "Threads") + } + + query = query. + Where("Threads.PostId = ThreadMemberships.PostId"). + Where(sq.Eq{"ThreadMemberships.UserId": userId}). + Where(sq.Or{sq.Eq{"Threads.TeamId": teamId}, sq.Eq{"Threads.TeamId": ""}}). Set("LastViewed", timestamp). Set("UnreadMentions", 0). - Set("LastUpdated", model.GetMillis()) + Set("LastUpdated", timestamp) - _, err = s.GetMasterX().ExecBuilder(query) + _, err := s.GetMasterX().ExecBuilder(query) if err != nil { return errors.Wrapf(err, "failed to update thread read state for user id=%s", userId) } + return nil } diff --git a/store/storetest/thread_store.go b/store/storetest/thread_store.go index 463c4396aa..e90e9f7d89 100644 --- a/store/storetest/thread_store.go +++ b/store/storetest/thread_store.go @@ -28,6 +28,7 @@ func TestThreadStore(t *testing.T, ss store.Store, s SqlStore) { t.Run("GetVarious", func(t *testing.T) { testVarious(t, ss) }) t.Run("MarkAllAsReadByChannels", func(t *testing.T) { testMarkAllAsReadByChannels(t, ss) }) t.Run("GetTopThreads", func(t *testing.T) { testGetTopThreads(t, ss) }) + t.Run("MarkAllAsReadByTeam", func(t *testing.T) { testMarkAllAsReadByTeam(t, ss) }) } func testThreadStorePopulation(t *testing.T, ss store.Store) { @@ -1603,5 +1604,230 @@ func testGetTopThreads(t *testing.T, ss store.Store) { // require first element to be post1 with 2 replyCount=2 require.Equal(t, topThreadsInTeamOlder.Items[1].PostId, post2.Id) }) - +} + +func testMarkAllAsReadByTeam(t *testing.T, ss store.Store) { + createThreadMembership := func(userID, postID string) { + t.Helper() + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: false, + UpdateFollowing: true, + UpdateViewedTimestamp: false, + UpdateParticipants: false, + } + _, err := ss.Thread().MaintainMembership(userID, postID, opts) + require.NoError(t, err) + } + + assertThreadReplyCount := func(t *testing.T, userID, teamID string, count int64, message string) { + t.Helper() + + teamsUnread, err := ss.Thread().GetTeamsUnreadForUser(userID, []string{teamID}) + require.NoError(t, err) + require.Lenf(t, teamsUnread, 1, "unexpected unread teams count: %s", message) + assert.Equalf(t, count, teamsUnread[teamID].ThreadCount, "unexpected thread count: %s", message) + } + + postingUserId := model.NewId() + userAID := model.NewId() + userBID := model.NewId() + + team1, err := ss.Team().Save(&model.Team{ + DisplayName: "Team1", + Name: "team1" + model.NewId(), + Email: MakeEmail(), + Type: model.TeamOpen, + }) + require.NoError(t, err) + + team1channel1, err := ss.Channel().Save(&model.Channel{ + TeamId: team1.Id, + DisplayName: "Team1: Channel1", + Name: "team1channel1" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + + team1channel2, err := ss.Channel().Save(&model.Channel{ + TeamId: team1.Id, + DisplayName: "Team1: Channel2", + Name: "team1channel2" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + + team2, err := ss.Team().Save(&model.Team{ + DisplayName: "Team2", + Name: "team2" + model.NewId(), + Email: MakeEmail(), + Type: model.TeamOpen, + }) + require.NoError(t, err) + + team2channel1, err := ss.Channel().Save(&model.Channel{ + TeamId: team2.Id, + DisplayName: "Team2: Channel1", + Name: "team2channel1" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + + team2channel2, err := ss.Channel().Save(&model.Channel{ + TeamId: team2.Id, + DisplayName: "Team2: Channel2", + Name: "team2channel2" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + + team1channel1post1, err := ss.Post().Save(&model.Post{ + ChannelId: team1channel1.Id, + UserId: postingUserId, + Message: "Root", + }) + require.NoError(t, err) + + _, err = ss.Post().Save(&model.Post{ + ChannelId: team1channel1.Id, + UserId: postingUserId, + RootId: team1channel1post1.Id, + Message: "Reply", + }) + require.NoError(t, err) + + team1channel2post1, err := ss.Post().Save(&model.Post{ + ChannelId: team1channel2.Id, + UserId: postingUserId, + Message: "Root", + }) + require.NoError(t, err) + + _, err = ss.Post().Save(&model.Post{ + ChannelId: team1channel1.Id, + UserId: postingUserId, + RootId: team1channel2post1.Id, + Message: "Reply", + }) + require.NoError(t, err) + + team2channel1post1, err := ss.Post().Save(&model.Post{ + ChannelId: team2channel1.Id, + UserId: postingUserId, + Message: "Root", + }) + require.NoError(t, err) + + _, err = ss.Post().Save(&model.Post{ + ChannelId: team2channel1.Id, + UserId: postingUserId, + RootId: team2channel1post1.Id, + Message: "Reply", + }) + require.NoError(t, err) + + team2channel2post1, err := ss.Post().Save(&model.Post{ + ChannelId: team2channel2.Id, + UserId: postingUserId, + Message: "Root", + }) + require.NoError(t, err) + + _, err = ss.Post().Save(&model.Post{ + ChannelId: team2channel1.Id, + UserId: postingUserId, + RootId: team2channel2post1.Id, + Message: "Reply", + }) + require.NoError(t, err) + + gm1, err := ss.Channel().Save(&model.Channel{ + DisplayName: "GM1", + Name: "gm1" + model.NewId(), + Type: model.ChannelTypeGroup, + }, -1) + require.NoError(t, err) + + gm1post1, err := ss.Post().Save(&model.Post{ + ChannelId: gm1.Id, + UserId: postingUserId, + Message: "Root", + }) + require.NoError(t, err) + + _, err = ss.Post().Save(&model.Post{ + ChannelId: gm1.Id, + UserId: postingUserId, + RootId: gm1post1.Id, + Message: "Reply", + }) + require.NoError(t, err) + + gm2, err := ss.Channel().Save(&model.Channel{ + DisplayName: "GM1", + Name: "gm1" + model.NewId(), + Type: model.ChannelTypeGroup, + }, -1) + require.NoError(t, err) + + gm2post1, err := ss.Post().Save(&model.Post{ + ChannelId: gm2.Id, + UserId: postingUserId, + Message: "Root", + }) + require.NoError(t, err) + + _, err = ss.Post().Save(&model.Post{ + ChannelId: gm2.Id, + UserId: postingUserId, + RootId: gm2post1.Id, + Message: "Reply", + }) + require.NoError(t, err) + + t.Run("empty team", func(t *testing.T) { + err = ss.Thread().MarkAllAsReadByTeam(model.NewId(), "") + require.NoError(t, err) + }) + + t.Run("unknown team", func(t *testing.T) { + err = ss.Thread().MarkAllAsReadByTeam(model.NewId(), model.NewId()) + require.NoError(t, err) + }) + + t.Run("team1", func(t *testing.T) { + createThreadMembership(userAID, team1channel1post1.Id) + createThreadMembership(userBID, team1channel1post1.Id) + createThreadMembership(userAID, team1channel2post1.Id) + createThreadMembership(userBID, team1channel2post1.Id) + createThreadMembership(userAID, team2channel1post1.Id) + createThreadMembership(userBID, team2channel1post1.Id) + + // Note that GMs (and similarly, DMs) don't count towards this API. + createThreadMembership(userAID, gm1.Id) + createThreadMembership(userBID, gm1.Id) + createThreadMembership(userAID, gm2.Id) + createThreadMembership(userBID, gm2.Id) + + assertThreadReplyCount(t, userAID, team1.Id, 2, "expected 2 unread messages in team1 for userA") + assertThreadReplyCount(t, userBID, team1.Id, 2, "expected 2 unread messages in team1 for userB") + assertThreadReplyCount(t, userAID, team2.Id, 1, "expected 1 unread message in team2 for userA") + assertThreadReplyCount(t, userBID, team2.Id, 1, "expected 1 unread message in team2 for userB") + + err = ss.Thread().MarkAllAsReadByTeam(userAID, team1.Id) + require.NoError(t, err) + + assertThreadReplyCount(t, userAID, team1.Id, 0, "expected 0 unread messages in team1 for userA") + assertThreadReplyCount(t, userBID, team1.Id, 2, "expected 2 unread messages in team1 for userB") + assertThreadReplyCount(t, userAID, team2.Id, 1, "expected 1 unread message in team2 for userA") + assertThreadReplyCount(t, userBID, team2.Id, 1, "expected 1 unread message in team2 for userB") + + err = ss.Thread().MarkAllAsReadByTeam(userBID, team1.Id) + require.NoError(t, err) + + assertThreadReplyCount(t, userAID, team1.Id, 0, "expected 0 unread messages in team1 for userA") + assertThreadReplyCount(t, userBID, team1.Id, 0, "expected 0 unread messages in team1 for userB") + assertThreadReplyCount(t, userAID, team2.Id, 1, "expected 1 unread message in team2 for userA") + assertThreadReplyCount(t, userBID, team2.Id, 1, "expected 1 unread message in team2 for userB") + }) }