From 981d1d869aa4578612d30ae7c1ed8b37ce53f5b7 Mon Sep 17 00:00:00 2001 From: Jesse Hallam Date: Tue, 22 Apr 2025 16:33:32 -0300 Subject: [PATCH] avoid SELECT * in channel member history store (#30828) --- .../sqlstore/channel_member_history_store.go | 49 +++++++---- .../storetest/channel_member_history_store.go | 82 ++++++++++++++++++- 2 files changed, 114 insertions(+), 17 deletions(-) diff --git a/server/channels/store/sqlstore/channel_member_history_store.go b/server/channels/store/sqlstore/channel_member_history_store.go index 5596c4c414..109eeb1d03 100644 --- a/server/channels/store/sqlstore/channel_member_history_store.go +++ b/server/channels/store/sqlstore/channel_member_history_store.go @@ -17,12 +17,25 @@ import ( type SqlChannelMemberHistoryStore struct { *SqlStore + + channelMemberHistoryQuery sq.SelectBuilder } func newSqlChannelMemberHistoryStore(sqlStore *SqlStore) store.ChannelMemberHistoryStore { - return &SqlChannelMemberHistoryStore{ + s := &SqlChannelMemberHistoryStore{ SqlStore: sqlStore, } + + s.channelMemberHistoryQuery = s.getQueryBuilder(). + Select( + "ChannelMemberHistory.ChannelId", + "ChannelMemberHistory.UserId", + "ChannelMemberHistory.JoinTime", + "ChannelMemberHistory.LeaveTime", + ). + From("ChannelMemberHistory") + + return s } func (s SqlChannelMemberHistoryStore) LogJoinEvent(userId string, channelId string, joinTime int64) error { @@ -69,7 +82,7 @@ func (s SqlChannelMemberHistoryStore) GetChannelsWithActivityDuring(startTime in // ChannelMemberHistory has been in production for long enough that we are assuming the export period // starts after the ChannelMemberHistory table was first introduced subqueryPosts := s.getSubQueryBuilder(). - Select("p.ChannelId"). + Select("p.ChannelId AS ChannelId"). Distinct(). From("Posts AS p"). Where( @@ -80,7 +93,7 @@ func (s SqlChannelMemberHistoryStore) GetChannelsWithActivityDuring(startTime in }) subqueryCMH := s.getSubQueryBuilder(). - Select("cmh.ChannelId"). + Select("cmh.ChannelId AS ChannelId"). Distinct(). From("ChannelMemberHistory AS cmh"). Where( @@ -100,9 +113,8 @@ func (s SqlChannelMemberHistoryStore) GetChannelsWithActivityDuring(startTime in return nil, errors.Wrap(err, "GetChannelsWithActivityDuring unionExpr to sql") } - // no bound args in this expression query, _, err := s.getQueryBuilder(). - Select("*"). + Select("ChannelId"). From(unionExpr).ToSql() if err != nil { return nil, errors.Wrap(err, "GetChannelsWithActivityDuring query to sql") @@ -160,25 +172,30 @@ func (s SqlChannelMemberHistoryStore) hasDataAtOrBefore(time int64) (bool, error } func (s SqlChannelMemberHistoryStore) getFromChannelMemberHistoryTable(startTime int64, endTime int64, channelIds []string) ([]*model.ChannelMemberHistoryResult, error) { - query, args, err := s.getQueryBuilder(). - Select(`cmh.*, u.Email AS "Email", u.Username, Bots.UserId IS NOT NULL AS IsBot, u.DeleteAt AS UserDeleteAt`). - From("ChannelMemberHistory cmh"). - Join("Users u ON cmh.UserId = u.Id"). + query := s.channelMemberHistoryQuery. + Column("u.Email AS \"Email\""). + Column("u.Username"). + Column("Bots.UserId IS NOT NULL AS IsBot"). + Column("u.DeleteAt AS UserDeleteAt"). + Join("Users u ON ChannelMemberHistory.UserId = u.Id"). LeftJoin("Bots ON Bots.UserId = u.Id"). Where(sq.And{ - sq.Eq{"cmh.ChannelId": channelIds}, - sq.LtOrEq{"cmh.JoinTime": endTime}, + sq.Eq{"ChannelMemberHistory.ChannelId": channelIds}, + sq.LtOrEq{"ChannelMemberHistory.JoinTime": endTime}, sq.Or{ - sq.Eq{"cmh.LeaveTime": nil}, - sq.GtOrEq{"cmh.LeaveTime": startTime}, + sq.Eq{"ChannelMemberHistory.LeaveTime": nil}, + sq.GtOrEq{"ChannelMemberHistory.LeaveTime": startTime}, }, }). - OrderBy("cmh.JoinTime ASC").ToSql() + OrderBy("ChannelMemberHistory.JoinTime ASC") + + queryString, args, err := query.ToSql() if err != nil { return nil, errors.Wrap(err, "channel_member_history_to_sql") } + histories := []*model.ChannelMemberHistoryResult{} - if err := s.GetReplica().Select(&histories, query, args...); err != nil { + if err := s.GetReplica().Select(&histories, queryString, args...); err != nil { return nil, err } @@ -234,7 +251,7 @@ func (s SqlChannelMemberHistoryStore) DeleteOrphanedRows(limit int) (deleted int // We need the extra level of nesting to deal with MySQL's locking const query = ` DELETE FROM ChannelMemberHistory WHERE (ChannelId, UserId, JoinTime) IN ( - SELECT * FROM ( + SELECT ChannelId, UserId, JoinTime FROM ( SELECT ChannelId, UserId, JoinTime FROM ChannelMemberHistory LEFT JOIN Channels ON ChannelMemberHistory.ChannelId = Channels.Id WHERE Channels.Id IS NULL diff --git a/server/channels/store/storetest/channel_member_history_store.go b/server/channels/store/storetest/channel_member_history_store.go index 3ee0240847..e6bbae7635 100644 --- a/server/channels/store/storetest/channel_member_history_store.go +++ b/server/channels/store/storetest/channel_member_history_store.go @@ -25,6 +25,7 @@ func TestChannelMemberHistoryStore(t *testing.T, rctx request.CTX, ss store.Stor t.Run("TestPermanentDeleteBatch", func(t *testing.T) { testPermanentDeleteBatch(t, rctx, ss) }) t.Run("TestPermanentDeleteBatchForRetentionPolicies", func(t *testing.T) { testPermanentDeleteBatchForRetentionPolicies(t, rctx, ss) }) t.Run("TestGetChannelsLeftSince", func(t *testing.T) { testGetChannelsLeftSince(t, rctx, ss) }) + t.Run("TestDeleteOrphanedRows", func(t *testing.T) { testDeleteOrphanedRows(t, rctx, ss) }) } func testLogJoinEvent(t *testing.T, rctx request.CTX, ss store.Store) { @@ -354,7 +355,7 @@ func testGetUsersInChannelAtChannelMembers(t *testing.T, rctx request.CTX, ss st user = *userPtr // clear any existing ChannelMemberHistory data that might interfere with our test - var tableDataTruncated = false + tableDataTruncated := false for !tableDataTruncated { var count int64 count, _, err = ss.ChannelMemberHistory().PermanentDeleteBatchForRetentionPolicies( @@ -600,3 +601,82 @@ func testGetChannelsLeftSince(t *testing.T, rctx request.CTX, ss store.Store) { require.NoError(t, err) assert.Equal(t, []string{channel.Id}, ids) } + +func testDeleteOrphanedRows(t *testing.T, rctx request.CTX, ss store.Store) { + // Create a channel + channelToKeep := &model.Channel{ + TeamId: model.NewId(), + DisplayName: "Channel to keep", + Name: model.NewId(), + Type: model.ChannelTypeOpen, + } + channelToKeep, err := ss.Channel().Save(rctx, channelToKeep, -1) + require.NoError(t, err) + + // Create a user + user := model.User{ + Email: MakeEmail(), + Nickname: model.NewId(), + Username: model.NewUsername(), + } + userPtr, err := ss.User().Save(rctx, &user) + require.NoError(t, err) + user = *userPtr + + // Add user to channel (via channel member history) + joinTime := model.GetMillis() + err = ss.ChannelMemberHistory().LogJoinEvent(user.Id, channelToKeep.Id, joinTime) + require.NoError(t, err) + + // Create multiple orphaned channel member history entries + // We'll use an ID that doesn't exist in the Channels table + nonExistentChannelId := model.NewId() + + // Create 3 orphaned entries + err = ss.ChannelMemberHistory().LogJoinEvent(user.Id, nonExistentChannelId, joinTime) + require.NoError(t, err) + + err = ss.ChannelMemberHistory().LogJoinEvent(model.NewId(), nonExistentChannelId, joinTime+100) + require.NoError(t, err) + + err = ss.ChannelMemberHistory().LogJoinEvent(model.NewId(), nonExistentChannelId, joinTime+200) + require.NoError(t, err) + + // Verify the data is setup correctly + channelIds, err := ss.ChannelMemberHistory().GetChannelsWithActivityDuring(joinTime-100, joinTime+300) + require.NoError(t, err) + assert.Contains(t, channelIds, channelToKeep.Id, "Channel to keep should still have history") + assert.Contains(t, channelIds, nonExistentChannelId, "Orphaned channel should still have history") + + // Test with limit of 0 (should delete nothing) + deletedCount, err := ss.ChannelMemberHistory().DeleteOrphanedRows(0) + require.NoError(t, err) + require.Equal(t, int64(0), deletedCount, "Should delete nothing with limit of 0") + + // Verify the data is unchanged + channelIds, err = ss.ChannelMemberHistory().GetChannelsWithActivityDuring(joinTime-100, joinTime+300) + require.NoError(t, err) + assert.Contains(t, channelIds, channelToKeep.Id, "Channel to keep should still have history") + assert.Contains(t, channelIds, nonExistentChannelId, "Orphaned channel should still have history") + + // Test limit parameter by deleting only 2 of the 3 orphaned rows + deletedCount, err = ss.ChannelMemberHistory().DeleteOrphanedRows(2) + require.NoError(t, err) + require.Equal(t, int64(2), deletedCount, "Should have deleted exactly 2 orphaned rows due to limit") + + // Delete the remaining orphaned row + deletedCount, err = ss.ChannelMemberHistory().DeleteOrphanedRows(100) + require.NoError(t, err) + require.Equal(t, int64(1), deletedCount, "Should have deleted the remaining orphaned row") + + // Verify the orphaned entries are removed and valid entries remain + channelIds, err = ss.ChannelMemberHistory().GetChannelsWithActivityDuring(joinTime-100, joinTime+300) + require.NoError(t, err) + assert.Contains(t, channelIds, channelToKeep.Id, "Channel to keep should still have history") + assert.NotContains(t, channelIds, nonExistentChannelId, "Orphaned channel should not have history") + + // Calling it again should delete nothing since orphans are gone + deletedCount, err = ss.ChannelMemberHistory().DeleteOrphanedRows(100) + require.NoError(t, err) + require.Equal(t, int64(0), deletedCount, "No rows should be deleted when no orphans exist") +}