From aefd6c8e0feecfca10ffa7fb42aa237faff2bb54 Mon Sep 17 00:00:00 2001 From: Kitae Kim Date: Thu, 3 Mar 2022 18:34:18 +0900 Subject: [PATCH] Migrate buildSearchPostFilterClause to use Squirrel (#19583) Automatic Merge --- store/sqlstore/post_store.go | 182 +++++++++++++++-------------------- store/sqlstore/store.go | 5 + 2 files changed, 82 insertions(+), 105 deletions(-) diff --git a/store/sqlstore/post_store.go b/store/sqlstore/post_store.go index b005abb21b..85e3cecdc2 100644 --- a/store/sqlstore/post_store.go +++ b/store/sqlstore/post_store.go @@ -9,7 +9,6 @@ import ( "fmt" "reflect" "regexp" - "strconv" "strings" "sync" @@ -1583,154 +1582,122 @@ var specialSearchChar = []string{ ":", } -func (s *SqlPostStore) buildCreateDateFilterClause(params *model.SearchParams, queryParams map[string]interface{}, builder sq.SelectBuilder) (sq.SelectBuilder, map[string]interface{}) { +func (s *SqlPostStore) buildCreateDateFilterClause(params *model.SearchParams, builder sq.SelectBuilder) sq.SelectBuilder { // handle after: before: on: filters if params.OnDate != "" { onDateStart, onDateEnd := params.GetOnDateMillis() - queryParams["OnDateStart"] = strconv.FormatInt(onDateStart, 10) - queryParams["OnDateEnd"] = strconv.FormatInt(onDateEnd, 10) - // between `on date` start of day and end of day - builder = builder.Where("CreateAt BETWEEN :OnDateStart AND :OnDateEnd") - return builder, queryParams + builder = builder.Where("CreateAt BETWEEN ? AND ?", onDateStart, onDateEnd) + return builder } if params.ExcludedDate != "" { excludedDateStart, excludedDateEnd := params.GetExcludedDateMillis() - queryParams["ExcludedDateStart"] = strconv.FormatInt(excludedDateStart, 10) - queryParams["ExcludedDateEnd"] = strconv.FormatInt(excludedDateEnd, 10) - - builder = builder.Where("CreateAt NOT BETWEEN :ExcludedDateStart AND :ExcludedDateEnd") + builder = builder.Where("CreateAt NOT BETWEEN ? AND ?", excludedDateStart, excludedDateEnd) } if params.AfterDate != "" { afterDate := params.GetAfterDateMillis() - queryParams["AfterDate"] = strconv.FormatInt(afterDate, 10) - // greater than `after date` - builder = builder.Where("CreateAt >= :AfterDate") + builder = builder.Where("CreateAt >= ?", afterDate) } if params.BeforeDate != "" { beforeDate := params.GetBeforeDateMillis() - queryParams["BeforeDate"] = strconv.FormatInt(beforeDate, 10) - // less than `before date` - builder = builder.Where("CreateAt <= :BeforeDate") + builder = builder.Where("CreateAt <= ?", beforeDate) } if params.ExcludedAfterDate != "" { afterDate := params.GetExcludedAfterDateMillis() - queryParams["ExcludedAfterDate"] = strconv.FormatInt(afterDate, 10) - - builder = builder.Where("CreateAt < :ExcludedAfterDate") + builder = builder.Where("CreateAt < ?", afterDate) } if params.ExcludedBeforeDate != "" { beforeDate := params.GetExcludedBeforeDateMillis() - queryParams["ExcludedBeforeDate"] = strconv.FormatInt(beforeDate, 10) - - builder = builder.Where("CreateAt > :ExcludedBeforeDate") + builder = builder.Where("CreateAt > ?", beforeDate) } - return builder, queryParams + return builder } -func (s *SqlPostStore) buildSearchTeamFilterClause(teamId string, queryParams map[string]interface{}, builder sq.SelectBuilder) (sq.SelectBuilder, map[string]interface{}) { +func (s *SqlPostStore) buildSearchTeamFilterClause(teamId string, builder sq.SelectBuilder) sq.SelectBuilder { if teamId == "" { - return builder, queryParams + return builder } - queryParams["TeamId"] = teamId - - return builder.Where("(TeamId = :TeamId OR TeamId = '')"), queryParams + return builder.Where(sq.Or{ + sq.Eq{"TeamId": teamId}, + sq.Eq{"TeamId": ""}, + }) } -func (s *SqlPostStore) buildSearchChannelFilterClause(channels []string, paramPrefix string, exclusion bool, queryParams map[string]interface{}, byName bool, builder sq.SelectBuilder) (sq.SelectBuilder, map[string]interface{}) { +func (s *SqlPostStore) buildSearchChannelFilterClause(channels []string, exclusion bool, byName bool, builder sq.SelectBuilder) sq.SelectBuilder { if len(channels) == 0 { - return builder, queryParams + return builder } - clauseSlice := []string{} - for i, channel := range channels { - paramName := paramPrefix + strconv.FormatInt(int64(i), 10) - clauseSlice = append(clauseSlice, ":"+paramName) - queryParams[paramName] = channel - } - clause := strings.Join(clauseSlice, ", ") if byName { if exclusion { - return builder.Where("Name NOT IN (" + clause + ")"), queryParams + return builder.Where(sq.NotEq{"Name": channels}) } - return builder.Where("Name IN (" + clause + ")"), queryParams + return builder.Where(sq.Eq{"Name": channels}) } if exclusion { - return builder.Where("Id NOT IN (" + clause + ")"), queryParams + return builder.Where(sq.NotEq{"Id": channels}) } - return builder.Where("Id IN (" + clause + ")"), queryParams + return builder.Where(sq.Eq{"Id": channels}) } -func (s *SqlPostStore) buildSearchUserFilterClause(users []string, paramPrefix string, exclusion bool, queryParams map[string]interface{}, byUsername bool) (string, map[string]interface{}) { +func (s *SqlPostStore) buildSearchUserFilterClause(users []string, exclusion bool, byUsername bool, builder sq.SelectBuilder) sq.SelectBuilder { if len(users) == 0 { - return "", queryParams + return builder } - clauseSlice := []string{} - for i, user := range users { - paramName := paramPrefix + strconv.FormatInt(int64(i), 10) - clauseSlice = append(clauseSlice, ":"+paramName) - queryParams[paramName] = user - } - clause := strings.Join(clauseSlice, ", ") + if byUsername { if exclusion { - return "AND Username NOT IN (" + clause + ")", queryParams + return builder.Where(sq.NotEq{"Username": users}) } - return "AND Username IN (" + clause + ")", queryParams + return builder.Where(sq.Eq{"Username": users}) } + if exclusion { - return "AND Id NOT IN (" + clause + ")", queryParams + return builder.Where(sq.NotEq{"Id": users}) } - return "AND Id IN (" + clause + ")", queryParams + return builder.Where(sq.Eq{"Id": users}) } -func (s *SqlPostStore) buildSearchPostFilterClause(fromUsers []string, excludedUsers []string, queryParams map[string]interface{}, userByUsername bool, builder sq.SelectBuilder) (sq.SelectBuilder, map[string]interface{}) { +func (s *SqlPostStore) buildSearchPostFilterClause(teamID string, fromUsers []string, excludedUsers []string, userByUsername bool, builder sq.SelectBuilder) (sq.SelectBuilder, error) { if len(fromUsers) == 0 && len(excludedUsers) == 0 { - return builder, queryParams + return builder, nil } - filterQuery := ` - UserId IN ( - SELECT - Id - FROM - Users, - TeamMembers - WHERE - TeamMembers.TeamId = :TeamId - AND Users.Id = TeamMembers.UserId - FROM_USER_FILTER - EXCLUDED_USER_FILTER)` + // Sub-query builder. + sb := s.getSubQueryBuilder().Select("Id").From("Users, TeamMembers").Where( + sq.And{ + sq.Eq{"TeamMembers.TeamId": teamID}, + sq.Expr("Users.Id = TeamMembers.UserId"), + }) + sb = s.buildSearchUserFilterClause(fromUsers, false, userByUsername, sb) + sb = s.buildSearchUserFilterClause(excludedUsers, true, userByUsername, sb) + subQuery, subQueryArgs, err := sb.ToSql() + if err != nil { + return sq.SelectBuilder{}, err + } - fromUserClause, queryParams := s.buildSearchUserFilterClause(fromUsers, "FromUser", false, queryParams, userByUsername) - filterQuery = strings.Replace(filterQuery, "FROM_USER_FILTER", fromUserClause, 1) - - excludedUserClause, queryParams := s.buildSearchUserFilterClause(excludedUsers, "ExcludedUser", true, queryParams, userByUsername) - filterQuery = strings.Replace(filterQuery, "EXCLUDED_USER_FILTER", excludedUserClause, 1) - - return builder.Where(filterQuery), queryParams + /* + * Squirrel does not support a sub-query in the WHERE condition. + * https://github.com/Masterminds/squirrel/issues/299 + */ + return builder.Where("UserId IN ("+subQuery+")", subQueryArgs...), nil } func (s *SqlPostStore) Search(teamId string, userId string, params *model.SearchParams) (*model.PostList, error) { return s.search(teamId, userId, params, true, true) } -// TODO: convert to squirrel func (s *SqlPostStore) search(teamId string, userId string, params *model.SearchParams, channelsByName bool, userByUsername bool) (*model.PostList, error) { - queryParams := map[string]interface{}{ - "UserId": userId, - } - list := model.NewPostList() if params.Terms == "" && params.ExcludedTerms == "" && len(params.InChannels) == 0 && len(params.ExcludedChannels) == 0 && @@ -1748,8 +1715,12 @@ func (s *SqlPostStore) search(teamId string, userId string, params *model.Search OrderByClause("CreateAt DESC"). Limit(100) - baseQuery, queryParams = s.buildSearchPostFilterClause(params.FromUsers, params.ExcludedUsers, queryParams, userByUsername, baseQuery) - baseQuery, queryParams = s.buildCreateDateFilterClause(params, queryParams, baseQuery) + var err error + baseQuery, err = s.buildSearchPostFilterClause(teamId, params.FromUsers, params.ExcludedUsers, userByUsername, baseQuery) + if err != nil { + return nil, errors.Wrap(err, "failed to build search post filter clause") + } + baseQuery = s.buildCreateDateFilterClause(params, baseQuery) termMap := map[string]bool{} terms := params.Terms @@ -1773,7 +1744,8 @@ func (s *SqlPostStore) search(teamId string, userId string, params *model.Search // we've already confirmed that we have a channel or user to search for } else if s.DriverName() == model.DatabaseDriverPostgres { // Parse text for wildcards - if wildcard, err := regexp.Compile(`\*($| )`); err == nil { + var wildcard *regexp.Regexp + if wildcard, err = regexp.Compile(`\*($| )`); err == nil { terms = wildcard.ReplaceAllLiteralString(terms, ":* ") excludedTerms = wildcard.ReplaceAllLiteralString(excludedTerms, ":* ") } @@ -1783,19 +1755,19 @@ func (s *SqlPostStore) search(teamId string, userId string, params *model.Search excludeClause = " & !(" + strings.Join(strings.Fields(excludedTerms), " | ") + ")" } + var termsClause string if params.OrTerms { - queryParams["Terms"] = "(" + strings.Join(strings.Fields(terms), " | ") + ")" + excludeClause + termsClause = "(" + strings.Join(strings.Fields(terms), " | ") + ")" + excludeClause } else if strings.HasPrefix(terms, `"`) && strings.HasSuffix(terms, `"`) { - queryParams["Terms"] = "(" + strings.Join(strings.Fields(terms), " <-> ") + ")" + excludeClause + termsClause = "(" + strings.Join(strings.Fields(terms), " <-> ") + ")" + excludeClause } else { - queryParams["Terms"] = "(" + strings.Join(strings.Fields(terms), " & ") + ")" + excludeClause + termsClause = "(" + strings.Join(strings.Fields(terms), " & ") + ")" + excludeClause } - searchClause := fmt.Sprintf("to_tsvector('english', %s) @@ to_tsquery('english', :Terms)", searchType) - baseQuery = baseQuery.Where(searchClause) + searchClause := fmt.Sprintf("to_tsvector('english', %s) @@ to_tsquery('english', ?)", searchType) + baseQuery = baseQuery.Where(searchClause, termsClause) } else if s.DriverName() == model.DatabaseDriverMysql { if searchType == "Message" { - var err error terms, err = removeMysqlStopWordsFromTerms(terms) if err != nil { return nil, errors.Wrap(err, "failed to remove Mysql stop-words from terms") @@ -1806,26 +1778,27 @@ func (s *SqlPostStore) search(teamId string, userId string, params *model.Search } } - searchClause := fmt.Sprintf("MATCH (%s) AGAINST (:Terms IN BOOLEAN MODE)", searchType) - baseQuery = baseQuery.Where(searchClause) - excludeClause := "" if excludedTerms != "" { excludeClause = " -(" + excludedTerms + ")" } + var termsClause string if params.OrTerms { - queryParams["Terms"] = terms + excludeClause + termsClause = terms + excludeClause } else { splitTerms := []string{} for _, t := range strings.Fields(terms) { splitTerms = append(splitTerms, "+"+t) } - queryParams["Terms"] = strings.Join(splitTerms, " ") + excludeClause + termsClause = strings.Join(splitTerms, " ") + excludeClause } + + searchClause := fmt.Sprintf("MATCH (%s) AGAINST (? IN BOOLEAN MODE)", searchType) + baseQuery = baseQuery.Where(searchClause, termsClause) } - inQuery := s.getQueryBuilder().Select("Id"). + inQuery := s.getSubQueryBuilder().Select("Id"). From("Channels, ChannelMembers"). Where("Id = ChannelId") @@ -1834,29 +1807,28 @@ func (s *SqlPostStore) search(teamId string, userId string, params *model.Search } if !params.SearchWithoutUserId { - inQuery = inQuery.Where("UserId = :UserId") + inQuery = inQuery.Where("UserId = ?", userId) } - inQuery, queryParams = s.buildSearchTeamFilterClause(teamId, queryParams, inQuery) - inQuery, queryParams = s.buildSearchChannelFilterClause(params.InChannels, "InChannel", false, queryParams, channelsByName, inQuery) - inQuery, queryParams = s.buildSearchChannelFilterClause(params.ExcludedChannels, "ExcludedChannel", true, queryParams, channelsByName, inQuery) + inQuery = s.buildSearchTeamFilterClause(teamId, inQuery) + inQuery = s.buildSearchChannelFilterClause(params.InChannels, false, channelsByName, inQuery) + inQuery = s.buildSearchChannelFilterClause(params.ExcludedChannels, true, channelsByName, inQuery) - inQueryClause, _, err := inQuery.ToSql() + inQueryClause, inQueryClauseArgs, err := inQuery.ToSql() if err != nil { return nil, err } - baseQuery = baseQuery.Where(fmt.Sprintf("ChannelId IN (%s)", inQueryClause)) + baseQuery = baseQuery.Where(fmt.Sprintf("ChannelId IN (%s)", inQueryClause), inQueryClauseArgs...) - searchQuery, _, err := baseQuery.ToSql() + searchQuery, searchQueryArgs, err := baseQuery.ToSql() if err != nil { return nil, err } var posts []*model.Post - _, err = s.GetSearchReplica().Select(&posts, searchQuery, queryParams) - if err != nil { + if err := s.GetSearchReplicaX().Select(&posts, searchQuery, searchQueryArgs...); err != nil { mlog.Warn("Query error searching posts.", mlog.Err(err)) // Don't return the error to the caller as it is of no use to the user. Instead return an empty set of search results. } else { diff --git a/store/sqlstore/store.go b/store/sqlstore/store.go index dd62beffa6..cc43626017 100644 --- a/store/sqlstore/store.go +++ b/store/sqlstore/store.go @@ -968,6 +968,11 @@ func (ss *SqlStore) getQueryBuilder() sq.StatementBuilderType { return builder } +// getSubQueryBuilder is necessary to generate the SQL query and args to pass to sub-queries because squirrel does not support WHERE clause in sub-queries. +func (ss *SqlStore) getSubQueryBuilder() sq.StatementBuilderType { + return sq.StatementBuilder.PlaceholderFormat(sq.Question) +} + func (ss *SqlStore) CheckIntegrity() <-chan model.IntegrityCheckResult { results := make(chan model.IntegrityCheckResult) go CheckRelationalIntegrity(ss, results)