diff --git a/server/channels/db/migrations/mysql/000027_create_status.up.sql b/server/channels/db/migrations/mysql/000027_create_status.up.sql index 0e6757ff05..7133628559 100644 --- a/server/channels/db/migrations/mysql/000027_create_status.up.sql +++ b/server/channels/db/migrations/mysql/000027_create_status.up.sql @@ -1,7 +1,7 @@ CREATE TABLE IF NOT EXISTS Status ( UserId varchar(26) NOT NULL, Status varchar(32) DEFAULT NULL, - Manual tinyint(1) DEFAULT NULL, + `Manual` tinyint(1) DEFAULT NULL, LastActivityAt bigint(20) DEFAULT NULL, PRIMARY KEY (UserId) ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4; diff --git a/server/channels/store/sqlstore/status_store.go b/server/channels/store/sqlstore/status_store.go index 590e5e92ee..e315fb7ce7 100644 --- a/server/channels/store/sqlstore/status_store.go +++ b/server/channels/store/sqlstore/status_store.go @@ -26,11 +26,11 @@ func newSqlStatusStore(sqlStore *SqlStore) store.StatusStore { func (s SqlStatusStore) SaveOrUpdate(st *model.Status) error { query := s.getQueryBuilder(). Insert("Status"). - Columns("UserId", "Status", "Manual", "LastActivityAt", "DNDEndTime", "PrevStatus"). + Columns("UserId", "Status", quoteColumnName(s.DriverName(), "Manual"), "LastActivityAt", "DNDEndTime", "PrevStatus"). Values(st.UserId, st.Status, st.Manual, st.LastActivityAt, st.DNDEndTime, st.PrevStatus) if s.DriverName() == model.DatabaseDriverMysql { - query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE Status = ?, Manual = ?, LastActivityAt = ?, DNDEndTime = ?, PrevStatus = ?", + query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE Status = ?, `Manual` = ?, LastActivityAt = ?, DNDEndTime = ?, PrevStatus = ?", st.Status, st.Manual, st.LastActivityAt, st.DNDEndTime, st.PrevStatus)) } else { query = query.SuffixExpr(sq.Expr("ON CONFLICT (userid) DO UPDATE SET Status = ?, Manual = ?, LastActivityAt = ?, DNDEndTime = ?, PrevStatus = ?", @@ -63,7 +63,7 @@ func (s SqlStatusStore) Get(userId string) (*model.Status, error) { func (s SqlStatusStore) GetByIds(userIds []string) ([]*model.Status, error) { query := s.getQueryBuilder(). - Select("UserId, Status, Manual, LastActivityAt"). + Select(fmt.Sprintf("UserId, Status, %s, LastActivityAt", quoteColumnName(s.DriverName(), "Manual"))). From("Status"). Where(sq.Eq{"UserId": userIds}) queryString, args, err := query.ToSql() @@ -123,7 +123,7 @@ func (s SqlStatusStore) updateExpiredStatuses(t *sqlxTxWrapper) ([]*model.Status Set("Status", sq.Expr("PrevStatus")). Set("PrevStatus", model.StatusDnd). Set("DNDEndTime", 0). - Set("Manual", false). + Set(quoteColumnName(s.DriverName(), "Manual"), false). ToSql() if err != nil { @@ -174,7 +174,7 @@ func (s SqlStatusStore) UpdateExpiredDNDStatuses() (_ []*model.Status, err error Set("Status", sq.Expr("PrevStatus")). Set("PrevStatus", model.StatusDnd). Set("DNDEndTime", 0). - Set("Manual", false). + Set(quoteColumnName(s.DriverName(), "Manual"), false). Suffix("RETURNING *"). ToSql() @@ -204,7 +204,7 @@ func (s SqlStatusStore) UpdateExpiredDNDStatuses() (_ []*model.Status, err error } func (s SqlStatusStore) ResetAll() error { - if _, err := s.GetMasterX().Exec("UPDATE Status SET Status = ? WHERE Manual = false", model.StatusOffline); err != nil { + if _, err := s.GetMasterX().Exec(fmt.Sprintf("UPDATE Status SET Status = ? WHERE %s = false", quoteColumnName(s.DriverName(), "Manual")), model.StatusOffline); err != nil { return errors.Wrap(err, "failed to update Statuses") } return nil diff --git a/server/channels/store/sqlstore/utils.go b/server/channels/store/sqlstore/utils.go index 75357f4c8a..0aa87cb7af 100644 --- a/server/channels/store/sqlstore/utils.go +++ b/server/channels/store/sqlstore/utils.go @@ -5,6 +5,7 @@ package sqlstore import ( "database/sql" + "fmt" "io" "net/url" "strconv" @@ -197,3 +198,16 @@ func maxInt64(a, b int64) int64 { } return b } + +// Adds backtiks to the column name for MySQL, this is required if +// the column name is a reserved keyword. +// +// `ColumnName` - MySQL +// ColumnName - Postgres +func quoteColumnName(driver string, columnName string) string { + if driver == model.DatabaseDriverMysql { + return fmt.Sprintf("`%s`", columnName) + } + + return columnName +}