diff --git a/store/sqlstore/sqlx_wrapper.go b/store/sqlstore/sqlx_wrapper.go index b29d56ce96..ef2236c121 100644 --- a/store/sqlstore/sqlx_wrapper.go +++ b/store/sqlstore/sqlx_wrapper.go @@ -35,17 +35,24 @@ func (w *StoreTestWrapper) DriverName() string { return w.orig.DriverName() } +type Builder interface { + ToSql() (string, []interface{}, error) +} + // sqlxExecutor exposes sqlx operations. It is used to enable some internal store methods to // accept both transactions (*sqlxTxWrapper) and common db handlers (*sqlxDbWrapper). type sqlxExecutor interface { Get(dest interface{}, query string, args ...interface{}) error + GetBuilder(dest interface{}, builder Builder) error NamedExec(query string, arg interface{}) (sql.Result, error) Exec(query string, args ...interface{}) (sql.Result, error) + ExecBuilder(builder Builder) (sql.Result, error) ExecRaw(query string, args ...interface{}) (sql.Result, error) NamedQuery(query string, arg interface{}) (*sqlx.Rows, error) QueryRowX(query string, args ...interface{}) *sqlx.Row QueryX(query string, args ...interface{}) (*sqlx.Rows, error) Select(dest interface{}, query string, args ...interface{}) error + SelectBuilder(dest interface{}, builder Builder) error } // namedParamRegex is used to capture all named parameters and convert them @@ -105,6 +112,15 @@ func (w *sqlxDBWrapper) Get(dest interface{}, query string, args ...interface{}) return w.DB.GetContext(ctx, dest, query, args...) } +func (w *sqlxDBWrapper) GetBuilder(dest interface{}, builder Builder) error { + query, args, err := builder.ToSql() + if err != nil { + return err + } + + return w.Get(dest, query, args...) +} + func (w *sqlxDBWrapper) NamedExec(query string, arg interface{}) (sql.Result, error) { if w.DB.DriverName() == model.DatabaseDriverPostgres { query = namedParamRegex.ReplaceAllStringFunc(query, strings.ToLower) @@ -127,6 +143,15 @@ func (w *sqlxDBWrapper) Exec(query string, args ...interface{}) (sql.Result, err return w.ExecRaw(query, args...) } +func (w *sqlxDBWrapper) ExecBuilder(builder Builder) (sql.Result, error) { + query, args, err := builder.ToSql() + if err != nil { + return nil, err + } + + return w.Exec(query, args...) +} + func (w *sqlxDBWrapper) ExecNoTimeout(query string, args ...interface{}) (sql.Result, error) { query = w.DB.Rebind(query) @@ -212,6 +237,15 @@ func (w *sqlxDBWrapper) Select(dest interface{}, query string, args ...interface return w.DB.SelectContext(ctx, dest, query, args...) } +func (w *sqlxDBWrapper) SelectBuilder(dest interface{}, builder Builder) error { + query, args, err := builder.ToSql() + if err != nil { + return err + } + + return w.Select(dest, query, args...) +} + type sqlxTxWrapper struct { *sqlx.Tx queryTimeout time.Duration @@ -240,6 +274,15 @@ func (w *sqlxTxWrapper) Get(dest interface{}, query string, args ...interface{}) return w.Tx.GetContext(ctx, dest, query, args...) } +func (w *sqlxTxWrapper) GetBuilder(dest interface{}, builder Builder) error { + query, args, err := builder.ToSql() + if err != nil { + return err + } + + return w.Get(dest, query, args...) +} + func (w *sqlxTxWrapper) Exec(query string, args ...interface{}) (sql.Result, error) { query = w.Tx.Rebind(query) @@ -258,6 +301,15 @@ func (w *sqlxTxWrapper) ExecNoTimeout(query string, args ...interface{}) (sql.Re return w.Tx.ExecContext(context.Background(), query, args...) } +func (w *sqlxTxWrapper) ExecBuilder(builder Builder) (sql.Result, error) { + query, args, err := builder.ToSql() + if err != nil { + return nil, err + } + + return w.Exec(query, args...) +} + // ExecRaw is like Exec but without any rebinding of params. You need to pass // the exact param types of your target database. func (w *sqlxTxWrapper) ExecRaw(query string, args ...interface{}) (sql.Result, error) { @@ -375,6 +427,15 @@ func (w *sqlxTxWrapper) Select(dest interface{}, query string, args ...interface return w.Tx.SelectContext(ctx, dest, query, args...) } +func (w *sqlxTxWrapper) SelectBuilder(dest interface{}, builder Builder) error { + query, args, err := builder.ToSql() + if err != nil { + return err + } + + return w.Select(dest, query, args...) +} + func removeSpace(r rune) rune { // Strip everything except ' ' // This also strips out more than one space, diff --git a/store/sqlstore/thread_store.go b/store/sqlstore/thread_store.go index 2147701df2..9d34b0ae5d 100644 --- a/store/sqlstore/thread_store.go +++ b/store/sqlstore/thread_store.go @@ -69,14 +69,10 @@ func (s *SqlThreadStore) initializeQueries() { func (s *SqlThreadStore) Get(id string) (*model.Thread, error) { var thread model.Thread - query, args, err := s.threadsSelectQuery. - Where(sq.Eq{"PostId": id}). - ToSql() - if err != nil { - return nil, errors.Wrap(err, "thread_tosql") - } + query := s.threadsSelectQuery. + Where(sq.Eq{"PostId": id}) - err = s.GetReplicaX().Get(&thread, query, args...) + err := s.GetReplicaX().GetBuilder(&thread, query) if err != nil { if err == sql.ErrNoRows { return nil, nil @@ -119,13 +115,8 @@ func (s *SqlThreadStore) GetTotalUnreadThreads(userId, teamId string, opts model query := s.getTotalThreadsQuery(userId, teamId, opts). Where(sq.Expr("ThreadMemberships.LastViewed < Threads.LastReplyAt")) - sql, args, err := query.ToSql() - if err != nil { - return 0, errors.Wrapf(err, "failed to build query to count unread threads for user id=%s", userId) - } - var totalUnreadThreads int64 - err = s.GetReplicaX().Get(&totalUnreadThreads, sql, args...) + err := s.GetReplicaX().GetBuilder(&totalUnreadThreads, query) if err != nil { return 0, errors.Wrapf(err, "failed to count unread threads for user id=%s", userId) } @@ -142,13 +133,8 @@ func (s *SqlThreadStore) GetTotalThreads(userId, teamId string, opts model.GetUs query := s.getTotalThreadsQuery(userId, teamId, opts) - sql, args, err := query.ToSql() - if err != nil { - return 0, errors.Wrapf(err, "failed to build query to count threads for user id=%s", userId) - } - var totalThreads int64 - err = s.GetReplicaX().Get(&totalThreads, sql, args...) + err := s.GetReplicaX().GetBuilder(&totalThreads, query) if err != nil { return 0, errors.Wrapf(err, "failed to count threads for user id=%s", userId) } @@ -183,12 +169,7 @@ func (s *SqlThreadStore) GetTotalUnreadMentions(userId, teamId string, opts mode query = query.Where(sq.Eq{"COALESCE(Threads.ThreadDeleteAt, 0)": 0}) } - sql, args, err := query.ToSql() - if err != nil { - return 0, errors.Wrapf(err, "failed to build query to count unread mentions for user id=%s", userId) - } - - err = s.GetReplicaX().Get(&totalUnreadMentions, sql, args...) + err := s.GetReplicaX().GetBuilder(&totalUnreadMentions, query) if err != nil { return 0, errors.Wrapf(err, "failed to count unread mentions for user id=%s", userId) } @@ -224,18 +205,13 @@ func (s *SqlThreadStore) GetThreadsForUser(userId, teamId string, opts model.Get unreadRepliesQuery = unreadRepliesQuery.Where(sq.Eq{"Posts.DeleteAt": 0}) } - unreadRepliesSql, unreadRepliesArgs, err := unreadRepliesQuery.ToSql() - if err != nil { - return nil, errors.Wrapf(err, "failed to build subquery to count unread replies when getting threads for user id=%s", userId) - } - query := s.threadsAndPostsSelectQuery. Column(postSliceCoalesceQuery()). Columns( "ThreadMemberships.LastViewed as LastViewedAt", "ThreadMemberships.UnreadMentions as UnreadMentions", ). - Column(sq.Alias(sq.Expr(unreadRepliesSql, unreadRepliesArgs...), "UnreadReplies")). + Column(sq.Alias(unreadRepliesQuery, "UnreadReplies")). Join("Posts ON Posts.Id = Threads.PostId"). Join("ThreadMemberships ON ThreadMemberships.PostId = Threads.PostId") @@ -282,12 +258,7 @@ func (s *SqlThreadStore) GetThreadsForUser(userId, teamId string, opts model.Get OrderBy("Threads.LastReplyAt " + order). Limit(pageSize) - sql, args, err := query.ToSql() - if err != nil { - return nil, errors.Wrapf(err, "failed to build query to fetch threads for user id=%s", userId) - } - - err = s.GetReplicaX().Select(&threads, sql, args...) + err := s.GetReplicaX().SelectBuilder(&threads, query) if err != nil { return nil, errors.Wrapf(err, "failed to fetch threads for user id=%s", userId) } @@ -373,21 +344,16 @@ func (s *SqlThreadStore) GetTeamsUnreadForUser(userID string, teamIDs []string) wg.Add(1) go func() { defer wg.Done() - repliesQuery, repliesQueryArgs, err := s.getQueryBuilder(). + repliesQuery := s.getQueryBuilder(). Select("COUNT(Threads.PostId) AS Count, TeamId"). From("Threads"). LeftJoin("ThreadMemberships ON Threads.PostId = ThreadMemberships.PostId"). LeftJoin("Channels ON Threads.ChannelId = Channels.Id"). Where(fetchConditions). Where("Threads.LastReplyAt > ThreadMemberships.LastViewed"). - GroupBy("Channels.TeamId"). - ToSql() - if err != nil { - err1 = errors.Wrap(err, "GetTotalUnreadThreads_Tosql") - return - } + GroupBy("Channels.TeamId") - err = s.GetReplicaX().Select(&unreadThreads, repliesQuery, repliesQueryArgs...) + err := s.GetReplicaX().SelectBuilder(&unreadThreads, repliesQuery) if err != nil { err1 = errors.Wrap(err, "failed to get total unread threads") } @@ -396,19 +362,15 @@ func (s *SqlThreadStore) GetTeamsUnreadForUser(userID string, teamIDs []string) wg.Add(1) go func() { defer wg.Done() - mentionsQuery, mentionsQueryArgs, err := s.getQueryBuilder(). + mentionsQuery := s.getQueryBuilder(). Select("COALESCE(SUM(ThreadMemberships.UnreadMentions),0) AS Count, TeamId"). From("ThreadMemberships"). LeftJoin("Threads ON Threads.PostId = ThreadMemberships.PostId"). LeftJoin("Channels ON Threads.ChannelId = Channels.Id"). Where(fetchConditions). - GroupBy("Channels.TeamId"). - ToSql() - if err != nil { - err2 = errors.Wrap(err, "GetTotalUnreadMentions_Tosql") - } + GroupBy("Channels.TeamId") - err = s.GetReplicaX().Select(&unreadMentions, mentionsQuery, mentionsQueryArgs...) + err := s.GetReplicaX().SelectBuilder(&unreadMentions, mentionsQuery) if err != nil { err2 = errors.Wrap(err, "failed to get total unread mentions") } @@ -459,16 +421,12 @@ func (s *SqlThreadStore) GetThreadFollowers(threadID string, fetchOnlyActive boo } } - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Select("ThreadMemberships.UserId"). From("ThreadMemberships"). - Where(fetchConditions). - ToSql() - if err != nil { - return nil, errors.Wrapf(err, "failed to build query to get thread followers for thread id=%s", threadID) - } + Where(fetchConditions) - err = s.GetReplicaX().Select(&users, query, args...) + err := s.GetReplicaX().SelectBuilder(&users, query) if err != nil { return nil, errors.Wrapf(err, "failed to get thread followers for thread id=%s", threadID) } @@ -494,17 +452,14 @@ func (s *SqlThreadStore) GetThreadForUser(teamId string, threadMembership *model model.Post } - unreadRepliesQuery, unreadRepliesArgs, err := sq. + unreadRepliesQuery := sq. Select("COUNT(Posts.Id)"). From("Posts"). Where(sq.And{ sq.Eq{"Posts.RootId": threadMembership.PostId}, sq.Gt{"Posts.CreateAt": threadMembership.LastViewed}, sq.Eq{"Posts.DeleteAt": 0}, - }).ToSql() - if err != nil { - return nil, errors.Wrapf(err, "failed to build subquery to count unread replies for getting thread for user id=%s, post id=%s", threadMembership.UserId, threadMembership.PostId) - } + }) fetchConditions := sq.And{ sq.Or{sq.Eq{"Channels.TeamId": teamId}, sq.Eq{"Channels.TeamId": ""}}, @@ -518,19 +473,13 @@ func (s *SqlThreadStore) GetThreadForUser(teamId string, threadMembership *model } var thread JoinedThread - querySQL, threadArgs, err := query. - Column(sq.Alias(sq.Expr(unreadRepliesQuery), "UnreadReplies")). + query = query. + Column(sq.Alias(unreadRepliesQuery, "UnreadReplies")). LeftJoin("Posts ON Posts.Id = Threads.PostId"). LeftJoin("Channels ON Posts.ChannelId = Channels.Id"). - Where(fetchConditions). - ToSql() - if err != nil { - return nil, errors.Wrapf(err, "failed to build query to get thread for user id=%s, post id=%s", threadMembership.UserId, threadMembership.PostId) - } + Where(fetchConditions) - args := append(unreadRepliesArgs, threadArgs...) - - err = s.GetReplicaX().Get(&thread, querySQL, args...) + err := s.GetReplicaX().GetBuilder(&thread, query) if err != nil { if err == sql.ErrNoRows { return nil, store.NewErrNotFound("Thread", threadMembership.PostId) @@ -609,12 +558,7 @@ func (s *SqlThreadStore) MarkAllAsReadByChannels(userID string, channelIDs []str Where(sq.Eq{"Threads.ChannelId": channelIDs}). Where(sq.Expr("Threads.LastReplyAt > ThreadMemberships.LastViewed")) - sql, args, err := query.ToSql() - if err != nil { - return errors.Wrapf(err, "failed to build query to mark all as read by %d channels for user id=%s", len(channelIDs), userID) - } - - if _, err := s.GetMasterX().Exec(sql, args...); err != nil { + if _, err := s.GetMasterX().ExecBuilder(query); err != nil { return errors.Wrapf(err, "failed to mark all threads as read by channels for user id=%s", userID) } @@ -624,19 +568,15 @@ func (s *SqlThreadStore) MarkAllAsReadByChannels(userID string, channelIDs []str func (s *SqlThreadStore) MarkAllAsRead(userId string, threadIds []string) error { timestamp := model.GetMillis() - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Update("ThreadMemberships"). Where(sq.Eq{"UserId": userId}). Where(sq.Eq{"PostId": threadIds}). Set("LastViewed", timestamp). Set("UnreadMentions", 0). - Set("LastUpdated", model.GetMillis()). - ToSql() - if err != nil { - return errors.Wrapf(err, "failed to build query to mark %d threads as read for user id=%s", len(threadIds), userId) - } + Set("LastUpdated", model.GetMillis()) - _, err = s.GetMasterX().Exec(query, args...) + _, err := s.GetMasterX().ExecBuilder(query) if err != nil { return errors.Wrapf(err, "failed to mark %d threads as read for user id=%s", len(threadIds), userId) } @@ -656,19 +596,15 @@ func (s *SqlThreadStore) MarkAllAsReadByTeam(userId, teamId string) error { membershipIds = append(membershipIds, m.PostId) } timestamp := model.GetMillis() - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Update("ThreadMemberships"). Where(sq.Eq{"PostId": membershipIds}). Where(sq.Eq{"UserId": userId}). Set("LastViewed", timestamp). Set("UnreadMentions", 0). - Set("LastUpdated", model.GetMillis()). - ToSql() - if err != nil { - return errors.Wrapf(err, "failed to build query to update thread read state for user id=%s", userId) - } + Set("LastUpdated", model.GetMillis()) - _, err = s.GetMasterX().Exec(query, args...) + _, err = s.GetMasterX().ExecBuilder(query) if err != nil { return errors.Wrapf(err, "failed to update thread read state for user id=%s", userId) } @@ -677,18 +613,14 @@ func (s *SqlThreadStore) MarkAllAsReadByTeam(userId, teamId string) error { // MarkAsRead marks the given thread for the given user as unread from the given timestamp. func (s *SqlThreadStore) MarkAsRead(userId, threadId string, timestamp int64) error { - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Update("ThreadMemberships"). Where(sq.Eq{"UserId": userId}). Where(sq.Eq{"PostId": threadId}). Set("LastViewed", timestamp). - Set("LastUpdated", model.GetMillis()). - ToSql() - if err != nil { - return errors.Wrapf(err, "failed to build query to update thread read state for user id=%s thread_id=%v", userId, threadId) - } + Set("LastUpdated", model.GetMillis()) - _, err = s.GetMasterX().Exec(query, args...) + _, err := s.GetMasterX().ExecBuilder(query) if err != nil { return errors.Wrapf(err, "failed to update thread read state for user id=%s thread_id=%v", userId, threadId) } @@ -696,16 +628,12 @@ func (s *SqlThreadStore) MarkAsRead(userId, threadId string, timestamp int64) er } func (s *SqlThreadStore) saveMembership(ex sqlxExecutor, membership *model.ThreadMembership) (*model.ThreadMembership, error) { - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Insert("ThreadMemberships"). Columns("PostId", "UserId", "Following", "LastViewed", "LastUpdated", "UnreadMentions"). - Values(membership.PostId, membership.UserId, membership.Following, membership.LastViewed, membership.LastUpdated, membership.UnreadMentions). - ToSql() - if err != nil { - return nil, errors.Wrapf(err, "failed to build query to save thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) - } + Values(membership.PostId, membership.UserId, membership.Following, membership.LastViewed, membership.LastUpdated, membership.UnreadMentions) - _, err = ex.Exec(query, args...) + _, err := ex.ExecBuilder(query) if err != nil { return nil, errors.Wrapf(err, "failed to save thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) } @@ -718,7 +646,7 @@ func (s *SqlThreadStore) UpdateMembership(membership *model.ThreadMembership) (* } func (s *SqlThreadStore) updateMembership(ex sqlxExecutor, membership *model.ThreadMembership) (*model.ThreadMembership, error) { - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Update("ThreadMemberships"). Set("Following", membership.Following). Set("LastViewed", membership.LastViewed). @@ -727,13 +655,9 @@ func (s *SqlThreadStore) updateMembership(ex sqlxExecutor, membership *model.Thr Where(sq.And{ sq.Eq{"PostId": membership.PostId}, sq.Eq{"UserId": membership.UserId}, - }). - ToSql() - if err != nil { - return nil, errors.Wrapf(err, "failed to build query to update thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) - } + }) - _, err = ex.Exec(query, args...) + _, err := ex.ExecBuilder(query) if err != nil { return nil, errors.Wrapf(err, "failed to update thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) } @@ -744,19 +668,15 @@ func (s *SqlThreadStore) updateMembership(ex sqlxExecutor, membership *model.Thr func (s *SqlThreadStore) GetMembershipsForUser(userId, teamId string) ([]*model.ThreadMembership, error) { memberships := []*model.ThreadMembership{} - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Select("ThreadMemberships.*"). Join("Threads ON Threads.PostId = ThreadMemberships.PostId"). Join("Channels ON Threads.ChannelId = Channels.Id"). From("ThreadMemberships"). Where(sq.Or{sq.Eq{"Channels.TeamId": teamId}, sq.Eq{"Channels.TeamId": ""}}). - Where(sq.Eq{"ThreadMemberships.UserId": userId}). - ToSql() - if err != nil { - return nil, errors.Wrapf(err, "failed to build query to get thread membership with userid=%s", userId) - } + Where(sq.Eq{"ThreadMemberships.UserId": userId}) - err = s.GetReplicaX().Select(&memberships, query, args...) + err := s.GetReplicaX().SelectBuilder(&memberships, query) if err != nil { return nil, errors.Wrapf(err, "failed to get thread membership with userid=%s", userId) } @@ -769,19 +689,15 @@ func (s *SqlThreadStore) GetMembershipForUser(userId, postId string) (*model.Thr func (s *SqlThreadStore) getMembershipForUser(ex sqlxExecutor, userId, postId string) (*model.ThreadMembership, error) { var membership model.ThreadMembership - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Select("*"). From("ThreadMemberships"). Where(sq.And{ sq.Eq{"PostId": postId}, sq.Eq{"UserId": userId}, - }). - ToSql() - if err != nil { - return nil, errors.Wrapf(err, "failed to build query to get thread membership with userid=%s postid=%s", userId, postId) - } + }) - err = ex.Get(&membership, query, args...) + err := ex.GetBuilder(&membership, query) if err != nil { if err == sql.ErrNoRows { return nil, store.NewErrNotFound("Thread", postId) @@ -793,18 +709,14 @@ func (s *SqlThreadStore) getMembershipForUser(ex sqlxExecutor, userId, postId st } func (s *SqlThreadStore) DeleteMembershipForUser(userId string, postId string) error { - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Delete("ThreadMemberships"). Where(sq.And{ sq.Eq{"PostId": postId}, sq.Eq{"UserId": userId}, - }). - ToSql() - if err != nil { - return errors.Wrap(err, "failed to build query to delete thread membership") - } + }) - _, err = s.GetMasterX().Exec(query, args...) + _, err := s.GetMasterX().ExecBuilder(query) if err != nil { return errors.Wrap(err, "failed to delete thread membership") } @@ -906,18 +818,15 @@ func (s *SqlThreadStore) MaintainMembership(userId, postId string, opts store.Th } func (s *SqlThreadStore) GetPosts(threadId string, since int64) ([]*model.Post, error) { - query, args, err := s.getQueryBuilder(). + query := s.getQueryBuilder(). Select("*"). From("Posts"). Where(sq.Eq{"RootId": threadId}). Where(sq.Eq{"DeleteAt": 0}). - Where(sq.GtOrEq{"CreateAt": since}).ToSql() - if err != nil { - return nil, errors.Wrap(err, "failed to build query to fetch thread posts") - } + Where(sq.GtOrEq{"CreateAt": since}) result := []*model.Post{} - err = s.GetReplicaX().Select(&result, query, args...) + err := s.GetReplicaX().SelectBuilder(&result, query) if err != nil { return nil, errors.Wrap(err, "failed to fetch thread posts") } @@ -1008,23 +917,21 @@ func (s *SqlThreadStore) DeleteOrphanedRows(limit int) (deleted int64, err error } // return number of unread replies for a single thread -func (s *SqlThreadStore) GetThreadUnreadReplyCount(threadMembership *model.ThreadMembership) (unreadReplies int64, err error) { - query, args, err := s.getQueryBuilder(). +func (s *SqlThreadStore) GetThreadUnreadReplyCount(threadMembership *model.ThreadMembership) (int64, error) { + query := s.getQueryBuilder(). Select("COUNT(Posts.Id)"). From("Posts"). Where(sq.And{ sq.Eq{"Posts.RootId": threadMembership.PostId}, sq.Gt{"Posts.CreateAt": threadMembership.LastViewed}, sq.Eq{"Posts.DeleteAt": 0}, - }).ToSql() - if err != nil { - return 0, errors.Wrapf(err, "failed to build query to count unread reply count for post id=%s", threadMembership.PostId) - } + }) - err = s.GetReplicaX().Get(&unreadReplies, query, args...) + var unreadReplies int64 + err := s.GetReplicaX().GetBuilder(&unreadReplies, query) if err != nil { return 0, errors.Wrapf(err, "failed to count unread reply count for post id=%s", threadMembership.PostId) } - return + return unreadReplies, nil }