diff --git a/app/team.go b/app/team.go index 292dcc0450..34966f0426 100644 --- a/app/team.go +++ b/app/team.go @@ -1565,11 +1565,13 @@ func (a *App) GetTeamsUnreadForUser(excludeTeamId string, userID string, include return tu } + teamIDs := make([]string, 0, len(data)) for i := range data { id := data[i].TeamId if mu, ok := membersMap[id]; ok { membersMap[id] = unreads(data[i], mu) } else { + teamIDs = append(teamIDs, id) membersMap[id] = unreads(data[i], &model.TeamUnread{ MsgCount: 0, MentionCount: 0, @@ -1584,18 +1586,22 @@ func (a *App) GetTeamsUnreadForUser(excludeTeamId string, userID string, include includeCollapsedThreads = includeCollapsedThreads && *a.Config().ServiceSettings.CollapsedThreads != model.CollapsedThreadsDisabled - for _, member := range membersMap { - if includeCollapsedThreads { - data, err := a.Srv().Store.Thread().GetThreadsForUser(userID, member.TeamId, model.GetUserThreadsOpts{TotalsOnly: true, TeamOnly: true}) - if err != nil { - return nil, model.NewAppError("GetTeamsUnreadForUser", "app.team.get_unread.app_error", nil, err.Error(), http.StatusInternalServerError) - } - member.ThreadCount = data.TotalUnreadThreads - member.ThreadMentionCount = data.TotalUnreadMentions + if includeCollapsedThreads { + teamUnreads, err := a.Srv().Store.Thread().GetTeamsUnreadForUser(userID, teamIDs) + if err != nil { + return nil, model.NewAppError("GetTeamsUnreadForUser", "app.team.get_unread.app_error", nil, err.Error(), http.StatusInternalServerError) + } + for teamID, member := range membersMap { + if _, ok := teamUnreads[teamID]; ok { + member.ThreadCount = teamUnreads[teamID].ThreadCount + member.ThreadMentionCount = teamUnreads[teamID].ThreadMentionCount + } } - members = append(members, member) } + for _, member := range membersMap { + members = append(members, member) + } return members, nil } diff --git a/store/opentracinglayer/opentracinglayer.go b/store/opentracinglayer/opentracinglayer.go index 52a21589cd..608c7d8aa2 100644 --- a/store/opentracinglayer/opentracinglayer.go +++ b/store/opentracinglayer/opentracinglayer.go @@ -9213,6 +9213,24 @@ func (s *OpenTracingLayerThreadStore) GetPosts(threadID string, since int64) ([] return result, err } +func (s *OpenTracingLayerThreadStore) GetTeamsUnreadForUser(userID string, teamIDs []string) (map[string]*model.TeamUnread, error) { + origCtx := s.Root.Store.Context() + span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.GetTeamsUnreadForUser") + s.Root.Store.SetContext(newCtx) + defer func() { + s.Root.Store.SetContext(origCtx) + }() + + defer span.Finish() + result, err := s.ThreadStore.GetTeamsUnreadForUser(userID, teamIDs) + if err != nil { + span.LogFields(spanlog.Error(err)) + ext.Error.Set(span, true) + } + + return result, err +} + func (s *OpenTracingLayerThreadStore) GetThreadFollowers(threadID string, fetchOnlyActive bool) ([]string, error) { origCtx := s.Root.Store.Context() span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ThreadStore.GetThreadFollowers") diff --git a/store/retrylayer/retrylayer.go b/store/retrylayer/retrylayer.go index 5949d22d6b..fcc504abcc 100644 --- a/store/retrylayer/retrylayer.go +++ b/store/retrylayer/retrylayer.go @@ -10519,6 +10519,27 @@ func (s *RetryLayerThreadStore) GetPosts(threadID string, since int64) ([]*model } +func (s *RetryLayerThreadStore) GetTeamsUnreadForUser(userID string, teamIDs []string) (map[string]*model.TeamUnread, error) { + + tries := 0 + for { + result, err := s.ThreadStore.GetTeamsUnreadForUser(userID, teamIDs) + 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) GetThreadFollowers(threadID string, fetchOnlyActive bool) ([]string, error) { tries := 0 diff --git a/store/sqlstore/thread_store.go b/store/sqlstore/thread_store.go index f0e8aea333..8186da3477 100644 --- a/store/sqlstore/thread_store.go +++ b/store/sqlstore/thread_store.go @@ -8,6 +8,7 @@ import ( "database/sql" "encoding/json" "strconv" + "sync" "time" sq "github.com/Masterminds/squirrel" @@ -131,17 +132,7 @@ func (s *SqlThreadStore) GetThreadsForUser(userId, teamId string, opts model.Get fetchConditions := sq.And{ sq.Eq{"ThreadMemberships.UserId": userId}, sq.Eq{"ThreadMemberships.Following": true}, - } - if opts.TeamOnly { - fetchConditions = sq.And{ - sq.Eq{"Channels.TeamId": teamId}, - fetchConditions, - } - } else { - fetchConditions = sq.And{ - sq.Or{sq.Eq{"Channels.TeamId": teamId}, sq.Eq{"Channels.TeamId": ""}}, - fetchConditions, - } + sq.Or{sq.Eq{"Channels.TeamId": teamId}, sq.Eq{"Channels.TeamId": ""}}, } if !opts.Deleted { fetchConditions = sq.And{ @@ -363,6 +354,107 @@ func (s *SqlThreadStore) GetThreadsForUser(userId, teamId string, opts model.Get return result, nil } +// GetTeamsUnreadForUser returns the total unread threads and unread mentions +// for a user from all teams. +func (s *SqlThreadStore) GetTeamsUnreadForUser(userID string, teamIDs []string) (map[string]*model.TeamUnread, error) { + fetchConditions := sq.And{ + sq.Eq{"ThreadMemberships.UserId": userID}, + sq.Eq{"ThreadMemberships.Following": true}, + sq.Eq{"Channels.TeamId": teamIDs}, + sq.Eq{"COALESCE(Posts.DeleteAt, 0)": 0}, + } + + var wg sync.WaitGroup + var err1, err2 error + + unreadThreads := []struct { + Count int64 + TeamId string + }{} + unreadMentions := []struct { + Count int64 + TeamId string + }{} + + // Running these concurrently hasn't shown any major downside + // than running them serially. So using a bit of perf boost. + // In any case, they will be replaced by computed columns later. + wg.Add(1) + go func() { + defer wg.Done() + repliesQuery, repliesQueryArgs, err := s.getQueryBuilder(). + Select("COUNT(DISTINCT(Posts.RootId)) AS Count, TeamId"). + From("Posts"). + LeftJoin("ThreadMemberships ON Posts.RootId = ThreadMemberships.PostId"). + LeftJoin("Channels ON Posts.ChannelId = Channels.Id"). + Where(fetchConditions). + Where("Posts.CreateAt > ThreadMemberships.LastViewed"). + GroupBy("Channels.TeamId"). + ToSql() + if err != nil { + err1 = errors.Wrap(err, "GetTotalUnreadThreads_Tosql") + return + } + + err = s.GetReplicaX().Select(&unreadThreads, repliesQuery, repliesQueryArgs...) + if err != nil { + err1 = errors.Wrap(err, "failed to get total unread threads") + } + }() + + wg.Add(1) + go func() { + defer wg.Done() + mentionsQuery, mentionsQueryArgs, err := s.getQueryBuilder(). + Select("COALESCE(SUM(ThreadMemberships.UnreadMentions),0) AS Count, TeamId"). + From("ThreadMemberships"). + LeftJoin("Threads ON Threads.PostId = ThreadMemberships.PostId"). + LeftJoin("Posts ON Posts.Id = ThreadMemberships.PostId"). + LeftJoin("Channels ON Threads.ChannelId = Channels.Id"). + Where(fetchConditions). + GroupBy("Channels.TeamId"). + ToSql() + if err != nil { + err2 = errors.Wrap(err, "GetTotalUnreadMentions_Tosql") + } + + err = s.GetReplicaX().Select(&unreadMentions, mentionsQuery, mentionsQueryArgs...) + if err != nil { + err2 = errors.Wrap(err, "failed to get total unread mentions") + } + }() + + // Wait for them to be over + wg.Wait() + + if err1 != nil { + return nil, err1 + } + if err2 != nil { + return nil, err2 + } + + res := make(map[string]*model.TeamUnread) + // A bit of linear complexity here to create and return the map. + // This makes it easy to consume the output in the app layer. + for _, item := range unreadThreads { + res[item.TeamId] = &model.TeamUnread{ + ThreadCount: item.Count, + } + } + for _, item := range unreadMentions { + if _, ok := res[item.TeamId]; ok { + res[item.TeamId].ThreadMentionCount = item.Count + } else { + res[item.TeamId] = &model.TeamUnread{ + ThreadMentionCount: item.Count, + } + } + } + + return res, nil +} + func (s *SqlThreadStore) GetThreadFollowers(threadID string, fetchOnlyActive bool) ([]string, error) { users := []string{} diff --git a/store/store.go b/store/store.go index 081f99370d..ba0aea695d 100644 --- a/store/store.go +++ b/store/store.go @@ -298,6 +298,7 @@ type ThreadStore interface { Get(id string) (*model.Thread, error) GetThreadsForUser(userId, teamID string, opts model.GetUserThreadsOpts) (*model.Threads, error) GetThreadForUser(teamID string, threadMembership *model.ThreadMembership, extended bool) (*model.ThreadResponse, error) + GetTeamsUnreadForUser(userID string, teamIDs []string) (map[string]*model.TeamUnread, error) Delete(postID string) error GetPosts(threadID string, since int64) ([]*model.Post, error) diff --git a/store/storetest/mocks/ThreadStore.go b/store/storetest/mocks/ThreadStore.go index 4745fe7e05..fbe2bc4251 100644 --- a/store/storetest/mocks/ThreadStore.go +++ b/store/storetest/mocks/ThreadStore.go @@ -179,6 +179,29 @@ func (_m *ThreadStore) GetPosts(threadID string, since int64) ([]*model.Post, er return r0, r1 } +// GetTeamsUnreadForUser provides a mock function with given fields: userID, teamIDs +func (_m *ThreadStore) GetTeamsUnreadForUser(userID string, teamIDs []string) (map[string]*model.TeamUnread, error) { + ret := _m.Called(userID, teamIDs) + + var r0 map[string]*model.TeamUnread + if rf, ok := ret.Get(0).(func(string, []string) map[string]*model.TeamUnread); ok { + r0 = rf(userID, teamIDs) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(map[string]*model.TeamUnread) + } + } + + var r1 error + if rf, ok := ret.Get(1).(func(string, []string) error); ok { + r1 = rf(userID, teamIDs) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // GetThreadFollowers provides a mock function with given fields: threadID, fetchOnlyActive func (_m *ThreadStore) GetThreadFollowers(threadID string, fetchOnlyActive bool) ([]string, error) { ret := _m.Called(threadID, fetchOnlyActive) diff --git a/store/storetest/thread_store.go b/store/storetest/thread_store.go index a896b65c61..46a1829bf7 100644 --- a/store/storetest/thread_store.go +++ b/store/storetest/thread_store.go @@ -24,6 +24,7 @@ func TestThreadStore(t *testing.T, ss store.Store, s SqlStore) { t.Run("ThreadStorePermanentDeleteBatchThreadMembershipsForRetentionPolicies", func(t *testing.T) { testThreadStorePermanentDeleteBatchThreadMembershipsForRetentionPolicies(t, ss) }) + t.Run("GetTeamsUnreadForUser", func(t *testing.T) { testGetTeamsUnreadForUser(t, ss) }) } func testThreadStorePopulation(t *testing.T, ss store.Store) { @@ -520,12 +521,13 @@ func testThreadSQLOperations(t *testing.T, ss store.Store, s SqlStore) { }) } -func threadStoreCreateReply(t *testing.T, ss store.Store, channelID, postID string, createAt int64) *model.Post { +func threadStoreCreateReply(t *testing.T, ss store.Store, channelID, postID, userID string, createAt int64) *model.Post { reply, err := ss.Post().Save(&model.Post{ ChannelId: channelID, - UserId: model.NewId(), + UserId: userID, CreateAt: createAt, RootId: postID, + Message: model.NewRandomString(10), }) require.NoError(t, err) return reply @@ -553,7 +555,7 @@ func testThreadStorePermanentDeleteBatchForRetentionPolicies(t *testing.T, ss st UserId: model.NewId(), }) require.NoError(t, err) - threadStoreCreateReply(t, ss, channel.Id, post.Id, 2000) + threadStoreCreateReply(t, ss, channel.Id, post.Id, post.UserId, 2000) thread, err := ss.Thread().Get(post.Id) require.NoError(t, err) @@ -575,7 +577,7 @@ func testThreadStorePermanentDeleteBatchForRetentionPolicies(t *testing.T, ss st assert.Nil(t, thread, "thread should have been deleted by channel policy") // create a new thread - threadStoreCreateReply(t, ss, channel.Id, post.Id, 2000) + threadStoreCreateReply(t, ss, channel.Id, post.Id, post.UserId, 2000) thread, err = ss.Thread().Get(post.Id) require.NoError(t, err) @@ -641,7 +643,7 @@ func testThreadStorePermanentDeleteBatchThreadMembershipsForRetentionPolicies(t UserId: model.NewId(), }) require.NoError(t, err) - threadStoreCreateReply(t, ss, channel.Id, post.Id, 2000) + threadStoreCreateReply(t, ss, channel.Id, post.Id, post.UserId, 2000) threadMembership := createThreadMembership(userID, post.Id) @@ -702,3 +704,102 @@ func testThreadStorePermanentDeleteBatchThreadMembershipsForRetentionPolicies(t _, err = ss.Thread().GetMembershipForUser(userID, post.Id) require.Error(t, err, "thread membership should have been deleted because thread no longer exists") } + +func testGetTeamsUnreadForUser(t *testing.T, ss store.Store) { + userID := model.NewId() + 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) + } + team1, 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(&model.Channel{ + TeamId: team1.Id, + DisplayName: "DisplayName", + Name: "channel" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + post, err := ss.Post().Save(&model.Post{ + ChannelId: channel1.Id, + UserId: userID, + Message: model.NewRandomString(10), + }) + require.NoError(t, err) + threadStoreCreateReply(t, ss, channel1.Id, post.Id, post.UserId, model.GetMillis()) + createThreadMembership(userID, post.Id) + + teamsUnread, err := ss.Thread().GetTeamsUnreadForUser(userID, []string{team1.Id}) + require.NoError(t, err) + assert.Len(t, teamsUnread, 1) + assert.Equal(t, int64(1), teamsUnread[team1.Id].ThreadCount) + + post, err = ss.Post().Save(&model.Post{ + ChannelId: channel1.Id, + UserId: userID, + Message: model.NewRandomString(10), + }) + require.NoError(t, err) + threadStoreCreateReply(t, ss, channel1.Id, post.Id, post.UserId, model.GetMillis()) + createThreadMembership(userID, post.Id) + + teamsUnread, err = ss.Thread().GetTeamsUnreadForUser(userID, []string{team1.Id}) + require.NoError(t, err) + assert.Len(t, teamsUnread, 1) + assert.Equal(t, int64(2), teamsUnread[team1.Id].ThreadCount) + + team2, err := ss.Team().Save(&model.Team{ + DisplayName: "DisplayName", + Name: "team" + model.NewId(), + Email: MakeEmail(), + Type: model.TeamOpen, + }) + require.NoError(t, err) + channel2, err := ss.Channel().Save(&model.Channel{ + TeamId: team2.Id, + DisplayName: "DisplayName", + Name: "channel" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + post2, err := ss.Post().Save(&model.Post{ + ChannelId: channel2.Id, + UserId: userID, + Message: model.NewRandomString(10), + }) + require.NoError(t, err) + threadStoreCreateReply(t, ss, channel2.Id, post2.Id, post2.UserId, model.GetMillis()) + createThreadMembership(userID, post2.Id) + + teamsUnread, err = ss.Thread().GetTeamsUnreadForUser(userID, []string{team1.Id, team2.Id}) + require.NoError(t, err) + assert.Len(t, teamsUnread, 2) + assert.Equal(t, int64(2), teamsUnread[team1.Id].ThreadCount) + assert.Equal(t, int64(1), teamsUnread[team2.Id].ThreadCount) + + opts := store.ThreadMembershipOpts{ + Following: true, + IncrementMentions: true, + } + _, err = ss.Thread().MaintainMembership(userID, post2.Id, opts) + require.NoError(t, err) + + teamsUnread, err = ss.Thread().GetTeamsUnreadForUser(userID, []string{team2.Id}) + require.NoError(t, err) + assert.Len(t, teamsUnread, 1) + assert.Equal(t, int64(1), teamsUnread[team2.Id].ThreadCount) + assert.Equal(t, int64(1), teamsUnread[team2.Id].ThreadMentionCount) +} diff --git a/store/timerlayer/timerlayer.go b/store/timerlayer/timerlayer.go index 3ddd4aa2dc..70ab563494 100644 --- a/store/timerlayer/timerlayer.go +++ b/store/timerlayer/timerlayer.go @@ -8295,6 +8295,22 @@ func (s *TimerLayerThreadStore) GetPosts(threadID string, since int64) ([]*model return result, err } +func (s *TimerLayerThreadStore) GetTeamsUnreadForUser(userID string, teamIDs []string) (map[string]*model.TeamUnread, error) { + start := timemodule.Now() + + result, err := s.ThreadStore.GetTeamsUnreadForUser(userID, teamIDs) + + elapsed := float64(timemodule.Since(start)) / float64(timemodule.Second) + if s.Root.Metrics != nil { + success := "false" + if err == nil { + success = "true" + } + s.Root.Metrics.ObserveStoreMethodDuration("ThreadStore.GetTeamsUnreadForUser", success, elapsed) + } + return result, err +} + func (s *TimerLayerThreadStore) GetThreadFollowers(threadID string, fetchOnlyActive bool) ([]string, error) { start := timemodule.Now()