Introduce (Get|Select|Exec)Builder (#20029)

Этот коммит содержится в:
Jesse Hallam
2022-05-16 14:48:21 -03:00
коммит произвёл GitHub
родитель 8802c6c9d6
Коммит 9c851e996c
2 изменённых файлов: 118 добавлений и 150 удалений

Просмотреть файл

@@ -35,17 +35,24 @@ func (w *StoreTestWrapper) DriverName() string {
return w.orig.DriverName() 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 // sqlxExecutor exposes sqlx operations. It is used to enable some internal store methods to
// accept both transactions (*sqlxTxWrapper) and common db handlers (*sqlxDbWrapper). // accept both transactions (*sqlxTxWrapper) and common db handlers (*sqlxDbWrapper).
type sqlxExecutor interface { type sqlxExecutor interface {
Get(dest interface{}, query string, args ...interface{}) error Get(dest interface{}, query string, args ...interface{}) error
GetBuilder(dest interface{}, builder Builder) error
NamedExec(query string, arg interface{}) (sql.Result, error) NamedExec(query string, arg interface{}) (sql.Result, error)
Exec(query string, args ...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) ExecRaw(query string, args ...interface{}) (sql.Result, error)
NamedQuery(query string, arg interface{}) (*sqlx.Rows, error) NamedQuery(query string, arg interface{}) (*sqlx.Rows, error)
QueryRowX(query string, args ...interface{}) *sqlx.Row QueryRowX(query string, args ...interface{}) *sqlx.Row
QueryX(query string, args ...interface{}) (*sqlx.Rows, error) QueryX(query string, args ...interface{}) (*sqlx.Rows, error)
Select(dest interface{}, query string, args ...interface{}) 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 // 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...) 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) { func (w *sqlxDBWrapper) NamedExec(query string, arg interface{}) (sql.Result, error) {
if w.DB.DriverName() == model.DatabaseDriverPostgres { if w.DB.DriverName() == model.DatabaseDriverPostgres {
query = namedParamRegex.ReplaceAllStringFunc(query, strings.ToLower) 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...) 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) { func (w *sqlxDBWrapper) ExecNoTimeout(query string, args ...interface{}) (sql.Result, error) {
query = w.DB.Rebind(query) 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...) 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 { type sqlxTxWrapper struct {
*sqlx.Tx *sqlx.Tx
queryTimeout time.Duration 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...) 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) { func (w *sqlxTxWrapper) Exec(query string, args ...interface{}) (sql.Result, error) {
query = w.Tx.Rebind(query) 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...) 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 // ExecRaw is like Exec but without any rebinding of params. You need to pass
// the exact param types of your target database. // the exact param types of your target database.
func (w *sqlxTxWrapper) ExecRaw(query string, args ...interface{}) (sql.Result, error) { 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...) 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 { func removeSpace(r rune) rune {
// Strip everything except ' ' // Strip everything except ' '
// This also strips out more than one space, // This also strips out more than one space,

Просмотреть файл

@@ -69,14 +69,10 @@ func (s *SqlThreadStore) initializeQueries() {
func (s *SqlThreadStore) Get(id string) (*model.Thread, error) { func (s *SqlThreadStore) Get(id string) (*model.Thread, error) {
var thread model.Thread var thread model.Thread
query, args, err := s.threadsSelectQuery. query := s.threadsSelectQuery.
Where(sq.Eq{"PostId": id}). Where(sq.Eq{"PostId": id})
ToSql()
if err != nil {
return nil, errors.Wrap(err, "thread_tosql")
}
err = s.GetReplicaX().Get(&thread, query, args...) err := s.GetReplicaX().GetBuilder(&thread, query)
if err != nil { if err != nil {
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
return nil, nil return nil, nil
@@ -119,13 +115,8 @@ func (s *SqlThreadStore) GetTotalUnreadThreads(userId, teamId string, opts model
query := s.getTotalThreadsQuery(userId, teamId, opts). query := s.getTotalThreadsQuery(userId, teamId, opts).
Where(sq.Expr("ThreadMemberships.LastViewed < Threads.LastReplyAt")) 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 var totalUnreadThreads int64
err = s.GetReplicaX().Get(&totalUnreadThreads, sql, args...) err := s.GetReplicaX().GetBuilder(&totalUnreadThreads, query)
if err != nil { if err != nil {
return 0, errors.Wrapf(err, "failed to count unread threads for user id=%s", userId) 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) 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 var totalThreads int64
err = s.GetReplicaX().Get(&totalThreads, sql, args...) err := s.GetReplicaX().GetBuilder(&totalThreads, query)
if err != nil { if err != nil {
return 0, errors.Wrapf(err, "failed to count threads for user id=%s", userId) 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}) query = query.Where(sq.Eq{"COALESCE(Threads.ThreadDeleteAt, 0)": 0})
} }
sql, args, err := query.ToSql() err := s.GetReplicaX().GetBuilder(&totalUnreadMentions, query)
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...)
if err != nil { if err != nil {
return 0, errors.Wrapf(err, "failed to count unread mentions for user id=%s", userId) 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}) 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. query := s.threadsAndPostsSelectQuery.
Column(postSliceCoalesceQuery()). Column(postSliceCoalesceQuery()).
Columns( Columns(
"ThreadMemberships.LastViewed as LastViewedAt", "ThreadMemberships.LastViewed as LastViewedAt",
"ThreadMemberships.UnreadMentions as UnreadMentions", "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("Posts ON Posts.Id = Threads.PostId").
Join("ThreadMemberships ON ThreadMemberships.PostId = 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). OrderBy("Threads.LastReplyAt " + order).
Limit(pageSize) Limit(pageSize)
sql, args, err := query.ToSql() err := s.GetReplicaX().SelectBuilder(&threads, query)
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...)
if err != nil { if err != nil {
return nil, errors.Wrapf(err, "failed to fetch threads for user id=%s", userId) 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) wg.Add(1)
go func() { go func() {
defer wg.Done() defer wg.Done()
repliesQuery, repliesQueryArgs, err := s.getQueryBuilder(). repliesQuery := s.getQueryBuilder().
Select("COUNT(Threads.PostId) AS Count, TeamId"). Select("COUNT(Threads.PostId) AS Count, TeamId").
From("Threads"). From("Threads").
LeftJoin("ThreadMemberships ON Threads.PostId = ThreadMemberships.PostId"). LeftJoin("ThreadMemberships ON Threads.PostId = ThreadMemberships.PostId").
LeftJoin("Channels ON Threads.ChannelId = Channels.Id"). LeftJoin("Channels ON Threads.ChannelId = Channels.Id").
Where(fetchConditions). Where(fetchConditions).
Where("Threads.LastReplyAt > ThreadMemberships.LastViewed"). Where("Threads.LastReplyAt > ThreadMemberships.LastViewed").
GroupBy("Channels.TeamId"). GroupBy("Channels.TeamId")
ToSql()
if err != nil {
err1 = errors.Wrap(err, "GetTotalUnreadThreads_Tosql")
return
}
err = s.GetReplicaX().Select(&unreadThreads, repliesQuery, repliesQueryArgs...) err := s.GetReplicaX().SelectBuilder(&unreadThreads, repliesQuery)
if err != nil { if err != nil {
err1 = errors.Wrap(err, "failed to get total unread threads") 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) wg.Add(1)
go func() { go func() {
defer wg.Done() defer wg.Done()
mentionsQuery, mentionsQueryArgs, err := s.getQueryBuilder(). mentionsQuery := s.getQueryBuilder().
Select("COALESCE(SUM(ThreadMemberships.UnreadMentions),0) AS Count, TeamId"). Select("COALESCE(SUM(ThreadMemberships.UnreadMentions),0) AS Count, TeamId").
From("ThreadMemberships"). From("ThreadMemberships").
LeftJoin("Threads ON Threads.PostId = ThreadMemberships.PostId"). LeftJoin("Threads ON Threads.PostId = ThreadMemberships.PostId").
LeftJoin("Channels ON Threads.ChannelId = Channels.Id"). LeftJoin("Channels ON Threads.ChannelId = Channels.Id").
Where(fetchConditions). Where(fetchConditions).
GroupBy("Channels.TeamId"). GroupBy("Channels.TeamId")
ToSql()
if err != nil {
err2 = errors.Wrap(err, "GetTotalUnreadMentions_Tosql")
}
err = s.GetReplicaX().Select(&unreadMentions, mentionsQuery, mentionsQueryArgs...) err := s.GetReplicaX().SelectBuilder(&unreadMentions, mentionsQuery)
if err != nil { if err != nil {
err2 = errors.Wrap(err, "failed to get total unread mentions") 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"). Select("ThreadMemberships.UserId").
From("ThreadMemberships"). From("ThreadMemberships").
Where(fetchConditions). Where(fetchConditions)
ToSql()
if err != nil {
return nil, errors.Wrapf(err, "failed to build query to get thread followers for thread id=%s", threadID)
}
err = s.GetReplicaX().Select(&users, query, args...) err := s.GetReplicaX().SelectBuilder(&users, query)
if err != nil { if err != nil {
return nil, errors.Wrapf(err, "failed to get thread followers for thread id=%s", threadID) 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 model.Post
} }
unreadRepliesQuery, unreadRepliesArgs, err := sq. unreadRepliesQuery := sq.
Select("COUNT(Posts.Id)"). Select("COUNT(Posts.Id)").
From("Posts"). From("Posts").
Where(sq.And{ Where(sq.And{
sq.Eq{"Posts.RootId": threadMembership.PostId}, sq.Eq{"Posts.RootId": threadMembership.PostId},
sq.Gt{"Posts.CreateAt": threadMembership.LastViewed}, sq.Gt{"Posts.CreateAt": threadMembership.LastViewed},
sq.Eq{"Posts.DeleteAt": 0}, 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{ fetchConditions := sq.And{
sq.Or{sq.Eq{"Channels.TeamId": teamId}, sq.Eq{"Channels.TeamId": ""}}, 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 var thread JoinedThread
querySQL, threadArgs, err := query. query = query.
Column(sq.Alias(sq.Expr(unreadRepliesQuery), "UnreadReplies")). Column(sq.Alias(unreadRepliesQuery, "UnreadReplies")).
LeftJoin("Posts ON Posts.Id = Threads.PostId"). LeftJoin("Posts ON Posts.Id = Threads.PostId").
LeftJoin("Channels ON Posts.ChannelId = Channels.Id"). LeftJoin("Channels ON Posts.ChannelId = Channels.Id").
Where(fetchConditions). 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)
}
args := append(unreadRepliesArgs, threadArgs...) err := s.GetReplicaX().GetBuilder(&thread, query)
err = s.GetReplicaX().Get(&thread, querySQL, args...)
if err != nil { if err != nil {
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
return nil, store.NewErrNotFound("Thread", threadMembership.PostId) 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.Eq{"Threads.ChannelId": channelIDs}).
Where(sq.Expr("Threads.LastReplyAt > ThreadMemberships.LastViewed")) Where(sq.Expr("Threads.LastReplyAt > ThreadMemberships.LastViewed"))
sql, args, err := query.ToSql() if _, err := s.GetMasterX().ExecBuilder(query); err != nil {
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 {
return errors.Wrapf(err, "failed to mark all threads as read by channels for user id=%s", userID) 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 { func (s *SqlThreadStore) MarkAllAsRead(userId string, threadIds []string) error {
timestamp := model.GetMillis() timestamp := model.GetMillis()
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Update("ThreadMemberships"). Update("ThreadMemberships").
Where(sq.Eq{"UserId": userId}). Where(sq.Eq{"UserId": userId}).
Where(sq.Eq{"PostId": threadIds}). Where(sq.Eq{"PostId": threadIds}).
Set("LastViewed", timestamp). Set("LastViewed", timestamp).
Set("UnreadMentions", 0). Set("UnreadMentions", 0).
Set("LastUpdated", model.GetMillis()). 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)
}
_, err = s.GetMasterX().Exec(query, args...) _, err := s.GetMasterX().ExecBuilder(query)
if err != nil { if err != nil {
return errors.Wrapf(err, "failed to mark %d threads as read for user id=%s", len(threadIds), userId) 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) membershipIds = append(membershipIds, m.PostId)
} }
timestamp := model.GetMillis() timestamp := model.GetMillis()
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Update("ThreadMemberships"). Update("ThreadMemberships").
Where(sq.Eq{"PostId": membershipIds}). Where(sq.Eq{"PostId": membershipIds}).
Where(sq.Eq{"UserId": userId}). Where(sq.Eq{"UserId": userId}).
Set("LastViewed", timestamp). Set("LastViewed", timestamp).
Set("UnreadMentions", 0). Set("UnreadMentions", 0).
Set("LastUpdated", model.GetMillis()). 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)
}
_, err = s.GetMasterX().Exec(query, args...) _, err = s.GetMasterX().ExecBuilder(query)
if err != nil { if err != nil {
return errors.Wrapf(err, "failed to update thread read state for user id=%s", userId) 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. // 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 { func (s *SqlThreadStore) MarkAsRead(userId, threadId string, timestamp int64) error {
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Update("ThreadMemberships"). Update("ThreadMemberships").
Where(sq.Eq{"UserId": userId}). Where(sq.Eq{"UserId": userId}).
Where(sq.Eq{"PostId": threadId}). Where(sq.Eq{"PostId": threadId}).
Set("LastViewed", timestamp). Set("LastViewed", timestamp).
Set("LastUpdated", model.GetMillis()). 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)
}
_, err = s.GetMasterX().Exec(query, args...) _, err := s.GetMasterX().ExecBuilder(query)
if err != nil { if err != nil {
return errors.Wrapf(err, "failed to update thread read state for user id=%s thread_id=%v", userId, threadId) 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) { func (s *SqlThreadStore) saveMembership(ex sqlxExecutor, membership *model.ThreadMembership) (*model.ThreadMembership, error) {
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Insert("ThreadMemberships"). Insert("ThreadMemberships").
Columns("PostId", "UserId", "Following", "LastViewed", "LastUpdated", "UnreadMentions"). Columns("PostId", "UserId", "Following", "LastViewed", "LastUpdated", "UnreadMentions").
Values(membership.PostId, membership.UserId, membership.Following, membership.LastViewed, membership.LastUpdated, membership.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)
}
_, err = ex.Exec(query, args...) _, err := ex.ExecBuilder(query)
if err != nil { if err != nil {
return nil, errors.Wrapf(err, "failed to save thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) 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) { func (s *SqlThreadStore) updateMembership(ex sqlxExecutor, membership *model.ThreadMembership) (*model.ThreadMembership, error) {
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Update("ThreadMemberships"). Update("ThreadMemberships").
Set("Following", membership.Following). Set("Following", membership.Following).
Set("LastViewed", membership.LastViewed). Set("LastViewed", membership.LastViewed).
@@ -727,13 +655,9 @@ func (s *SqlThreadStore) updateMembership(ex sqlxExecutor, membership *model.Thr
Where(sq.And{ Where(sq.And{
sq.Eq{"PostId": membership.PostId}, sq.Eq{"PostId": membership.PostId},
sq.Eq{"UserId": membership.UserId}, 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 { if err != nil {
return nil, errors.Wrapf(err, "failed to update thread membership with postid=%s userid=%s", membership.PostId, membership.UserId) 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) { func (s *SqlThreadStore) GetMembershipsForUser(userId, teamId string) ([]*model.ThreadMembership, error) {
memberships := []*model.ThreadMembership{} memberships := []*model.ThreadMembership{}
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Select("ThreadMemberships.*"). Select("ThreadMemberships.*").
Join("Threads ON Threads.PostId = ThreadMemberships.PostId"). Join("Threads ON Threads.PostId = ThreadMemberships.PostId").
Join("Channels ON Threads.ChannelId = Channels.Id"). Join("Channels ON Threads.ChannelId = Channels.Id").
From("ThreadMemberships"). From("ThreadMemberships").
Where(sq.Or{sq.Eq{"Channels.TeamId": teamId}, sq.Eq{"Channels.TeamId": ""}}). Where(sq.Or{sq.Eq{"Channels.TeamId": teamId}, sq.Eq{"Channels.TeamId": ""}}).
Where(sq.Eq{"ThreadMemberships.UserId": userId}). 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)
}
err = s.GetReplicaX().Select(&memberships, query, args...) err := s.GetReplicaX().SelectBuilder(&memberships, query)
if err != nil { if err != nil {
return nil, errors.Wrapf(err, "failed to get thread membership with userid=%s", userId) 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) { func (s *SqlThreadStore) getMembershipForUser(ex sqlxExecutor, userId, postId string) (*model.ThreadMembership, error) {
var membership model.ThreadMembership var membership model.ThreadMembership
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Select("*"). Select("*").
From("ThreadMemberships"). From("ThreadMemberships").
Where(sq.And{ Where(sq.And{
sq.Eq{"PostId": postId}, sq.Eq{"PostId": postId},
sq.Eq{"UserId": userId}, 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 != nil {
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
return nil, store.NewErrNotFound("Thread", postId) 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 { func (s *SqlThreadStore) DeleteMembershipForUser(userId string, postId string) error {
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Delete("ThreadMemberships"). Delete("ThreadMemberships").
Where(sq.And{ Where(sq.And{
sq.Eq{"PostId": postId}, sq.Eq{"PostId": postId},
sq.Eq{"UserId": userId}, 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 { if err != nil {
return errors.Wrap(err, "failed to delete thread membership") 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) { func (s *SqlThreadStore) GetPosts(threadId string, since int64) ([]*model.Post, error) {
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Select("*"). Select("*").
From("Posts"). From("Posts").
Where(sq.Eq{"RootId": threadId}). Where(sq.Eq{"RootId": threadId}).
Where(sq.Eq{"DeleteAt": 0}). Where(sq.Eq{"DeleteAt": 0}).
Where(sq.GtOrEq{"CreateAt": since}).ToSql() Where(sq.GtOrEq{"CreateAt": since})
if err != nil {
return nil, errors.Wrap(err, "failed to build query to fetch thread posts")
}
result := []*model.Post{} result := []*model.Post{}
err = s.GetReplicaX().Select(&result, query, args...) err := s.GetReplicaX().SelectBuilder(&result, query)
if err != nil { if err != nil {
return nil, errors.Wrap(err, "failed to fetch thread posts") 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 // return number of unread replies for a single thread
func (s *SqlThreadStore) GetThreadUnreadReplyCount(threadMembership *model.ThreadMembership) (unreadReplies int64, err error) { func (s *SqlThreadStore) GetThreadUnreadReplyCount(threadMembership *model.ThreadMembership) (int64, error) {
query, args, err := s.getQueryBuilder(). query := s.getQueryBuilder().
Select("COUNT(Posts.Id)"). Select("COUNT(Posts.Id)").
From("Posts"). From("Posts").
Where(sq.And{ Where(sq.And{
sq.Eq{"Posts.RootId": threadMembership.PostId}, sq.Eq{"Posts.RootId": threadMembership.PostId},
sq.Gt{"Posts.CreateAt": threadMembership.LastViewed}, sq.Gt{"Posts.CreateAt": threadMembership.LastViewed},
sq.Eq{"Posts.DeleteAt": 0}, 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 { if err != nil {
return 0, errors.Wrapf(err, "failed to count unread reply count for post id=%s", threadMembership.PostId) return 0, errors.Wrapf(err, "failed to count unread reply count for post id=%s", threadMembership.PostId)
} }
return return unreadReplies, nil
} }