diff --git a/go.mod b/go.mod index 3fe2dd45bb..ac2f4b5b92 100644 --- a/go.mod +++ b/go.mod @@ -126,6 +126,8 @@ require ( gopkg.in/yaml.v2 v2.4.0 ) +replace github.com/Masterminds/squirrel v1.5.2 => github.com/lieut-data/squirrel v1.5.4 + // Hack to prevent the willf/bitset module from being upgraded to 1.2.0. // They changed the module path from github.com/willf/bitset to // github.com/bits-and-blooms/bitset and a couple of dependent repos are yet diff --git a/go.sum b/go.sum index 4fa04fd124..4c6f922b0c 100644 --- a/go.sum +++ b/go.sum @@ -978,6 +978,10 @@ github.com/lib/pq v1.8.0/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/lib/pq v1.10.0/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/lib/pq v1.10.4 h1:SO9z7FRPzA03QhHKJrH5BXA6HU1rS4V2nIVrrNC1iYk= github.com/lib/pq v1.10.4/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= +github.com/lieut-data/squirrel v1.5.3 h1:c6RI29VQOkUvqjpUsa45HeZ+wShtzDcSrmTUom/F65g= +github.com/lieut-data/squirrel v1.5.3/go.mod h1:NNaOrjSoIDfDA40n7sr2tPNZRfjzjA400rg+riTZj10= +github.com/lieut-data/squirrel v1.5.4 h1:OGzJNl0/ZxdjLEHuFzDo797zB2V7i8wQXBVThcOzbHE= +github.com/lieut-data/squirrel v1.5.4/go.mod h1:NNaOrjSoIDfDA40n7sr2tPNZRfjzjA400rg+riTZj10= github.com/lunixbochs/vtclean v1.0.0/go.mod h1:pHhQNgMf3btfWnGBVipUOjRYhoOsdGqdm/+2c2E2WMI= github.com/lyft/protoc-gen-star v0.5.3/go.mod h1:V0xaHgaf5oCCqmcxYcWiDfTiKsZsRc87/1qhoTACD8w= github.com/magiconair/properties v1.8.0/go.mod h1:PppfXfuXeibc/6YijjN8zIbojt8czPbwD3XqdrwzmxQ= diff --git a/store/sqlstore/thread_store.go b/store/sqlstore/thread_store.go index d926b8accc..8795d96e2b 100644 --- a/store/sqlstore/thread_store.go +++ b/store/sqlstore/thread_store.go @@ -571,36 +571,28 @@ func (s *SqlThreadStore) MarkAllAsReadByChannels(userID string, channelIDs []str now := model.GetMillis() - // TODO: Fork squirrel to include https://github.com/Masterminds/squirrel/pull/256 and - // support FROM in an UPDATE query. - channelIDsSql, channelIDsArgs := constructArrayArgs(channelIDs) - - var query string + var query sq.UpdateBuilder if s.DriverName() == model.DatabaseDriverPostgres { - query = ` - UPDATE ThreadMemberships - SET LastViewed = ?, UnreadMentions = ?, LastUpdated = ? - FROM Threads - WHERE ThreadMemberships.UserId = ? - AND Threads.PostId = ThreadMemberships.PostId - AND Threads.ChannelID IN ` + channelIDsSql + ` - AND Threads.LastReplyAt > ThreadMemberships.LastViewed - ` + query = s.getQueryBuilder().Update("ThreadMemberships").From("Threads") + } else { - query = ` - UPDATE ThreadMemberships, Threads - SET ThreadMemberships.LastViewed = ?, ThreadMemberships.UnreadMentions = ?, ThreadMemberships.LastUpdated = ? - WHERE ThreadMemberships.UserId = ? - AND Threads.PostId = ThreadMemberships.PostId - AND Threads.ChannelID IN ` + channelIDsSql + ` - AND Threads.LastReplyAt > ThreadMemberships.LastViewed - ` + query = s.getQueryBuilder().Update("ThreadMemberships", "Threads") } - args := []interface{}{now, 0, now, userID} - args = append(args, channelIDsArgs...) + query = query.Set("LastViewed", now). + Set("UnreadMentions", 0). + Set("LastUpdated", now). + Where(sq.Eq{"ThreadMemberships.UserId": userID}). + Where(sq.Expr("Threads.PostId = ThreadMemberships.PostId")). + Where(sq.Eq{"Threads.ChannelId": channelIDs}). + Where(sq.Expr("Threads.LastReplyAt > ThreadMemberships.LastViewed")) - if _, err := s.GetMasterX().Exec(query, args...); err != nil { + 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 { return errors.Wrapf(err, "failed to mark all threads as read by channels for user id=%s", userID) } diff --git a/vendor/github.com/Masterminds/squirrel/squirrel_ctx.go b/vendor/github.com/Masterminds/squirrel/squirrel_ctx.go index 504e763d2e..c20148ad33 100644 --- a/vendor/github.com/Masterminds/squirrel/squirrel_ctx.go +++ b/vendor/github.com/Masterminds/squirrel/squirrel_ctx.go @@ -32,7 +32,7 @@ type QueryRowerContext interface { QueryRowContext(ctx context.Context, query string, args ...interface{}) RowScanner } -// RunnerContext groups the Runner interface, along with the Contect versions of each of +// RunnerContext groups the Runner interface, along with the Context versions of each of // its methods type RunnerContext interface { Runner diff --git a/vendor/github.com/Masterminds/squirrel/statement.go b/vendor/github.com/Masterminds/squirrel/statement.go index 9420c67f8e..1c481be25d 100644 --- a/vendor/github.com/Masterminds/squirrel/statement.go +++ b/vendor/github.com/Masterminds/squirrel/statement.go @@ -22,8 +22,8 @@ func (b StatementBuilderType) Replace(into string) InsertBuilder { } // Update returns a UpdateBuilder for this StatementBuilderType. -func (b StatementBuilderType) Update(table string) UpdateBuilder { - return UpdateBuilder(b).Table(table) +func (b StatementBuilderType) Update(tables ...string) UpdateBuilder { + return UpdateBuilder(b).Table(tables...) } // Delete returns a DeleteBuilder for this StatementBuilderType. @@ -76,8 +76,8 @@ func Replace(into string) InsertBuilder { // Update returns a new UpdateBuilder with the given table name. // // See UpdateBuilder.Table. -func Update(table string) UpdateBuilder { - return StatementBuilder.Update(table) +func Update(tables ...string) UpdateBuilder { + return StatementBuilder.Update(tables...) } // Delete returns a new DeleteBuilder with the given table name. diff --git a/vendor/github.com/Masterminds/squirrel/update.go b/vendor/github.com/Masterminds/squirrel/update.go index 8d658d7219..86f9c9aa9c 100644 --- a/vendor/github.com/Masterminds/squirrel/update.go +++ b/vendor/github.com/Masterminds/squirrel/update.go @@ -14,8 +14,9 @@ type updateData struct { PlaceholderFormat PlaceholderFormat RunWith BaseRunner Prefixes []Sqlizer - Table string + Tables []string SetClauses []setClause + From []Sqlizer WhereParts []Sqlizer OrderBys []string Limit string @@ -54,7 +55,7 @@ func (d *updateData) QueryRow() RowScanner { } func (d *updateData) ToSql() (sqlStr string, args []interface{}, err error) { - if len(d.Table) == 0 { + if len(d.Tables) == 0 { err = fmt.Errorf("update statements must specify a table") return } @@ -75,7 +76,7 @@ func (d *updateData) ToSql() (sqlStr string, args []interface{}, err error) { } sql.WriteString("UPDATE ") - sql.WriteString(d.Table) + sql.WriteString(strings.Join(d.Tables, ", ")) sql.WriteString(" SET ") setSqls := make([]string, len(d.SetClauses)) @@ -100,6 +101,14 @@ func (d *updateData) ToSql() (sqlStr string, args []interface{}, err error) { } sql.WriteString(strings.Join(setSqls, ", ")) + if len(d.From) > 0 { + sql.WriteString(" FROM ") + args, err = appendToSql(d.From, sql, ", ", args) + if err != nil { + return + } + } + if len(d.WhereParts) > 0 { sql.WriteString(" WHERE ") args, err = appendToSql(d.WhereParts, sql, " AND ", args) @@ -208,8 +217,16 @@ func (b UpdateBuilder) PrefixExpr(expr Sqlizer) UpdateBuilder { } // Table sets the table to be updated. -func (b UpdateBuilder) Table(table string) UpdateBuilder { - return builder.Set(b, "Table", table).(UpdateBuilder) +// Additional tables are used with supporting databases to implicitly join. +func (b UpdateBuilder) Table(tables ...string) UpdateBuilder { + nonEmptyTables := make([]string, 0, len(tables)) + for _, table := range tables { + if table != "" { + nonEmptyTables = append(nonEmptyTables, table) + } + } + + return builder.Set(b, "Tables", nonEmptyTables).(UpdateBuilder) } // Set adds SET clauses to the query. @@ -233,6 +250,19 @@ func (b UpdateBuilder) SetMap(clauses map[string]interface{}) UpdateBuilder { return b } +// From adds FROM clause to the query +// FROM is valid construct in postgresql only. +func (b UpdateBuilder) From(from string) UpdateBuilder { + return builder.Append(b, "From", newPart(from)).(UpdateBuilder) +} + +// FromSelect sets a subquery into the FROM clause of the query. +func (b UpdateBuilder) FromSelect(from SelectBuilder, alias string) UpdateBuilder { + // Prevent misnumbered parameters in nested selects (#183). + from = from.PlaceholderFormat(Question) + return builder.Append(b, "From", Alias(from, alias)).(UpdateBuilder) +} + // Where adds WHERE expressions to the query. // // See SelectBuilder.Where for more information. diff --git a/vendor/modules.txt b/vendor/modules.txt index b952b6b26a..4de552f5ce 100644 --- a/vendor/modules.txt +++ b/vendor/modules.txt @@ -10,7 +10,7 @@ github.com/JalfResi/justext # github.com/Masterminds/semver/v3 v3.1.1 ## explicit github.com/Masterminds/semver/v3 -# github.com/Masterminds/squirrel v1.5.2 +# github.com/Masterminds/squirrel v1.5.2 => github.com/lieut-data/squirrel v1.5.4 ## explicit github.com/Masterminds/squirrel # github.com/PuerkitoBio/goquery v1.8.0