From b5e78a0ce12e555e0a433fa0e11d27462947b67e Mon Sep 17 00:00:00 2001 From: Jesse Hallam Date: Fri, 11 Mar 2022 10:57:42 -0400 Subject: [PATCH] MM-42282: handle teamId parameter correctly (#19685) As per https://community-daily.mattermost.com/core/pl/ugs7ue6e4j8a7cgegk1bxje8to, `ThreadStore.GetThreadsForUser` accepts a `teamId` parameter, but incorrectly handles an empty value of `""` as looking only for channels with an empty `teamId` (aka DMs and GMs) instead of finding all channels and effectively ignoring the team property. Fixes: https://mattermost.atlassian.net/browse/MM-42282 --- store/sqlstore/thread_store.go | 11 ++- store/storetest/thread_store.go | 153 ++++++++++++++++++++++++++++++++ 2 files changed, 163 insertions(+), 1 deletion(-) diff --git a/store/sqlstore/thread_store.go b/store/sqlstore/thread_store.go index d398222c7c..602456cac2 100644 --- a/store/sqlstore/thread_store.go +++ b/store/sqlstore/thread_store.go @@ -67,7 +67,16 @@ func (s *SqlThreadStore) GetThreadsForUser(userId, teamId string, opts model.Get fetchConditions := sq.And{ sq.Eq{"ThreadMemberships.UserId": userId}, sq.Eq{"ThreadMemberships.Following": true}, - sq.Or{sq.Eq{"Channels.TeamId": teamId}, sq.Eq{"Channels.TeamId": ""}}, + } + + if teamId != "" { + fetchConditions = sq.And{ + fetchConditions, + sq.Or{ + sq.Eq{"Channels.TeamId": teamId}, + sq.Eq{"Channels.TeamId": ""}, + }, + } } if !opts.Deleted { fetchConditions = sq.And{ diff --git a/store/storetest/thread_store.go b/store/storetest/thread_store.go index 933d78f3b9..01110b3e14 100644 --- a/store/storetest/thread_store.go +++ b/store/storetest/thread_store.go @@ -23,6 +23,7 @@ func TestThreadStore(t *testing.T, ss store.Store, s SqlStore) { testThreadStorePermanentDeleteBatchThreadMembershipsForRetentionPolicies(t, ss, s) }) t.Run("GetTeamsUnreadForUser", func(t *testing.T) { testGetTeamsUnreadForUser(t, ss) }) + t.Run("GetThreadsForUser", func(t *testing.T) { testGetThreadsForUser(t, ss) }) t.Run("MarkAllAsReadByChannels", func(t *testing.T) { testMarkAllAsReadByChannels(t, ss) }) } @@ -678,6 +679,158 @@ func testGetTeamsUnreadForUser(t *testing.T, ss store.Store) { assert.Equal(t, int64(1), teamsUnread[team2.Id].ThreadMentionCount) } +func testGetThreadsForUser(t *testing.T, ss store.Store) { + user1, err := ss.User().Save(&model.User{ + Username: "user1" + model.NewId(), + Email: MakeEmail(), + }) + require.NoError(t, err) + user2, err := ss.User().Save(&model.User{ + Username: "user2" + model.NewId(), + Email: MakeEmail(), + }) + require.NoError(t, err) + + user1ID := user1.Id + user2ID := user2.Id + + team1, err := ss.Team().Save(&model.Team{ + DisplayName: "Team1", + Name: "team" + model.NewId(), + Email: MakeEmail(), + Type: model.TeamOpen, + }) + require.NoError(t, err) + + team2, err := ss.Team().Save(&model.Team{ + DisplayName: "Team2", + Name: "team" + model.NewId(), + Email: MakeEmail(), + Type: model.TeamOpen, + }) + require.NoError(t, err) + + team1channel1, err := ss.Channel().Save(&model.Channel{ + TeamId: team1.Id, + DisplayName: "Channel1", + Name: "channel" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + + team2channel1, err := ss.Channel().Save(&model.Channel{ + TeamId: team2.Id, + DisplayName: "Channel2", + Name: "channel" + model.NewId(), + Type: model.ChannelTypeOpen, + }, -1) + require.NoError(t, err) + + dm1, err := ss.Channel().CreateDirectChannel(&model.User{Id: user1ID}, &model.User{Id: user2ID}) + require.NoError(t, err) + + gm1, err := ss.Channel().Save(&model.Channel{ + DisplayName: "GM", + Name: "gm" + model.NewId(), + Type: model.ChannelTypeGroup, + }, -1) + require.NoError(t, err) + + team1channel1post1, err := ss.Post().Save(&model.Post{ + ChannelId: team1channel1.Id, + UserId: user1ID, + Message: model.NewRandomString(10), + }) + require.NoError(t, err) + + team1channel1post2, err := ss.Post().Save(&model.Post{ + ChannelId: team1channel1.Id, + UserId: user1ID, + Message: model.NewRandomString(10), + }) + require.NoError(t, err) + + team2channel1post1, err := ss.Post().Save(&model.Post{ + ChannelId: team2channel1.Id, + UserId: user1ID, + Message: model.NewRandomString(10), + }) + require.NoError(t, err) + + dm1post1, err := ss.Post().Save(&model.Post{ + ChannelId: dm1.Id, + UserId: user1ID, + Message: model.NewRandomString(10), + }) + require.NoError(t, err) + + gm1post1, err := ss.Post().Save(&model.Post{ + ChannelId: gm1.Id, + UserId: user1ID, + Message: model.NewRandomString(10), + }) + require.NoError(t, err) + + threadStoreCreateReply(t, ss, team1channel1.Id, team1channel1post1.Id, user2ID, model.GetMillis()) + threadStoreCreateReply(t, ss, team1channel1.Id, team1channel1post2.Id, user2ID, model.GetMillis()) + threadStoreCreateReply(t, ss, team2channel1.Id, team2channel1post1.Id, user2ID, model.GetMillis()) + threadStoreCreateReply(t, ss, dm1.Id, dm1post1.Id, user2ID, model.GetMillis()) + threadStoreCreateReply(t, ss, gm1.Id, gm1post1.Id, user2ID, model.GetMillis()) + + 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) + } + + createThreadMembership(user1ID, team1channel1post1.Id) + createThreadMembership(user1ID, team1channel1post2.Id) + createThreadMembership(user1ID, team2channel1post1.Id) + createThreadMembership(user1ID, dm1post1.Id) + createThreadMembership(user1ID, gm1post1.Id) + + t.Run("no team specified, user1", func(t *testing.T) { + threads, err := ss.Thread().GetThreadsForUser(user1ID, "", model.GetUserThreadsOpts{}) + require.NoError(t, err) + + // 2 threads from team1, 1 threads from team2, 1 dm thread, 1 gm thread + assert.EqualValues(t, 5, threads.Total) + assert.EqualValues(t, 5, threads.TotalUnreadThreads) + assert.EqualValues(t, 0, threads.TotalUnreadMentions) + assert.Len(t, threads.Threads, 5) + }) + + t.Run("team1 specified, user1", func(t *testing.T) { + threads, err := ss.Thread().GetThreadsForUser(user1ID, team1.Id, model.GetUserThreadsOpts{}) + require.NoError(t, err) + + // 2 threads from team1, 1 dm thread, 1 gm thread + assert.EqualValues(t, 4, threads.Total) + assert.EqualValues(t, 4, threads.TotalUnreadThreads) + assert.EqualValues(t, 0, threads.TotalUnreadMentions) + assert.Len(t, threads.Threads, 4) + }) + + t.Run("team2 specified, user1", func(t *testing.T) { + threads, err := ss.Thread().GetThreadsForUser(user1ID, team2.Id, model.GetUserThreadsOpts{}) + require.NoError(t, err) + + // 1 thread from team1, 1 dm thread, 1 gm thread + assert.EqualValues(t, 3, threads.Total) + assert.EqualValues(t, 3, threads.TotalUnreadThreads) + assert.EqualValues(t, 0, threads.TotalUnreadMentions) + assert.Len(t, threads.Threads, 3) + }) +} + func testMarkAllAsReadByChannels(t *testing.T, ss store.Store) { postingUserId := model.NewId() userAID := model.NewId()