Mono repo -> Master (#22553)
Combines the following repositories into one: https://github.com/mattermost/mattermost-server https://github.com/mattermost/mattermost-webapp https://github.com/mattermost/focalboard https://github.com/mattermost/mattermost-plugin-playbooks
Этот коммит содержится в:
65
server/channels/store/sqlstore/adapters.go
Обычный файл
65
server/channels/store/sqlstore/adapters.go
Обычный файл
@@ -0,0 +1,65 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql/driver"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
)
|
||||
|
||||
type jsonArray []string
|
||||
|
||||
func (a jsonArray) Value() (driver.Value, error) {
|
||||
var out bytes.Buffer
|
||||
if err := out.WriteByte('['); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for i, item := range a {
|
||||
if _, err := out.WriteString(strconv.Quote(item)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Skip the last element.
|
||||
if i < len(a)-1 {
|
||||
if err := out.WriteByte(','); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
err := out.WriteByte(']')
|
||||
return out.Bytes(), err
|
||||
}
|
||||
|
||||
type jsonStringVal string
|
||||
|
||||
func (str jsonStringVal) Value() (driver.Value, error) {
|
||||
return strconv.Quote(string(str)), nil
|
||||
}
|
||||
|
||||
type jsonKeyPath string
|
||||
|
||||
func (str jsonKeyPath) Value() (driver.Value, error) {
|
||||
return "{" + string(str) + "}", nil
|
||||
}
|
||||
|
||||
type TraceOnAdapter struct{}
|
||||
|
||||
func (t *TraceOnAdapter) Printf(format string, v ...any) {
|
||||
originalString := fmt.Sprintf(format, v...)
|
||||
newString := strings.ReplaceAll(originalString, "\n", " ")
|
||||
newString = strings.ReplaceAll(newString, "\t", " ")
|
||||
newString = strings.ReplaceAll(newString, "\"", "")
|
||||
mlog.Debug(newString)
|
||||
}
|
||||
|
||||
type JSONSerializable interface {
|
||||
ToJSON() string
|
||||
}
|
||||
21
server/channels/store/sqlstore/adapters_test.go
Обычный файл
21
server/channels/store/sqlstore/adapters_test.go
Обычный файл
@@ -0,0 +1,21 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestJSONArray(t *testing.T) {
|
||||
input := []string{"a", "b"}
|
||||
|
||||
out, err := jsonArray(input).Value()
|
||||
require.NoError(t, err)
|
||||
outBuf, ok := out.([]byte)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, []byte(`["a","b"]`), outBuf)
|
||||
}
|
||||
68
server/channels/store/sqlstore/audit_store.go
Обычный файл
68
server/channels/store/sqlstore/audit_store.go
Обычный файл
@@ -0,0 +1,68 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlAuditStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlAuditStore(sqlStore *SqlStore) store.AuditStore {
|
||||
return &SqlAuditStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlAuditStore) Save(audit *model.Audit) error {
|
||||
audit.Id = model.NewId()
|
||||
audit.CreateAt = model.GetMillis()
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(`INSERT INTO Audits
|
||||
(Id, CreateAt, UserId, Action, ExtraInfo, IpAddress, SessionId)
|
||||
VALUES
|
||||
(:Id, :CreateAt, :UserId, :Action, :ExtraInfo, :IpAddress, :SessionId)`, audit); err != nil {
|
||||
return errors.Wrapf(err, "failed to save Audit with userId=%s and action=%s", audit.UserId, audit.Action)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlAuditStore) Get(userId string, offset int, limit int) (model.Audits, error) {
|
||||
if limit > 1000 {
|
||||
return nil, store.NewErrOutOfBounds(limit)
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Audits").
|
||||
OrderBy("CreateAt DESC").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset))
|
||||
|
||||
if userId != "" {
|
||||
query = query.Where(sq.Eq{"UserId": userId})
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "audits_tosql")
|
||||
}
|
||||
|
||||
var audits model.Audits
|
||||
if err := s.GetReplicaX().Select(&audits, queryString, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get Audit list for userId=%s", userId)
|
||||
}
|
||||
return audits, nil
|
||||
}
|
||||
|
||||
func (s SqlAuditStore) PermanentDeleteByUser(userId string) error {
|
||||
if _, err := s.GetMasterX().Exec("DELETE FROM Audits WHERE UserId = ?", userId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Audit with userId=%s", userId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/audit_store_test.go
Обычный файл
14
server/channels/store/sqlstore/audit_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestAuditStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestAuditStore)
|
||||
}
|
||||
221
server/channels/store/sqlstore/bot_store.go
Обычный файл
221
server/channels/store/sqlstore/bot_store.go
Обычный файл
@@ -0,0 +1,221 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/einterfaces"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
// bot is a subset of the model.Bot type, omitting the model.User fields.
|
||||
type bot struct {
|
||||
UserId string `json:"user_id"`
|
||||
Description string `json:"description"`
|
||||
OwnerId string `json:"owner_id"`
|
||||
LastIconUpdate int64 `json:"last_icon_update"`
|
||||
CreateAt int64 `json:"create_at"`
|
||||
UpdateAt int64 `json:"update_at"`
|
||||
DeleteAt int64 `json:"delete_at"`
|
||||
}
|
||||
|
||||
func botFromModel(b *model.Bot) *bot {
|
||||
return &bot{
|
||||
UserId: b.UserId,
|
||||
Description: b.Description,
|
||||
OwnerId: b.OwnerId,
|
||||
LastIconUpdate: b.LastIconUpdate,
|
||||
CreateAt: b.CreateAt,
|
||||
UpdateAt: b.UpdateAt,
|
||||
DeleteAt: b.DeleteAt,
|
||||
}
|
||||
}
|
||||
|
||||
// SqlBotStore is a store for managing bots in the database.
|
||||
// Bots are otherwise normal users with extra metadata record in the Bots table. The primary key
|
||||
// for a bot matches the primary key value for corresponding User record.
|
||||
type SqlBotStore struct {
|
||||
*SqlStore
|
||||
metrics einterfaces.MetricsInterface
|
||||
}
|
||||
|
||||
// newSqlBotStore creates an instance of SqlBotStore, registering the table schema in question.
|
||||
func newSqlBotStore(sqlStore *SqlStore, metrics einterfaces.MetricsInterface) store.BotStore {
|
||||
return &SqlBotStore{
|
||||
SqlStore: sqlStore,
|
||||
metrics: metrics,
|
||||
}
|
||||
}
|
||||
|
||||
// Get fetches the given bot in the database.
|
||||
func (us SqlBotStore) Get(botUserId string, includeDeleted bool) (*model.Bot, error) {
|
||||
var excludeDeletedSql = "AND b.DeleteAt = 0"
|
||||
if includeDeleted {
|
||||
excludeDeletedSql = ""
|
||||
}
|
||||
|
||||
query := `
|
||||
SELECT
|
||||
b.UserId,
|
||||
u.Username,
|
||||
u.FirstName AS DisplayName,
|
||||
b.Description,
|
||||
b.OwnerId,
|
||||
COALESCE(b.LastIconUpdate, 0) AS LastIconUpdate,
|
||||
b.CreateAt,
|
||||
b.UpdateAt,
|
||||
b.DeleteAt
|
||||
FROM
|
||||
Bots b
|
||||
JOIN
|
||||
Users u ON (u.Id = b.UserId)
|
||||
WHERE
|
||||
b.UserId = ?
|
||||
` + excludeDeletedSql + `
|
||||
`
|
||||
|
||||
var bot model.Bot
|
||||
if err := us.GetReplicaX().Get(&bot, query, botUserId); err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Bot", botUserId)
|
||||
} else if err != nil {
|
||||
return nil, errors.Wrapf(err, "selectone: user_id=%s", botUserId)
|
||||
}
|
||||
|
||||
return &bot, nil
|
||||
}
|
||||
|
||||
// GetAll fetches from all bots in the database.
|
||||
func (us SqlBotStore) GetAll(options *model.BotGetOptions) ([]*model.Bot, error) {
|
||||
var conditions []string
|
||||
var conditionsSql string
|
||||
var additionalJoin string
|
||||
var args []any
|
||||
|
||||
if !options.IncludeDeleted {
|
||||
conditions = append(conditions, "b.DeleteAt = 0")
|
||||
}
|
||||
if options.OwnerId != "" {
|
||||
conditions = append(conditions, "b.OwnerId = ?")
|
||||
args = append(args, options.OwnerId)
|
||||
}
|
||||
if options.OnlyOrphaned {
|
||||
additionalJoin = "JOIN Users o ON (o.Id = b.OwnerId)"
|
||||
conditions = append(conditions, "o.DeleteAt != 0")
|
||||
}
|
||||
|
||||
if len(conditions) > 0 {
|
||||
conditionsSql = "WHERE " + strings.Join(conditions, " AND ")
|
||||
}
|
||||
|
||||
sql := `
|
||||
SELECT
|
||||
b.UserId,
|
||||
u.Username,
|
||||
u.FirstName AS DisplayName,
|
||||
b.Description,
|
||||
b.OwnerId,
|
||||
COALESCE(b.LastIconUpdate, 0) AS LastIconUpdate,
|
||||
b.CreateAt,
|
||||
b.UpdateAt,
|
||||
b.DeleteAt
|
||||
FROM
|
||||
Bots b
|
||||
JOIN
|
||||
Users u ON (u.Id = b.UserId)
|
||||
` + additionalJoin + `
|
||||
` + conditionsSql + `
|
||||
ORDER BY
|
||||
b.CreateAt ASC,
|
||||
u.Username ASC
|
||||
LIMIT
|
||||
?
|
||||
OFFSET
|
||||
?
|
||||
`
|
||||
// append limit, offset
|
||||
args = append(args, options.PerPage, options.Page*options.PerPage)
|
||||
|
||||
bots := []*model.Bot{}
|
||||
if err := us.GetReplicaX().Select(&bots, sql, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "error selecting all bots")
|
||||
}
|
||||
|
||||
return bots, nil
|
||||
}
|
||||
|
||||
// Save persists a new bot to the database.
|
||||
// It assumes the corresponding user was saved via the user store.
|
||||
func (us SqlBotStore) Save(bot *model.Bot) (*model.Bot, error) {
|
||||
bot = bot.Clone()
|
||||
bot.PreSave()
|
||||
|
||||
if err := bot.IsValid(); err != nil { // TODO: change to return error in v6.
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := us.GetMasterX().NamedExec(`INSERT INTO Bots
|
||||
(UserId, Description, OwnerId, LastIconUpdate, CreateAt, UpdateAt, DeleteAt)
|
||||
VALUES
|
||||
(:UserId, :Description, :OwnerId, :LastIconUpdate, :CreateAt, :UpdateAt, :DeleteAt)`, botFromModel(bot)); err != nil {
|
||||
return nil, errors.Wrapf(err, "insert: user_id=%s", bot.UserId)
|
||||
}
|
||||
|
||||
return bot, nil
|
||||
}
|
||||
|
||||
// Update persists an updated bot to the database.
|
||||
// It assumes the corresponding user was updated via the user store.
|
||||
func (us SqlBotStore) Update(bot *model.Bot) (*model.Bot, error) {
|
||||
bot = bot.Clone()
|
||||
|
||||
bot.PreUpdate()
|
||||
if err := bot.IsValid(); err != nil { // TODO: needs to return error in v6
|
||||
return nil, err
|
||||
}
|
||||
|
||||
oldBot, err := us.Get(bot.UserId, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
oldBot.Description = bot.Description
|
||||
oldBot.OwnerId = bot.OwnerId
|
||||
oldBot.LastIconUpdate = bot.LastIconUpdate
|
||||
oldBot.UpdateAt = bot.UpdateAt
|
||||
oldBot.DeleteAt = bot.DeleteAt
|
||||
bot = oldBot
|
||||
|
||||
res, err := us.GetMasterX().NamedExec(`UPDATE Bots
|
||||
SET Description=:Description, OwnerId=:OwnerId, LastIconUpdate=:LastIconUpdate,
|
||||
UpdateAt=:UpdateAt, DeleteAt=:DeleteAt
|
||||
WHERE UserId=:UserId`, botFromModel(bot))
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "update: user_id=%s", bot.UserId)
|
||||
}
|
||||
count, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "error while getting rows_affected")
|
||||
}
|
||||
if count > 1 {
|
||||
return nil, fmt.Errorf("unexpected count while updating bot: count=%d, userId=%s", count, bot.UserId)
|
||||
}
|
||||
|
||||
return bot, nil
|
||||
}
|
||||
|
||||
// PermanentDelete removes the bot from the database altogether.
|
||||
// If the corresponding user is to be deleted, it must be done via the user store.
|
||||
func (us SqlBotStore) PermanentDelete(botUserId string) error {
|
||||
query := "DELETE FROM Bots WHERE UserId = ?"
|
||||
if _, err := us.GetMasterX().Exec(query, botUserId); err != nil {
|
||||
return store.NewErrInvalidInput("Bot", "UserId", botUserId).Wrap(err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/bot_store_test.go
Обычный файл
14
server/channels/store/sqlstore/bot_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestBotStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestBotStore)
|
||||
}
|
||||
269
server/channels/store/sqlstore/channel_member_history_store.go
Обычный файл
269
server/channels/store/sqlstore/channel_member_history_store.go
Обычный файл
@@ -0,0 +1,269 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
)
|
||||
|
||||
type SqlChannelMemberHistoryStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlChannelMemberHistoryStore(sqlStore *SqlStore) store.ChannelMemberHistoryStore {
|
||||
return &SqlChannelMemberHistoryStore{
|
||||
SqlStore: sqlStore,
|
||||
}
|
||||
}
|
||||
|
||||
func (s SqlChannelMemberHistoryStore) LogJoinEvent(userId string, channelId string, joinTime int64) error {
|
||||
channelMemberHistory := &model.ChannelMemberHistory{
|
||||
UserId: userId,
|
||||
ChannelId: channelId,
|
||||
JoinTime: joinTime,
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(`INSERT INTO ChannelMemberHistory
|
||||
(UserId, ChannelId, JoinTime)
|
||||
VALUES
|
||||
(:UserId, :ChannelId, :JoinTime)`, channelMemberHistory); err != nil {
|
||||
return errors.Wrapf(err, "LogJoinEvent userId=%s channelId=%s joinTime=%d", userId, channelId, joinTime)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlChannelMemberHistoryStore) LogLeaveEvent(userId string, channelId string, leaveTime int64) error {
|
||||
query, params, err := s.getQueryBuilder().
|
||||
Update("ChannelMemberHistory").
|
||||
Set("LeaveTime", leaveTime).
|
||||
Where(sq.And{
|
||||
sq.Eq{"UserId": userId},
|
||||
sq.Eq{"ChannelId": channelId},
|
||||
sq.Eq{"LeaveTime": nil},
|
||||
}).ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "channel_member_history_to_sql")
|
||||
}
|
||||
sqlResult, err := s.GetMasterX().Exec(query, params...)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "LogLeaveEvent userId=%s channelId=%s leaveTime=%d", userId, channelId, leaveTime)
|
||||
}
|
||||
|
||||
if rows, err := sqlResult.RowsAffected(); err == nil && rows != 1 {
|
||||
// there was no join event to update - this is best effort, so no need to raise an error
|
||||
mlog.Warn("Channel join event for user and channel not found", mlog.String("user", userId), mlog.String("channel", channelId))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlChannelMemberHistoryStore) GetUsersInChannelDuring(startTime int64, endTime int64, channelId string) ([]*model.ChannelMemberHistoryResult, error) {
|
||||
useChannelMemberHistory, err := s.hasDataAtOrBefore(startTime)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "hasDataAtOrBefore startTime=%d endTime=%d channelId=%s", startTime, endTime, channelId)
|
||||
}
|
||||
|
||||
if useChannelMemberHistory {
|
||||
// the export period starts after the ChannelMemberHistory table was first introduced, so we can use the
|
||||
// data from it for our export
|
||||
channelMemberHistories, err2 := s.getFromChannelMemberHistoryTable(startTime, endTime, channelId)
|
||||
if err2 != nil {
|
||||
return nil, errors.Wrapf(err2, "getFromChannelMemberHistoryTable startTime=%d endTime=%d channelId=%s", startTime, endTime, channelId)
|
||||
}
|
||||
return channelMemberHistories, nil
|
||||
}
|
||||
// the export period starts before the ChannelMemberHistory table was introduced, so we need to fake the
|
||||
// data by assuming that anybody who has ever joined the channel in question was present during the export period.
|
||||
// this may not always be true, but it's better than saying that somebody wasn't there when they were
|
||||
channelMemberHistories, err := s.getFromChannelMembersTable(startTime, endTime, channelId)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "getFromChannelMembersTable startTime=%d endTime=%d channelId=%s", startTime, endTime, channelId)
|
||||
}
|
||||
return channelMemberHistories, nil
|
||||
}
|
||||
|
||||
func (s SqlChannelMemberHistoryStore) hasDataAtOrBefore(time int64) (bool, error) {
|
||||
type NullableCountResult struct {
|
||||
Min sql.NullInt64
|
||||
}
|
||||
query, _, err := s.getQueryBuilder().Select("MIN(JoinTime) as Min").From("ChannelMemberHistory").ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "channel_member_history_to_sql")
|
||||
}
|
||||
var result NullableCountResult
|
||||
if err := s.GetReplicaX().Get(&result, query); err != nil {
|
||||
return false, err
|
||||
} else if result.Min.Valid {
|
||||
return result.Min.Int64 <= time, nil
|
||||
} else {
|
||||
// if the result was null, there are no rows in the table, so there is no data from before
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (s SqlChannelMemberHistoryStore) getFromChannelMemberHistoryTable(startTime int64, endTime int64, channelId string) ([]*model.ChannelMemberHistoryResult, error) {
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select(`cmh.*, u.Email AS "Email", u.Username, Bots.UserId IS NOT NULL AS IsBot, u.DeleteAt AS UserDeleteAt`).
|
||||
From("ChannelMemberHistory cmh").
|
||||
Join("Users u ON cmh.UserId = u.Id").
|
||||
LeftJoin("Bots ON Bots.UserId = u.Id").
|
||||
Where(sq.And{
|
||||
sq.Eq{"cmh.ChannelId": channelId},
|
||||
sq.LtOrEq{"cmh.JoinTime": endTime},
|
||||
sq.Or{
|
||||
sq.Eq{"cmh.LeaveTime": nil},
|
||||
sq.GtOrEq{"cmh.LeaveTime": startTime},
|
||||
},
|
||||
}).
|
||||
OrderBy("cmh.JoinTime ASC").ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "channel_member_history_to_sql")
|
||||
}
|
||||
histories := []*model.ChannelMemberHistoryResult{}
|
||||
if err := s.GetReplicaX().Select(&histories, query, args...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return histories, nil
|
||||
}
|
||||
|
||||
func (s SqlChannelMemberHistoryStore) getFromChannelMembersTable(startTime int64, endTime int64, channelId string) ([]*model.ChannelMemberHistoryResult, error) {
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select(`ch.ChannelId, ch.UserId, u.Email AS "Email", u.Username, Bots.UserId IS NOT NULL AS IsBot, u.DeleteAt AS UserDeleteAt`).
|
||||
Distinct().
|
||||
From("ChannelMembers ch").
|
||||
Join("Users u ON ch.UserId = u.id").
|
||||
LeftJoin("Bots ON Bots.UserId = u.id").
|
||||
Where(sq.Eq{"ch.ChannelId": channelId}).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "channel_member_history_to_sql")
|
||||
}
|
||||
|
||||
histories := []*model.ChannelMemberHistoryResult{}
|
||||
if err := s.GetReplicaX().Select(&histories, query, args...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// we have to fill in the join/leave times, because that data doesn't exist in the channel members table
|
||||
for _, channelMemberHistory := range histories {
|
||||
channelMemberHistory.JoinTime = startTime
|
||||
channelMemberHistory.LeaveTime = model.NewInt64(endTime)
|
||||
}
|
||||
return histories, nil
|
||||
}
|
||||
|
||||
// PermanentDeleteBatchForRetentionPolicies deletes a batch of records which are affected by
|
||||
// the global or a granular retention policy.
|
||||
// See `genericPermanentDeleteBatchForRetentionPolicies` for details.
|
||||
func (s SqlChannelMemberHistoryStore) PermanentDeleteBatchForRetentionPolicies(now, globalPolicyEndTime, limit int64, cursor model.RetentionPolicyCursor) (int64, model.RetentionPolicyCursor, error) {
|
||||
builder := s.getQueryBuilder().
|
||||
Select("ChannelMemberHistory.ChannelId, ChannelMemberHistory.UserId, ChannelMemberHistory.JoinTime").
|
||||
From("ChannelMemberHistory")
|
||||
return genericPermanentDeleteBatchForRetentionPolicies(RetentionPolicyBatchDeletionInfo{
|
||||
BaseBuilder: builder,
|
||||
Table: "ChannelMemberHistory",
|
||||
TimeColumn: "LeaveTime",
|
||||
PrimaryKeys: []string{"ChannelId", "UserId", "JoinTime"},
|
||||
ChannelIDTable: "ChannelMemberHistory",
|
||||
NowMillis: now,
|
||||
GlobalPolicyEndTime: globalPolicyEndTime,
|
||||
Limit: limit,
|
||||
}, s.SqlStore, cursor)
|
||||
}
|
||||
|
||||
// DeleteOrphanedRows removes entries from ChannelMemberHistory when a corresponding channel no longer exists.
|
||||
func (s SqlChannelMemberHistoryStore) DeleteOrphanedRows(limit int) (deleted int64, err error) {
|
||||
// We need the extra level of nesting to deal with MySQL's locking
|
||||
const query = `
|
||||
DELETE FROM ChannelMemberHistory WHERE (ChannelId, UserId, JoinTime) IN (
|
||||
SELECT * FROM (
|
||||
SELECT ChannelId, UserId, JoinTime FROM ChannelMemberHistory
|
||||
LEFT JOIN Channels ON ChannelMemberHistory.ChannelId = Channels.Id
|
||||
WHERE Channels.Id IS NULL
|
||||
LIMIT ?
|
||||
) AS A
|
||||
)`
|
||||
result, err := s.GetMasterX().Exec(query, limit)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
return result.RowsAffected()
|
||||
}
|
||||
|
||||
func (s SqlChannelMemberHistoryStore) PermanentDeleteBatch(endTime int64, limit int64) (int64, error) {
|
||||
var (
|
||||
query string
|
||||
args []any
|
||||
err error
|
||||
)
|
||||
|
||||
if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
var innerSelect string
|
||||
innerSelect, args, err = s.getQueryBuilder().
|
||||
Select("ctid").
|
||||
From("ChannelMemberHistory").
|
||||
Where(sq.And{
|
||||
sq.NotEq{"LeaveTime": nil},
|
||||
sq.LtOrEq{"LeaveTime": endTime},
|
||||
}).Limit(uint64(limit)).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "channel_member_history_to_sql")
|
||||
}
|
||||
query, _, err = s.getQueryBuilder().
|
||||
Delete("ChannelMemberHistory").
|
||||
Where(fmt.Sprintf(
|
||||
"ctid IN (%s)", innerSelect,
|
||||
)).ToSql()
|
||||
} else {
|
||||
query, args, err = s.getQueryBuilder().
|
||||
Delete("ChannelMemberHistory").
|
||||
Where(sq.And{
|
||||
sq.NotEq{"LeaveTime": nil},
|
||||
sq.LtOrEq{"LeaveTime": endTime},
|
||||
}).
|
||||
Limit(uint64(limit)).ToSql()
|
||||
}
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "channel_member_history_to_sql")
|
||||
}
|
||||
sqlResult, err := s.GetMasterX().Exec(query, args...)
|
||||
if err != nil {
|
||||
return 0, errors.Wrapf(err, "PermanentDeleteBatch endTime=%d limit=%d", endTime, limit)
|
||||
}
|
||||
|
||||
rowsAffected, err := sqlResult.RowsAffected()
|
||||
if err != nil {
|
||||
return 0, errors.Wrapf(err, "PermanentDeleteBatch endTime=%d limit=%d", endTime, limit)
|
||||
}
|
||||
return rowsAffected, nil
|
||||
}
|
||||
|
||||
// GetChannelsLeftSince returns list of channels that the user has left after a given time,
|
||||
// but has not rejoined again.
|
||||
func (s SqlChannelMemberHistoryStore) GetChannelsLeftSince(userID string, since int64) ([]string, error) {
|
||||
query, params, err := s.getQueryBuilder().
|
||||
Select("ChannelId").
|
||||
From("ChannelMemberHistory").
|
||||
GroupBy("ChannelId").
|
||||
Where(sq.Eq{"UserId": userID}).
|
||||
Having("MAX(LeaveTime) > MAX(JoinTime) AND MAX(LeaveTime) IS NOT NULL AND MAX(LeaveTime) >= ?", since).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "channel_member_history_to_sql")
|
||||
}
|
||||
channelIds := []string{}
|
||||
err = s.GetReplicaX().Select(&channelIds, query, params...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "GetChannelsLeftSince userId=%s since=%d", userID, since)
|
||||
}
|
||||
|
||||
return channelIds, nil
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestChannelMemberHistoryStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestChannelMemberHistoryStore)
|
||||
}
|
||||
4714
server/channels/store/sqlstore/channel_store.go
Обычный файл
4714
server/channels/store/sqlstore/channel_store.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
1127
server/channels/store/sqlstore/channel_store_categories.go
Обычный файл
1127
server/channels/store/sqlstore/channel_store_categories.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
14
server/channels/store/sqlstore/channel_store_categories_test.go
Обычный файл
14
server/channels/store/sqlstore/channel_store_categories_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestChannelStoreCategories(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestChannelStoreCategories)
|
||||
}
|
||||
1397
server/channels/store/sqlstore/channel_store_test.go
Обычный файл
1397
server/channels/store/sqlstore/channel_store_test.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
139
server/channels/store/sqlstore/cluster_discovery_store.go
Обычный файл
139
server/channels/store/sqlstore/cluster_discovery_store.go
Обычный файл
@@ -0,0 +1,139 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type sqlClusterDiscoveryStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlClusterDiscoveryStore(sqlStore *SqlStore) store.ClusterDiscoveryStore {
|
||||
return &sqlClusterDiscoveryStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s sqlClusterDiscoveryStore) Save(ClusterDiscovery *model.ClusterDiscovery) error {
|
||||
ClusterDiscovery.PreSave()
|
||||
if err := ClusterDiscovery.IsValid(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(`
|
||||
INSERT INTO
|
||||
ClusterDiscovery
|
||||
(Id, Type, ClusterName, Hostname, GossipPort, Port, CreateAt, LastPingAt)
|
||||
VALUES
|
||||
(:Id, :Type, :ClusterName, :Hostname, :GossipPort, :Port, :CreateAt, :LastPingAt)
|
||||
`, ClusterDiscovery); err != nil {
|
||||
return errors.Wrap(err, "failed to save ClusterDiscovery")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s sqlClusterDiscoveryStore) Delete(ClusterDiscovery *model.ClusterDiscovery) (bool, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Delete("ClusterDiscovery").
|
||||
Where(sq.Eq{"Type": ClusterDiscovery.Type}).
|
||||
Where(sq.Eq{"ClusterName": ClusterDiscovery.ClusterName}).
|
||||
Where(sq.Eq{"Hostname": ClusterDiscovery.Hostname})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "cluster_discovery_tosql")
|
||||
}
|
||||
|
||||
res, err := s.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to delete ClusterDiscovery")
|
||||
}
|
||||
|
||||
count, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to count rows affected")
|
||||
}
|
||||
|
||||
return count != 0, nil
|
||||
}
|
||||
|
||||
func (s sqlClusterDiscoveryStore) Exists(ClusterDiscovery *model.ClusterDiscovery) (bool, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("COUNT(*)").
|
||||
From("ClusterDiscovery").
|
||||
Where(sq.Eq{"Type": ClusterDiscovery.Type}).
|
||||
Where(sq.Eq{"ClusterName": ClusterDiscovery.ClusterName}).
|
||||
Where(sq.Eq{"Hostname": ClusterDiscovery.Hostname})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "cluster_discovery_tosql")
|
||||
}
|
||||
|
||||
var count int
|
||||
if err := s.GetMasterX().Get(&count, queryString, args...); err != nil {
|
||||
return false, errors.Wrap(err, "failed to count ClusterDiscovery")
|
||||
}
|
||||
|
||||
return count != 0, nil
|
||||
}
|
||||
|
||||
func (s sqlClusterDiscoveryStore) GetAll(ClusterDiscoveryType, clusterName string) ([]*model.ClusterDiscovery, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("ClusterDiscovery").
|
||||
Where(sq.Eq{"Type": ClusterDiscoveryType}).
|
||||
Where(sq.Eq{"ClusterName": clusterName}).
|
||||
Where(sq.Gt{"LastPingAt": model.GetMillis() - model.CDSOfflineAfterMillis})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "cluster_discovery_tosql")
|
||||
}
|
||||
|
||||
list := []*model.ClusterDiscovery{}
|
||||
if err := s.GetMasterX().Select(&list, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find ClusterDiscovery")
|
||||
}
|
||||
return list, nil
|
||||
}
|
||||
|
||||
func (s sqlClusterDiscoveryStore) SetLastPingAt(ClusterDiscovery *model.ClusterDiscovery) error {
|
||||
query := s.getQueryBuilder().
|
||||
Update("ClusterDiscovery").
|
||||
Set("LastPingAt", model.GetMillis()).
|
||||
Where(sq.Eq{"Type": ClusterDiscovery.Type}).
|
||||
Where(sq.Eq{"ClusterName": ClusterDiscovery.ClusterName}).
|
||||
Where(sq.Eq{"Hostname": ClusterDiscovery.Hostname})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "cluster_discovery_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to update ClusterDiscovery")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s sqlClusterDiscoveryStore) Cleanup() error {
|
||||
query := s.getQueryBuilder().
|
||||
Delete("ClusterDiscovery").
|
||||
Where(sq.Lt{"LastPingAt": model.GetMillis() - model.CDSOfflineAfterMillis})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "cluster_discovery_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to delete ClusterDiscoveries")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/cluster_discovery_store_test.go
Обычный файл
14
server/channels/store/sqlstore/cluster_discovery_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestClusterDiscoveryStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestClusterDiscoveryStore)
|
||||
}
|
||||
230
server/channels/store/sqlstore/command_store.go
Обычный файл
230
server/channels/store/sqlstore/command_store.go
Обычный файл
@@ -0,0 +1,230 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlCommandStore struct {
|
||||
*SqlStore
|
||||
|
||||
commandsQuery sq.SelectBuilder
|
||||
}
|
||||
|
||||
func newSqlCommandStore(sqlStore *SqlStore) store.CommandStore {
|
||||
s := &SqlCommandStore{SqlStore: sqlStore}
|
||||
|
||||
s.commandsQuery = s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Commands")
|
||||
return s
|
||||
}
|
||||
|
||||
func (s SqlCommandStore) Save(command *model.Command) (*model.Command, error) {
|
||||
if command.Id != "" {
|
||||
return nil, store.NewErrInvalidInput("Command", "CommandId", command.Id)
|
||||
}
|
||||
|
||||
command.PreSave()
|
||||
if err := command.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Trigger is a keyword
|
||||
trigger := s.toReserveCase("trigger")
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(`INSERT INTO Commands (Id, Token, CreateAt,
|
||||
UpdateAt, DeleteAt, CreatorId, TeamId, `+trigger+`, Method, Username,
|
||||
IconURL, AutoComplete, AutoCompleteDesc, AutoCompleteHint, DisplayName, Description,
|
||||
URL, PluginId)
|
||||
VALUES (:Id, :Token, :CreateAt, :UpdateAt, :DeleteAt, :CreatorId, :TeamId, :Trigger, :Method,
|
||||
:Username, :IconURL, :AutoComplete, :AutoCompleteDesc, :AutoCompleteHint, :DisplayName,
|
||||
:Description, :URL, :PluginId)`, command); err != nil {
|
||||
return nil, errors.Wrapf(err, "insert: command_id=%s", command.Id)
|
||||
}
|
||||
|
||||
return command, nil
|
||||
}
|
||||
|
||||
func (s SqlCommandStore) Get(id string) (*model.Command, error) {
|
||||
var command model.Command
|
||||
|
||||
query, args, err := s.commandsQuery.
|
||||
Where(sq.Eq{"Id": id, "DeleteAt": 0}).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "commands_tosql")
|
||||
}
|
||||
if err = s.GetReplicaX().Get(&command, query, args...); err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Command", id)
|
||||
} else if err != nil {
|
||||
return nil, errors.Wrapf(err, "selectone: command_id=%s", id)
|
||||
}
|
||||
|
||||
return &command, nil
|
||||
}
|
||||
|
||||
func (s SqlCommandStore) GetByTeam(teamId string) ([]*model.Command, error) {
|
||||
commands := []*model.Command{}
|
||||
|
||||
sql, args, err := s.commandsQuery.
|
||||
Where(sq.Eq{"TeamId": teamId, "DeleteAt": 0}).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "commands_tosql")
|
||||
}
|
||||
if err := s.GetReplicaX().Select(&commands, sql, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "select: team_id=%s", teamId)
|
||||
}
|
||||
|
||||
return commands, nil
|
||||
}
|
||||
|
||||
func (s SqlCommandStore) GetByTrigger(teamId string, trigger string) (*model.Command, error) {
|
||||
var command model.Command
|
||||
var triggerStr string
|
||||
if s.DriverName() == "mysql" {
|
||||
triggerStr = "`Trigger`"
|
||||
} else {
|
||||
triggerStr = "\"trigger\""
|
||||
}
|
||||
|
||||
query, args, err := s.commandsQuery.
|
||||
Where(sq.Eq{"TeamId": teamId, "DeleteAt": 0, triggerStr: trigger}).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "commands_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Get(&command, query, args...); err == sql.ErrNoRows {
|
||||
errorId := "teamId=" + teamId + ", trigger=" + trigger
|
||||
return nil, store.NewErrNotFound("Command", errorId)
|
||||
} else if err != nil {
|
||||
return nil, errors.Wrapf(err, "selectone: team_id=%s, trigger=%s", teamId, trigger)
|
||||
}
|
||||
|
||||
return &command, nil
|
||||
}
|
||||
|
||||
func (s SqlCommandStore) Delete(commandId string, time int64) error {
|
||||
sql, args, err := s.getQueryBuilder().
|
||||
Update("Commands").
|
||||
SetMap(sq.Eq{"DeleteAt": time, "UpdateAt": time}).
|
||||
Where(sq.Eq{"Id": commandId}).ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "commands_tosql")
|
||||
}
|
||||
|
||||
_, err = s.GetMasterX().Exec(sql, args...)
|
||||
if err != nil {
|
||||
errors.Wrapf(err, "delete: command_id=%s", commandId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlCommandStore) PermanentDeleteByTeam(teamId string) error {
|
||||
sql, args, err := s.getQueryBuilder().
|
||||
Delete("Commands").
|
||||
Where(sq.Eq{"TeamId": teamId}).ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "commands_tosql")
|
||||
}
|
||||
_, err = s.GetMasterX().Exec(sql, args...)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "delete: team_id=%s", teamId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlCommandStore) PermanentDeleteByUser(userId string) error {
|
||||
sql, args, err := s.getQueryBuilder().
|
||||
Delete("Commands").
|
||||
Where(sq.Eq{"CreatorId": userId}).ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "commands_tosql")
|
||||
}
|
||||
_, err = s.GetMasterX().Exec(sql, args...)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "delete: user_id=%s", userId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlCommandStore) Update(cmd *model.Command) (*model.Command, error) {
|
||||
cmd.UpdateAt = model.GetMillis()
|
||||
|
||||
if err := cmd.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Update("Commands").
|
||||
Set("Token", cmd.Token).
|
||||
Set("CreateAt", cmd.CreateAt).
|
||||
Set("UpdateAt", cmd.UpdateAt).
|
||||
Set("CreatorId", cmd.CreatorId).
|
||||
Set("TeamId", cmd.TeamId).
|
||||
Set("Method", cmd.Method).
|
||||
Set("Username", cmd.Username).
|
||||
Set("IconURL", cmd.IconURL).
|
||||
Set("AutoComplete", cmd.AutoComplete).
|
||||
Set("AutoCompleteDesc", cmd.AutoCompleteDesc).
|
||||
Set("AutoCompleteHint", cmd.AutoCompleteHint).
|
||||
Set("DisplayName", cmd.DisplayName).
|
||||
Set("Description", cmd.Description).
|
||||
Set("URL", cmd.URL).
|
||||
Set("PluginId", cmd.PluginId).
|
||||
Where(sq.Eq{"Id": cmd.Id})
|
||||
|
||||
// Trigger is a keyword
|
||||
query = query.Set(s.toReserveCase("trigger"), cmd.Trigger)
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "commands_tosql")
|
||||
}
|
||||
|
||||
res, err := s.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to update commands")
|
||||
}
|
||||
count, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "error while getting rows_affected")
|
||||
}
|
||||
if count > 1 {
|
||||
return nil, fmt.Errorf("unexpected count while updating commands: count=%d, Id=%s", count, cmd.Id)
|
||||
}
|
||||
return cmd, nil
|
||||
}
|
||||
|
||||
func (s SqlCommandStore) AnalyticsCommandCount(teamId string) (int64, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("COUNT(*)").
|
||||
From("Commands").
|
||||
Where(sq.Eq{"DeleteAt": 0})
|
||||
|
||||
if teamId != "" {
|
||||
query = query.Where(sq.Eq{"TeamId": teamId})
|
||||
}
|
||||
|
||||
sql, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrapf(err, "commands_tosql")
|
||||
}
|
||||
|
||||
var c int64
|
||||
err = s.GetReplicaX().Get(&c, sql, args...)
|
||||
if err != nil {
|
||||
return 0, errors.Wrapf(err, "unable to count the commands: team_id=%s", teamId)
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/command_store_test.go
Обычный файл
14
server/channels/store/sqlstore/command_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestCommandStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestCommandStore)
|
||||
}
|
||||
109
server/channels/store/sqlstore/command_webhook_store.go
Обычный файл
109
server/channels/store/sqlstore/command_webhook_store.go
Обычный файл
@@ -0,0 +1,109 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
)
|
||||
|
||||
type SqlCommandWebhookStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlCommandWebhookStore(sqlStore *SqlStore) store.CommandWebhookStore {
|
||||
return &SqlCommandWebhookStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlCommandWebhookStore) Save(webhook *model.CommandWebhook) (*model.CommandWebhook, error) {
|
||||
if webhook.Id != "" {
|
||||
return nil, store.NewErrInvalidInput("CommandWebhook", "id", webhook.Id)
|
||||
}
|
||||
|
||||
webhook.PreSave()
|
||||
if err := webhook.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(`INSERT INTO CommandWebhooks
|
||||
(Id,CreateAt,CommandId,UserId,ChannelId,RootId,UseCount)
|
||||
Values
|
||||
(:Id, :CreateAt, :CommandId, :UserId, :ChannelId, :RootId, :UseCount)`, webhook); err != nil {
|
||||
return nil, errors.Wrapf(err, "save: id=%s", webhook.Id)
|
||||
}
|
||||
|
||||
return webhook, nil
|
||||
}
|
||||
|
||||
func (s SqlCommandWebhookStore) Get(id string) (*model.CommandWebhook, error) {
|
||||
var webhook model.CommandWebhook
|
||||
|
||||
exptime := model.GetMillis() - model.CommandWebhookLifetime
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("CommandWebhooks").
|
||||
Where(sq.Eq{"Id": id}).
|
||||
Where(sq.Gt{"CreateAt": exptime})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "get_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Get(&webhook, queryString, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("CommandWebhook", id)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "get: id=%s", id)
|
||||
}
|
||||
|
||||
return &webhook, nil
|
||||
}
|
||||
|
||||
func (s SqlCommandWebhookStore) TryUse(id string, limit int) error {
|
||||
query := s.getQueryBuilder().
|
||||
Update("CommandWebhooks").
|
||||
Set("UseCount", sq.Expr("UseCount + 1")).
|
||||
Where(sq.Eq{"Id": id}).
|
||||
Where(sq.Lt{"UseCount": limit})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "tryuse_tosql")
|
||||
}
|
||||
|
||||
if sqlResult, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrapf(err, "tryuse: id=%s limit=%d", id, limit)
|
||||
} else if rows, err := sqlResult.RowsAffected(); rows == 0 {
|
||||
return store.NewErrInvalidInput("CommandWebhook", "id", id).Wrap(err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlCommandWebhookStore) Cleanup() {
|
||||
mlog.Debug("Cleaning up command webhook store.")
|
||||
exptime := model.GetMillis() - model.CommandWebhookLifetime
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Delete("CommandWebhooks").
|
||||
Where(sq.Lt{"CreateAt": exptime})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
mlog.Error("Failed to build query when trying to perform a cleanup in command webhook store.")
|
||||
return
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
mlog.Error("Unable to cleanup command webhook store.")
|
||||
}
|
||||
}
|
||||
14
server/channels/store/sqlstore/command_webhook_store_test.go
Обычный файл
14
server/channels/store/sqlstore/command_webhook_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestCommandWebhookStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestCommandWebhookStore)
|
||||
}
|
||||
329
server/channels/store/sqlstore/compliance_store.go
Обычный файл
329
server/channels/store/sqlstore/compliance_store.go
Обычный файл
@@ -0,0 +1,329 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlComplianceStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlComplianceStore(sqlStore *SqlStore) store.ComplianceStore {
|
||||
return &SqlComplianceStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlComplianceStore) Save(compliance *model.Compliance) (*model.Compliance, error) {
|
||||
compliance.PreSave()
|
||||
if err := compliance.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// DESC is a keyword
|
||||
desc := s.toReserveCase("desc")
|
||||
|
||||
query := `INSERT INTO Compliances (Id, CreateAt, UserId, Status, Count, ` + desc + `, Type, StartAt, EndAt, Keywords, Emails)
|
||||
VALUES
|
||||
(:Id, :CreateAt, :UserId, :Status, :Count, :Desc, :Type, :StartAt, :EndAt, :Keywords, :Emails)`
|
||||
if _, err := s.GetMasterX().NamedExec(query, compliance); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save Compliance")
|
||||
}
|
||||
return compliance, nil
|
||||
}
|
||||
|
||||
func (s SqlComplianceStore) Update(compliance *model.Compliance) (*model.Compliance, error) {
|
||||
if err := compliance.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Update("Compliances").
|
||||
Set("CreateAt", compliance.CreateAt).
|
||||
Set("UserId", compliance.UserId).
|
||||
Set("Status", compliance.Status).
|
||||
Set("Count", compliance.Count).
|
||||
Set("Type", compliance.Type).
|
||||
Set("StartAt", compliance.StartAt).
|
||||
Set("EndAt", compliance.EndAt).
|
||||
Set("Keywords", compliance.Keywords).
|
||||
Set("Emails", compliance.Emails).
|
||||
Where(sq.Eq{"Id": compliance.Id})
|
||||
|
||||
// DESC is a keyword
|
||||
query = query.Set(s.toReserveCase("desc"), compliance.Desc)
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "compliances_tosql")
|
||||
}
|
||||
|
||||
res, err := s.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to update Compliance")
|
||||
}
|
||||
count, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "error while getting rows_affected")
|
||||
}
|
||||
if count > 1 {
|
||||
return nil, fmt.Errorf("unexpected count while updating compliances: count=%d, Id=%s", count, compliance.Id)
|
||||
}
|
||||
return compliance, nil
|
||||
}
|
||||
|
||||
func (s SqlComplianceStore) GetAll(offset, limit int) (model.Compliances, error) {
|
||||
query := "SELECT * FROM Compliances ORDER BY CreateAt DESC LIMIT ? OFFSET ?"
|
||||
compliances := model.Compliances{}
|
||||
if err := s.GetReplicaX().Select(&compliances, query, limit, offset); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find all Compliances")
|
||||
}
|
||||
return compliances, nil
|
||||
}
|
||||
|
||||
func (s SqlComplianceStore) Get(id string) (*model.Compliance, error) {
|
||||
var compliance model.Compliance
|
||||
if err := s.GetReplicaX().Get(&compliance, `SELECT * FROM Compliances WHERE Id = ?`, id); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Compliances", id)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get Compliance with id=%s", id)
|
||||
}
|
||||
if compliance.Id == "" {
|
||||
return nil, store.NewErrNotFound("Compliance", id)
|
||||
}
|
||||
return &compliance, nil
|
||||
}
|
||||
|
||||
func (s SqlComplianceStore) ComplianceExport(job *model.Compliance, cursor model.ComplianceExportCursor, limit int) ([]*model.CompliancePost, model.ComplianceExportCursor, error) {
|
||||
keywordQuery := ""
|
||||
var argsKeywords []any
|
||||
keywords := strings.Fields(strings.TrimSpace(strings.ToLower(strings.Replace(job.Keywords, ",", " ", -1))))
|
||||
if len(keywords) > 0 {
|
||||
clauses := make([]string, len(keywords))
|
||||
|
||||
for i, keyword := range keywords {
|
||||
keyword = sanitizeSearchTerm(keyword, "\\")
|
||||
clauses[i] = "LOWER(Posts.Message) LIKE ?"
|
||||
argsKeywords = append(argsKeywords, "%"+keyword+"%")
|
||||
}
|
||||
|
||||
keywordQuery = "AND (" + strings.Join(clauses, " OR ") + ")"
|
||||
}
|
||||
|
||||
emailQuery := ""
|
||||
var argsEmails []any
|
||||
emails := strings.Fields(strings.TrimSpace(strings.ToLower(strings.Replace(job.Emails, ",", " ", -1))))
|
||||
if len(emails) > 0 {
|
||||
clauses := make([]string, len(emails))
|
||||
|
||||
for i, email := range emails {
|
||||
clauses[i] = "Users.Email = ?"
|
||||
argsEmails = append(argsEmails, email)
|
||||
}
|
||||
|
||||
emailQuery = "AND (" + strings.Join(clauses, " OR ") + ")"
|
||||
}
|
||||
|
||||
// The idea is to first iterate over the channel posts, and then when we run out of those,
|
||||
// start iterating over the direct message posts.
|
||||
|
||||
channelPosts := []*model.CompliancePost{}
|
||||
channelsQuery := ""
|
||||
var argsChannelsQuery []any
|
||||
if !cursor.ChannelsQueryCompleted {
|
||||
if cursor.LastChannelsQueryPostCreateAt == 0 {
|
||||
cursor.LastChannelsQueryPostCreateAt = job.StartAt
|
||||
}
|
||||
// append the named parameters of SQL query in the correct order to argsChannelsQuery
|
||||
argsChannelsQuery = append(argsChannelsQuery, cursor.LastChannelsQueryPostCreateAt, cursor.LastChannelsQueryPostCreateAt, cursor.LastChannelsQueryPostID, job.EndAt)
|
||||
argsChannelsQuery = append(argsChannelsQuery, argsEmails...)
|
||||
argsChannelsQuery = append(argsChannelsQuery, argsKeywords...)
|
||||
argsChannelsQuery = append(argsChannelsQuery, limit)
|
||||
channelsQuery = `
|
||||
SELECT
|
||||
Teams.Name AS TeamName,
|
||||
Teams.DisplayName AS TeamDisplayName,
|
||||
Channels.Name AS ChannelName,
|
||||
Channels.DisplayName AS ChannelDisplayName,
|
||||
Channels.Type AS ChannelType,
|
||||
Users.Username AS UserUsername,
|
||||
Users.Email AS UserEmail,
|
||||
Users.Nickname AS UserNickname,
|
||||
Posts.Id AS PostId,
|
||||
Posts.CreateAt AS PostCreateAt,
|
||||
Posts.UpdateAt AS PostUpdateAt,
|
||||
Posts.DeleteAt AS PostDeleteAt,
|
||||
Posts.RootId AS PostRootId,
|
||||
Posts.OriginalId AS PostOriginalId,
|
||||
Posts.Message AS PostMessage,
|
||||
Posts.Type AS PostType,
|
||||
Posts.Props AS PostProps,
|
||||
Posts.Hashtags AS PostHashtags,
|
||||
Posts.FileIds AS PostFileIds,
|
||||
Bots.UserId IS NOT NULL AS IsBot
|
||||
FROM
|
||||
Teams,
|
||||
Channels,
|
||||
Users,
|
||||
Posts
|
||||
LEFT JOIN
|
||||
Bots ON Bots.UserId = Posts.UserId
|
||||
WHERE
|
||||
Teams.Id = Channels.TeamId
|
||||
AND Posts.ChannelId = Channels.Id
|
||||
AND Posts.UserId = Users.Id
|
||||
AND (
|
||||
Posts.CreateAt > ?
|
||||
OR (Posts.CreateAt = ? AND Posts.Id > ?)
|
||||
)
|
||||
AND Posts.CreateAt < ?
|
||||
` + emailQuery + `
|
||||
` + keywordQuery + `
|
||||
ORDER BY Posts.CreateAt, Posts.Id
|
||||
LIMIT ?`
|
||||
if err := s.GetReplicaX().Select(&channelPosts, channelsQuery, argsChannelsQuery...); err != nil {
|
||||
return nil, cursor, errors.Wrap(err, "unable to export compliance")
|
||||
}
|
||||
if len(channelPosts) < limit {
|
||||
cursor.ChannelsQueryCompleted = true
|
||||
} else {
|
||||
cursor.LastChannelsQueryPostCreateAt = channelPosts[len(channelPosts)-1].PostCreateAt
|
||||
cursor.LastChannelsQueryPostID = channelPosts[len(channelPosts)-1].PostId
|
||||
}
|
||||
}
|
||||
|
||||
directMessagePosts := []*model.CompliancePost{}
|
||||
directMessagesQuery := ""
|
||||
var argsDirectMessagesQuery []any
|
||||
if !cursor.DirectMessagesQueryCompleted && len(channelPosts) < limit {
|
||||
if cursor.LastDirectMessagesQueryPostCreateAt == 0 {
|
||||
cursor.LastDirectMessagesQueryPostCreateAt = job.StartAt
|
||||
}
|
||||
// append the named parameters of SQL query in the correct order to argsDirectMessagesQuery
|
||||
argsDirectMessagesQuery = append(argsDirectMessagesQuery, cursor.LastDirectMessagesQueryPostCreateAt, cursor.LastDirectMessagesQueryPostCreateAt, cursor.LastDirectMessagesQueryPostID, job.EndAt)
|
||||
argsDirectMessagesQuery = append(argsDirectMessagesQuery, argsEmails...)
|
||||
argsDirectMessagesQuery = append(argsDirectMessagesQuery, argsKeywords...)
|
||||
argsDirectMessagesQuery = append(argsDirectMessagesQuery, limit-len(channelPosts))
|
||||
directMessagesQuery = `
|
||||
SELECT
|
||||
'direct-messages' AS TeamName,
|
||||
'Direct Messages' AS TeamDisplayName,
|
||||
Channels.Name AS ChannelName,
|
||||
Channels.DisplayName AS ChannelDisplayName,
|
||||
Channels.Type AS ChannelType,
|
||||
Users.Username AS UserUsername,
|
||||
Users.Email AS UserEmail,
|
||||
Users.Nickname AS UserNickname,
|
||||
Posts.Id AS PostId,
|
||||
Posts.CreateAt AS PostCreateAt,
|
||||
Posts.UpdateAt AS PostUpdateAt,
|
||||
Posts.DeleteAt AS PostDeleteAt,
|
||||
Posts.RootId AS PostRootId,
|
||||
Posts.OriginalId AS PostOriginalId,
|
||||
Posts.Message AS PostMessage,
|
||||
Posts.Type AS PostType,
|
||||
Posts.Props AS PostProps,
|
||||
Posts.Hashtags AS PostHashtags,
|
||||
Posts.FileIds AS PostFileIds,
|
||||
Bots.UserId IS NOT NULL AS IsBot
|
||||
FROM
|
||||
Channels,
|
||||
Users,
|
||||
Posts
|
||||
LEFT JOIN
|
||||
Bots ON Bots.UserId = Posts.UserId
|
||||
WHERE
|
||||
Channels.TeamId = ''
|
||||
AND Posts.ChannelId = Channels.Id
|
||||
AND Posts.UserId = Users.Id
|
||||
AND (
|
||||
Posts.CreateAt > ?
|
||||
OR (Posts.CreateAt = ? AND Posts.Id > ?)
|
||||
)
|
||||
AND Posts.CreateAt < ?
|
||||
` + emailQuery + `
|
||||
` + keywordQuery + `
|
||||
ORDER BY Posts.CreateAt, Posts.Id
|
||||
LIMIT ?`
|
||||
|
||||
if err := s.GetReplicaX().Select(&directMessagePosts, directMessagesQuery, argsDirectMessagesQuery...); err != nil {
|
||||
return nil, cursor, errors.Wrap(err, "unable to export compliance")
|
||||
}
|
||||
if len(directMessagePosts) < limit {
|
||||
cursor.DirectMessagesQueryCompleted = true
|
||||
} else {
|
||||
cursor.LastDirectMessagesQueryPostCreateAt = directMessagePosts[len(directMessagePosts)-1].PostCreateAt
|
||||
cursor.LastDirectMessagesQueryPostID = directMessagePosts[len(directMessagePosts)-1].PostId
|
||||
}
|
||||
}
|
||||
|
||||
return append(channelPosts, directMessagePosts...), cursor, nil
|
||||
}
|
||||
|
||||
func (s SqlComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
|
||||
var args []any
|
||||
args = append(args, model.ChannelTypeDirect, model.ChannelTypeGroup, cursor.LastPostUpdateAt, cursor.LastPostUpdateAt, cursor.LastPostId, limit)
|
||||
query :=
|
||||
`SELECT
|
||||
Posts.Id AS PostId,
|
||||
Posts.CreateAt AS PostCreateAt,
|
||||
Posts.UpdateAt AS PostUpdateAt,
|
||||
Posts.DeleteAt AS PostDeleteAt,
|
||||
Posts.Message AS PostMessage,
|
||||
Posts.Type AS PostType,
|
||||
Posts.Props AS PostProps,
|
||||
Posts.OriginalId AS PostOriginalId,
|
||||
Posts.RootId AS PostRootId,
|
||||
Posts.FileIds AS PostFileIds,
|
||||
Teams.Id AS TeamId,
|
||||
Teams.Name AS TeamName,
|
||||
Teams.DisplayName AS TeamDisplayName,
|
||||
Channels.Id AS ChannelId,
|
||||
CASE
|
||||
WHEN Channels.Type = ? THEN 'Direct Message'
|
||||
WHEN Channels.Type = ? THEN 'Group Message'
|
||||
ELSE Channels.DisplayName
|
||||
END AS ChannelDisplayName,
|
||||
Channels.Name AS ChannelName,
|
||||
Channels.Type AS ChannelType,
|
||||
Users.Id AS UserId,
|
||||
Users.Email AS UserEmail,
|
||||
Users.Username,
|
||||
Bots.UserId IS NOT NULL AS IsBot
|
||||
FROM
|
||||
Posts
|
||||
LEFT OUTER JOIN Channels ON Posts.ChannelId = Channels.Id
|
||||
LEFT OUTER JOIN Teams ON Channels.TeamId = Teams.Id
|
||||
LEFT OUTER JOIN Users ON Posts.UserId = Users.Id
|
||||
LEFT JOIN Bots ON Bots.UserId = Posts.UserId
|
||||
WHERE (
|
||||
Posts.UpdateAt > ?
|
||||
OR (
|
||||
Posts.UpdateAt = ?
|
||||
AND Posts.Id > ?
|
||||
)
|
||||
) AND Posts.Type NOT LIKE 'system_%'
|
||||
ORDER BY PostUpdateAt, PostId
|
||||
LIMIT ?`
|
||||
|
||||
cposts := []*model.MessageExport{}
|
||||
if err := s.GetReplicaX().SelectCtx(ctx, &cposts, query, args...); err != nil {
|
||||
return nil, cursor, errors.Wrap(err, "unable to export messages")
|
||||
}
|
||||
if len(cposts) > 0 {
|
||||
cursor.LastPostUpdateAt = *cposts[len(cposts)-1].PostUpdateAt
|
||||
cursor.LastPostId = *cposts[len(cposts)-1].PostId
|
||||
}
|
||||
return cposts, cursor, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/compliance_store_test.go
Обычный файл
14
server/channels/store/sqlstore/compliance_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestComplianceStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestComplianceStore)
|
||||
}
|
||||
42
server/channels/store/sqlstore/context.go
Обычный файл
42
server/channels/store/sqlstore/context.go
Обычный файл
@@ -0,0 +1,42 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
)
|
||||
|
||||
// storeContextKey is the base type for all context keys for the store.
|
||||
type storeContextKey string
|
||||
|
||||
// contextValue is a type to hold some pre-determined context values.
|
||||
type contextValue string
|
||||
|
||||
// Different possible values of contextValue.
|
||||
const (
|
||||
useMaster contextValue = "useMaster"
|
||||
)
|
||||
|
||||
// WithMaster adds the context value that master DB should be selected for this request.
|
||||
func WithMaster(ctx context.Context) context.Context {
|
||||
return context.WithValue(ctx, storeContextKey(useMaster), true)
|
||||
}
|
||||
|
||||
// hasMaster is a helper function to check whether master DB should be selected or not.
|
||||
func hasMaster(ctx context.Context) bool {
|
||||
if v := ctx.Value(storeContextKey(useMaster)); v != nil {
|
||||
if res, ok := v.(bool); ok && res {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// DBXFromContext is a helper utility that returns the sqlx DB handle from a given context.
|
||||
func (ss *SqlStore) DBXFromContext(ctx context.Context) *sqlxDBWrapper {
|
||||
if hasMaster(ctx) {
|
||||
return ss.GetMasterX()
|
||||
}
|
||||
return ss.GetReplicaX()
|
||||
}
|
||||
18
server/channels/store/sqlstore/context_test.go
Обычный файл
18
server/channels/store/sqlstore/context_test.go
Обычный файл
@@ -0,0 +1,18 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestContextMaster(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
m := WithMaster(ctx)
|
||||
assert.True(t, hasMaster(m))
|
||||
}
|
||||
260
server/channels/store/sqlstore/draft_store.go
Обычный файл
260
server/channels/store/sqlstore/draft_store.go
Обычный файл
@@ -0,0 +1,260 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"sync"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/einterfaces"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
)
|
||||
|
||||
type SqlDraftStore struct {
|
||||
*SqlStore
|
||||
metrics einterfaces.MetricsInterface
|
||||
maxDraftSizeOnce sync.Once
|
||||
maxDraftSizeCached int
|
||||
}
|
||||
|
||||
func draftSliceColumns() []string {
|
||||
return []string{
|
||||
"CreateAt",
|
||||
"UpdateAt",
|
||||
"DeleteAt",
|
||||
"Message",
|
||||
"RootId",
|
||||
"ChannelId",
|
||||
"UserId",
|
||||
"FileIds",
|
||||
"Props",
|
||||
"Priority",
|
||||
}
|
||||
}
|
||||
|
||||
func draftToSlice(draft *model.Draft) []interface{} {
|
||||
return []interface{}{
|
||||
draft.CreateAt,
|
||||
draft.UpdateAt,
|
||||
draft.DeleteAt,
|
||||
draft.Message,
|
||||
draft.RootId,
|
||||
draft.ChannelId,
|
||||
draft.UserId,
|
||||
model.ArrayToJSON(draft.FileIds),
|
||||
model.StringInterfaceToJSON(draft.Props),
|
||||
model.StringInterfaceToJSON(draft.Priority),
|
||||
}
|
||||
}
|
||||
|
||||
func newSqlDraftStore(sqlStore *SqlStore, metrics einterfaces.MetricsInterface) store.DraftStore {
|
||||
return &SqlDraftStore{
|
||||
SqlStore: sqlStore,
|
||||
metrics: metrics,
|
||||
maxDraftSizeCached: model.PostMessageMaxRunesV1,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *SqlDraftStore) Get(userId, channelId, rootId string, includeDeleted bool) (*model.Draft, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select(draftSliceColumns()...).
|
||||
From("Drafts").
|
||||
Where(sq.Eq{
|
||||
"UserId": userId,
|
||||
"ChannelId": channelId,
|
||||
"RootId": rootId,
|
||||
})
|
||||
|
||||
if !includeDeleted {
|
||||
query = query.Where(sq.Eq{"DeleteAt": 0})
|
||||
}
|
||||
|
||||
dt := model.Draft{}
|
||||
err := s.GetReplicaX().GetBuilder(&dt, query)
|
||||
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Draft", channelId)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to find draft with channelid = %s", channelId)
|
||||
}
|
||||
|
||||
return &dt, nil
|
||||
}
|
||||
|
||||
func (s *SqlDraftStore) Save(draft *model.Draft) (*model.Draft, error) {
|
||||
draft.PreSave()
|
||||
maxDraftSize := s.GetMaxDraftSize()
|
||||
if err := draft.IsValid(maxDraftSize); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
builder := s.getQueryBuilder().Insert("Drafts").Columns(draftSliceColumns()...).Values(draftToSlice(draft)...)
|
||||
query, args, err := builder.ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "save_draft_tosql")
|
||||
}
|
||||
|
||||
if _, err = s.GetMasterX().Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save Draft")
|
||||
}
|
||||
|
||||
return draft, nil
|
||||
}
|
||||
|
||||
func (s *SqlDraftStore) Update(draft *model.Draft) (*model.Draft, error) {
|
||||
draft.PreUpdate()
|
||||
|
||||
maxDraftSize := s.GetMaxDraftSize()
|
||||
if err := draft.IsValid(maxDraftSize); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Update("Drafts").
|
||||
Set("UpdateAt", draft.UpdateAt).
|
||||
Set("Message", draft.Message).
|
||||
Set("Props", draft.Props).
|
||||
Set("FileIds", draft.FileIds).
|
||||
Set("Priority", draft.Priority).
|
||||
Set("DeleteAt", 0).
|
||||
Where(sq.Eq{
|
||||
"UserId": draft.UserId,
|
||||
"ChannelId": draft.ChannelId,
|
||||
"RootId": draft.RootId,
|
||||
})
|
||||
|
||||
if _, err := s.GetMasterX().ExecBuilder(query); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update Draft with channelid=%s", draft.ChannelId)
|
||||
}
|
||||
|
||||
return draft, nil
|
||||
}
|
||||
|
||||
func (s *SqlDraftStore) GetDraftsForUser(userID, teamID string) ([]*model.Draft, error) {
|
||||
var drafts []*model.Draft
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select(
|
||||
"Drafts.CreateAt",
|
||||
"Drafts.UpdateAt",
|
||||
"Drafts.Message",
|
||||
"Drafts.RootId",
|
||||
"Drafts.ChannelId",
|
||||
"Drafts.UserId",
|
||||
"Drafts.FileIds",
|
||||
"Drafts.Props",
|
||||
"Drafts.Priority",
|
||||
).
|
||||
From("Drafts").
|
||||
InnerJoin("ChannelMembers ON ChannelMembers.ChannelId = Drafts.ChannelId").
|
||||
Where(sq.And{
|
||||
sq.Eq{"Drafts.DeleteAt": 0},
|
||||
sq.Eq{"Drafts.UserId": userID},
|
||||
sq.Eq{"ChannelMembers.UserId": userID},
|
||||
}).
|
||||
OrderBy("Drafts.UpdateAt DESC")
|
||||
|
||||
if teamID != "" {
|
||||
query = query.
|
||||
Join("Channels ON Drafts.ChannelId = Channels.Id").
|
||||
Where(sq.Or{
|
||||
sq.Eq{"Channels.TeamId": teamID},
|
||||
sq.Eq{"Channels.TeamId": ""},
|
||||
})
|
||||
}
|
||||
|
||||
err := s.GetReplicaX().SelectBuilder(&drafts, query)
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get user drafts")
|
||||
}
|
||||
|
||||
return drafts, nil
|
||||
}
|
||||
|
||||
func (s *SqlDraftStore) Delete(userID, channelID, rootID string) error {
|
||||
time := model.GetMillis()
|
||||
query := s.getQueryBuilder().
|
||||
Update("Drafts").
|
||||
Set("UpdateAt", time).
|
||||
Set("DeleteAt", time).
|
||||
Where(sq.Eq{
|
||||
"UserId": userID,
|
||||
"ChannelId": channelID,
|
||||
"RootId": rootID,
|
||||
})
|
||||
|
||||
sql, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to convert to sql")
|
||||
}
|
||||
|
||||
_, err = s.GetMasterX().Exec(sql, args...)
|
||||
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to delete Draft")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetMaxDraftSize returns the maximum number of runes that may be stored in a post.
|
||||
func (s *SqlDraftStore) GetMaxDraftSize() int {
|
||||
s.maxDraftSizeOnce.Do(func() {
|
||||
s.maxDraftSizeCached = s.determineMaxDraftSize()
|
||||
})
|
||||
return s.maxDraftSizeCached
|
||||
}
|
||||
|
||||
func (s *SqlDraftStore) determineMaxDraftSize() int {
|
||||
var maxDraftSizeBytes int32
|
||||
|
||||
if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
// The Draft.Message column in Postgres has historically been VARCHAR(4000), but
|
||||
// may be manually enlarged to support longer drafts.
|
||||
if err := s.GetReplicaX().Get(&maxDraftSizeBytes, `
|
||||
SELECT
|
||||
COALESCE(character_maximum_length, 0)
|
||||
FROM
|
||||
information_schema.columns
|
||||
WHERE
|
||||
table_name = 'drafts'
|
||||
AND column_name = 'message'
|
||||
`); err != nil {
|
||||
mlog.Warn("Unable to determine the maximum supported draft size", mlog.Err(err))
|
||||
}
|
||||
} else if s.DriverName() == model.DatabaseDriverMysql {
|
||||
// The Draft.Message column in MySQL has historically been TEXT, with a maximum
|
||||
// limit of 65535.
|
||||
if err := s.GetReplicaX().Get(&maxDraftSizeBytes, `
|
||||
SELECT
|
||||
COALESCE(CHARACTER_MAXIMUM_LENGTH, 0)
|
||||
FROM
|
||||
INFORMATION_SCHEMA.COLUMNS
|
||||
WHERE
|
||||
table_schema = DATABASE()
|
||||
AND table_name = 'Drafts'
|
||||
AND column_name = 'Message'
|
||||
LIMIT 0, 1
|
||||
`); err != nil {
|
||||
mlog.Warn("Unable to determine the maximum supported draft size", mlog.Err(err))
|
||||
}
|
||||
} else {
|
||||
mlog.Warn("No implementation found to determine the maximum supported draft size")
|
||||
}
|
||||
|
||||
// Assume a worst-case representation of four bytes per rune.
|
||||
maxDraftSize := int(maxDraftSizeBytes) / 4
|
||||
|
||||
mlog.Info("Draft.Message has size restrictions", mlog.Int("max_characters", maxDraftSize), mlog.Int32("max_bytes", maxDraftSizeBytes))
|
||||
|
||||
return maxDraftSize
|
||||
}
|
||||
14
server/channels/store/sqlstore/draft_store_test.go
Обычный файл
14
server/channels/store/sqlstore/draft_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestDraftStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestDraftStore)
|
||||
}
|
||||
157
server/channels/store/sqlstore/emoji_store.go
Обычный файл
157
server/channels/store/sqlstore/emoji_store.go
Обычный файл
@@ -0,0 +1,157 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/einterfaces"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlEmojiStore struct {
|
||||
*SqlStore
|
||||
metrics einterfaces.MetricsInterface
|
||||
}
|
||||
|
||||
func newSqlEmojiStore(sqlStore *SqlStore, metrics einterfaces.MetricsInterface) store.EmojiStore {
|
||||
return &SqlEmojiStore{
|
||||
SqlStore: sqlStore,
|
||||
metrics: metrics,
|
||||
}
|
||||
}
|
||||
|
||||
func (es SqlEmojiStore) Save(emoji *model.Emoji) (*model.Emoji, error) {
|
||||
emoji.PreSave()
|
||||
if err := emoji.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := es.GetMasterX().NamedExec(`INSERT INTO Emoji
|
||||
(Id, CreateAt, UpdateAt, DeleteAt, CreatorId, Name)
|
||||
VALUES
|
||||
(:Id, :CreateAt, :UpdateAt, :DeleteAt, :CreatorId, :Name)`, emoji); err != nil {
|
||||
return nil, errors.Wrap(err, "error saving emoji")
|
||||
}
|
||||
|
||||
return emoji, nil
|
||||
}
|
||||
|
||||
func (es SqlEmojiStore) Get(ctx context.Context, id string, allowFromCache bool) (*model.Emoji, error) {
|
||||
return es.getBy(ctx, "Id", id)
|
||||
}
|
||||
|
||||
func (es SqlEmojiStore) GetByName(ctx context.Context, name string, allowFromCache bool) (*model.Emoji, error) {
|
||||
return es.getBy(ctx, "Name", name)
|
||||
}
|
||||
|
||||
func (es SqlEmojiStore) GetMultipleByName(names []string) ([]*model.Emoji, error) {
|
||||
// Creating (?, ?, ?) len(names) number of times.
|
||||
keys := strings.Join(strings.Fields(strings.Repeat("? ", len(names))), ",")
|
||||
args := makeStringArgs(names)
|
||||
|
||||
emojis := []*model.Emoji{}
|
||||
if err := es.GetReplicaX().Select(&emojis,
|
||||
`SELECT
|
||||
*
|
||||
FROM
|
||||
Emoji
|
||||
WHERE
|
||||
Name IN (`+keys+`)
|
||||
AND DeleteAt = 0`, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "error getting emoji by names %v", names)
|
||||
}
|
||||
return emojis, nil
|
||||
}
|
||||
|
||||
func (es SqlEmojiStore) GetList(offset, limit int, sort string) ([]*model.Emoji, error) {
|
||||
emojis := []*model.Emoji{}
|
||||
|
||||
query := "SELECT * FROM Emoji WHERE DeleteAt = 0"
|
||||
|
||||
if sort == model.EmojiSortByName {
|
||||
query += " ORDER BY Name"
|
||||
}
|
||||
|
||||
query += " LIMIT ? OFFSET ?"
|
||||
|
||||
if err := es.GetReplicaX().Select(&emojis, query, limit, offset); err != nil {
|
||||
return nil, errors.Wrap(err, "could not get list of emojis")
|
||||
}
|
||||
return emojis, nil
|
||||
}
|
||||
|
||||
func (es SqlEmojiStore) Delete(emoji *model.Emoji, time int64) error {
|
||||
if sqlResult, err := es.GetMasterX().Exec(
|
||||
`UPDATE
|
||||
Emoji
|
||||
SET
|
||||
DeleteAt = ?,
|
||||
UpdateAt = ?
|
||||
WHERE
|
||||
Id = ?
|
||||
AND DeleteAt = 0`, time, time, emoji.Id); err != nil {
|
||||
return errors.Wrap(err, "could not delete emoji")
|
||||
} else if rows, err := sqlResult.RowsAffected(); rows == 0 {
|
||||
return store.NewErrNotFound("Emoji", emoji.Id).Wrap(err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (es SqlEmojiStore) Search(name string, prefixOnly bool, limit int) ([]*model.Emoji, error) {
|
||||
emojis := []*model.Emoji{}
|
||||
|
||||
name = sanitizeSearchTerm(name, "\\")
|
||||
|
||||
term := ""
|
||||
if !prefixOnly {
|
||||
term = "%"
|
||||
}
|
||||
|
||||
term += name + "%"
|
||||
|
||||
if err := es.GetReplicaX().Select(&emojis,
|
||||
`SELECT
|
||||
*
|
||||
FROM
|
||||
Emoji
|
||||
WHERE
|
||||
Name LIKE ?
|
||||
AND DeleteAt = 0
|
||||
ORDER BY Name
|
||||
LIMIT ?`, term, limit); err != nil {
|
||||
return nil, errors.Wrapf(err, "could not search emojis by name %s", name)
|
||||
}
|
||||
return emojis, nil
|
||||
}
|
||||
|
||||
// getBy returns one active (not deleted) emoji, found by any one column (what/key).
|
||||
func (es SqlEmojiStore) getBy(ctx context.Context, what, key string) (*model.Emoji, error) {
|
||||
var emoji model.Emoji
|
||||
|
||||
err := es.DBXFromContext(ctx).Get(&emoji,
|
||||
`SELECT
|
||||
*
|
||||
FROM
|
||||
Emoji
|
||||
WHERE
|
||||
`+what+` = ?
|
||||
AND DeleteAt = 0`, key)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Emoji", fmt.Sprintf("%s=%s", what, key))
|
||||
}
|
||||
|
||||
return nil, errors.Wrapf(err, "could not get emoji by %s with value %s", what, key)
|
||||
}
|
||||
|
||||
return &emoji, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/emoji_store_test.go
Обычный файл
14
server/channels/store/sqlstore/emoji_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestEmojiStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestEmojiStore)
|
||||
}
|
||||
801
server/channels/store/sqlstore/file_info_store.go
Обычный файл
801
server/channels/store/sqlstore/file_info_store.go
Обычный файл
@@ -0,0 +1,801 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/einterfaces"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
)
|
||||
|
||||
type fileInfoWithChannelID struct {
|
||||
Id string
|
||||
CreatorId string
|
||||
PostId string
|
||||
ChannelId string
|
||||
CreateAt int64
|
||||
UpdateAt int64
|
||||
DeleteAt int64
|
||||
Path string
|
||||
ThumbnailPath string
|
||||
PreviewPath string
|
||||
Name string
|
||||
Extension string
|
||||
Size int64
|
||||
MimeType string
|
||||
Width int
|
||||
Height int
|
||||
HasPreviewImage bool
|
||||
MiniPreview *[]byte
|
||||
Content string
|
||||
RemoteId *string
|
||||
Archived bool
|
||||
}
|
||||
|
||||
func (fi fileInfoWithChannelID) ToModel() *model.FileInfo {
|
||||
return &model.FileInfo{
|
||||
Id: fi.Id,
|
||||
CreatorId: fi.CreatorId,
|
||||
PostId: fi.PostId,
|
||||
ChannelId: fi.ChannelId,
|
||||
CreateAt: fi.CreateAt,
|
||||
UpdateAt: fi.UpdateAt,
|
||||
DeleteAt: fi.DeleteAt,
|
||||
Path: fi.Path,
|
||||
ThumbnailPath: fi.ThumbnailPath,
|
||||
PreviewPath: fi.PreviewPath,
|
||||
Name: fi.Name,
|
||||
Extension: fi.Extension,
|
||||
Size: fi.Size,
|
||||
MimeType: fi.MimeType,
|
||||
Width: fi.Width,
|
||||
Height: fi.Height,
|
||||
HasPreviewImage: fi.HasPreviewImage,
|
||||
MiniPreview: fi.MiniPreview,
|
||||
Content: fi.Content,
|
||||
RemoteId: fi.RemoteId,
|
||||
}
|
||||
}
|
||||
|
||||
type SqlFileInfoStore struct {
|
||||
*SqlStore
|
||||
metrics einterfaces.MetricsInterface
|
||||
queryFields []string
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) ClearCaches() {
|
||||
}
|
||||
|
||||
func newSqlFileInfoStore(sqlStore *SqlStore, metrics einterfaces.MetricsInterface) store.FileInfoStore {
|
||||
s := &SqlFileInfoStore{
|
||||
SqlStore: sqlStore,
|
||||
metrics: metrics,
|
||||
}
|
||||
|
||||
s.queryFields = []string{
|
||||
"FileInfo.Id",
|
||||
"FileInfo.CreatorId",
|
||||
"FileInfo.PostId",
|
||||
"FileInfo.CreateAt",
|
||||
"FileInfo.UpdateAt",
|
||||
"FileInfo.DeleteAt",
|
||||
"FileInfo.Path",
|
||||
"FileInfo.ThumbnailPath",
|
||||
"FileInfo.PreviewPath",
|
||||
"FileInfo.Name",
|
||||
"FileInfo.Extension",
|
||||
"FileInfo.Size",
|
||||
"FileInfo.MimeType",
|
||||
"FileInfo.Width",
|
||||
"FileInfo.Height",
|
||||
"FileInfo.HasPreviewImage",
|
||||
"FileInfo.MiniPreview",
|
||||
"Coalesce(FileInfo.Content, '') AS Content",
|
||||
"Coalesce(FileInfo.RemoteId, '') AS RemoteId",
|
||||
"FileInfo.Archived",
|
||||
}
|
||||
|
||||
return s
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) Save(info *model.FileInfo) (*model.FileInfo, error) {
|
||||
info.PreSave()
|
||||
if err := info.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query := `
|
||||
INSERT INTO FileInfo
|
||||
(Id, CreatorId, PostId, CreateAt, UpdateAt, DeleteAt, Path, ThumbnailPath, PreviewPath,
|
||||
Name, Extension, Size, MimeType, Width, Height, HasPreviewImage, MiniPreview, Content, RemoteId)
|
||||
VALUES
|
||||
(:Id, :CreatorId, :PostId, :CreateAt, :UpdateAt, :DeleteAt, :Path, :ThumbnailPath, :PreviewPath,
|
||||
:Name, :Extension, :Size, :MimeType, :Width, :Height, :HasPreviewImage, :MiniPreview, :Content, :RemoteId)
|
||||
`
|
||||
|
||||
if _, err := fs.GetMasterX().NamedExec(query, info); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save FileInfo")
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) GetByIds(ids []string) ([]*model.FileInfo, error) {
|
||||
query := fs.getQueryBuilder().
|
||||
Select(append(fs.queryFields, "COALESCE(P.ChannelId, '') as ChannelId")...).
|
||||
From("FileInfo").
|
||||
LeftJoin("Posts as P ON FileInfo.PostId=P.Id").
|
||||
Where(sq.Eq{"FileInfo.Id": ids}).
|
||||
Where(sq.Eq{"FileInfo.DeleteAt": 0}).
|
||||
OrderBy("FileInfo.CreateAt DESC")
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
|
||||
items := []fileInfoWithChannelID{}
|
||||
if err := fs.GetReplicaX().Select(&items, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find FileInfos")
|
||||
}
|
||||
if len(items) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
infos := make([]*model.FileInfo, 0, len(items))
|
||||
for _, item := range items {
|
||||
infos = append(infos, item.ToModel())
|
||||
}
|
||||
return infos, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) Upsert(info *model.FileInfo) (*model.FileInfo, error) {
|
||||
info.PreSave()
|
||||
if err := info.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
queryString, args, err := fs.getQueryBuilder().
|
||||
Update("FileInfo").
|
||||
SetMap(map[string]any{
|
||||
"UpdateAt": info.UpdateAt,
|
||||
"DeleteAt": info.DeleteAt,
|
||||
"Path": info.Path,
|
||||
"ThumbnailPath": info.ThumbnailPath,
|
||||
"PreviewPath": info.PreviewPath,
|
||||
"Name": info.Name,
|
||||
"Extension": info.Extension,
|
||||
"Size": info.Size,
|
||||
"MimeType": info.MimeType,
|
||||
"Width": info.Width,
|
||||
"Height": info.Height,
|
||||
"HasPreviewImage": info.HasPreviewImage,
|
||||
"MiniPreview": info.MiniPreview,
|
||||
"Content": info.Content,
|
||||
"RemoteId": info.RemoteId,
|
||||
}).
|
||||
Where(sq.Eq{"Id": info.Id}).
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
|
||||
sqlResult, err := fs.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to update FileInfo")
|
||||
}
|
||||
count, err := sqlResult.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "unable to retrieve rows affected")
|
||||
}
|
||||
if count == 0 {
|
||||
return fs.Save(info)
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) get(id string, fromMaster bool) (*model.FileInfo, error) {
|
||||
info := &model.FileInfo{}
|
||||
|
||||
query := fs.getQueryBuilder().
|
||||
Select(fs.queryFields...).
|
||||
From("FileInfo").
|
||||
Where(sq.Eq{"Id": id}).
|
||||
Where(sq.Eq{"DeleteAt": 0})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
|
||||
db := fs.GetReplicaX()
|
||||
if fromMaster {
|
||||
db = fs.GetMasterX()
|
||||
}
|
||||
|
||||
if err := db.Get(info, queryString, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("FileInfo", id)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get FileInfo with id=%s", id)
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) Get(id string) (*model.FileInfo, error) {
|
||||
return fs.get(id, false)
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) GetFromMaster(id string) (*model.FileInfo, error) {
|
||||
return fs.get(id, true)
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) GetWithOptions(page, perPage int, opt *model.GetFileInfosOptions) ([]*model.FileInfo, error) {
|
||||
if perPage < 0 {
|
||||
return nil, store.NewErrLimitExceeded("perPage", perPage, "value used in pagination while getting FileInfos")
|
||||
} else if page < 0 {
|
||||
return nil, store.NewErrLimitExceeded("page", page, "value used in pagination while getting FileInfos")
|
||||
}
|
||||
if perPage == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if opt == nil {
|
||||
opt = &model.GetFileInfosOptions{}
|
||||
}
|
||||
|
||||
query := fs.getQueryBuilder().
|
||||
Select(fs.queryFields...).
|
||||
From("FileInfo")
|
||||
|
||||
if len(opt.ChannelIds) > 0 {
|
||||
query = query.Join("Posts ON FileInfo.PostId = Posts.Id").
|
||||
Where(sq.Eq{"Posts.ChannelId": opt.ChannelIds})
|
||||
}
|
||||
|
||||
if len(opt.UserIds) > 0 {
|
||||
query = query.Where(sq.Eq{"FileInfo.CreatorId": opt.UserIds})
|
||||
}
|
||||
|
||||
if opt.Since > 0 {
|
||||
query = query.Where(sq.GtOrEq{"FileInfo.CreateAt": opt.Since})
|
||||
}
|
||||
|
||||
if !opt.IncludeDeleted {
|
||||
query = query.Where("FileInfo.DeleteAt = 0")
|
||||
}
|
||||
|
||||
if opt.SortBy == "" {
|
||||
opt.SortBy = model.FileinfoSortByCreated
|
||||
}
|
||||
sortDirection := "ASC"
|
||||
if opt.SortDescending {
|
||||
sortDirection = "DESC"
|
||||
}
|
||||
|
||||
switch opt.SortBy {
|
||||
case model.FileinfoSortByCreated:
|
||||
query = query.OrderBy("FileInfo.CreateAt " + sortDirection)
|
||||
case model.FileinfoSortBySize:
|
||||
query = query.OrderBy("FileInfo.Size " + sortDirection)
|
||||
default:
|
||||
return nil, store.NewErrInvalidInput("FileInfo", "<sortOption>", opt.SortBy)
|
||||
}
|
||||
|
||||
query = query.OrderBy("FileInfo.Id ASC") // secondary sort for sort stability
|
||||
|
||||
query = query.Limit(uint64(perPage)).Offset(uint64(perPage * page))
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
infos := []*model.FileInfo{}
|
||||
if err := fs.GetReplicaX().Select(&infos, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find FileInfos")
|
||||
}
|
||||
return infos, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) GetByPath(path string) (*model.FileInfo, error) {
|
||||
info := &model.FileInfo{}
|
||||
|
||||
query := fs.getQueryBuilder().
|
||||
Select(fs.queryFields...).
|
||||
From("FileInfo").
|
||||
Where(sq.Eq{"Path": path}).
|
||||
Where(sq.Eq{"DeleteAt": 0}).
|
||||
Limit(1)
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
|
||||
if err := fs.GetReplicaX().Get(info, queryString, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("FileInfo", fmt.Sprintf("path=%s", path))
|
||||
}
|
||||
|
||||
return nil, errors.Wrapf(err, "failed to get FileInfo with path=%s", path)
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) InvalidateFileInfosForPostCache(postId string, deleted bool) {
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) GetForPost(postId string, readFromMaster, includeDeleted, allowFromCache bool) ([]*model.FileInfo, error) {
|
||||
infos := []*model.FileInfo{}
|
||||
|
||||
dbmap := fs.GetReplicaX()
|
||||
|
||||
if readFromMaster {
|
||||
dbmap = fs.GetMasterX()
|
||||
}
|
||||
|
||||
query := fs.getQueryBuilder().
|
||||
Select(fs.queryFields...).
|
||||
From("FileInfo").
|
||||
Where(sq.Eq{"PostId": postId}).
|
||||
OrderBy("CreateAt")
|
||||
|
||||
if !includeDeleted {
|
||||
query = query.Where("DeleteAt = 0")
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
|
||||
if err := dbmap.Select(&infos, queryString, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find FileInfos with postId=%s", postId)
|
||||
}
|
||||
return infos, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) GetForUser(userId string) ([]*model.FileInfo, error) {
|
||||
infos := []*model.FileInfo{}
|
||||
|
||||
query := fs.getQueryBuilder().
|
||||
Select(fs.queryFields...).
|
||||
From("FileInfo").
|
||||
Where(sq.Eq{"CreatorId": userId}).
|
||||
Where(sq.Eq{"DeleteAt": 0}).
|
||||
OrderBy("CreateAt")
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
|
||||
if err := fs.GetReplicaX().Select(&infos, queryString, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find FileInfos with creatorId=%s", userId)
|
||||
}
|
||||
return infos, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) AttachToPost(fileId, postId, creatorId string) error {
|
||||
query := fs.getQueryBuilder().
|
||||
Update("FileInfo").
|
||||
Set("PostId", postId).
|
||||
Where(sq.And{
|
||||
sq.Eq{"Id": fileId},
|
||||
sq.Eq{"PostId": ""},
|
||||
sq.Or{
|
||||
sq.Eq{"CreatorId": creatorId},
|
||||
sq.Eq{"CreatorId": "nouser"},
|
||||
},
|
||||
})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
sqlResult, err := fs.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to update FileInfo with id=%s and postId=%s", fileId, postId)
|
||||
}
|
||||
|
||||
count, err := sqlResult.RowsAffected()
|
||||
if err != nil {
|
||||
// RowsAffected should never fail with the MySQL or Postgres drivers
|
||||
return errors.Wrap(err, "unable to retrieve rows affected")
|
||||
} else if count == 0 {
|
||||
// Could not attach the file to the post
|
||||
return store.NewErrInvalidInput("FileInfo", "<id, postId, creatorId>", fmt.Sprintf("<%s, %s, %s>", fileId, postId, creatorId))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) SetContent(fileId, content string) error {
|
||||
query := fs.getQueryBuilder().
|
||||
Update("FileInfo").
|
||||
Set("Content", content).
|
||||
Where(sq.Eq{"Id": fileId})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
|
||||
_, err = fs.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to update FileInfo content with id=%s", fileId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) DeleteForPost(postId string) (string, error) {
|
||||
if _, err := fs.GetMasterX().Exec(
|
||||
`UPDATE
|
||||
FileInfo
|
||||
SET
|
||||
DeleteAt = ?
|
||||
WHERE
|
||||
PostId = ?`, model.GetMillis(), postId); err != nil {
|
||||
return "", errors.Wrapf(err, "failed to update FileInfo with postId=%s", postId)
|
||||
}
|
||||
return postId, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) PermanentDelete(fileId string) error {
|
||||
if _, err := fs.GetMasterX().Exec(`DELETE FROM FileInfo WHERE Id = ?`, fileId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete FileInfo with id=%s", fileId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) PermanentDeleteBatch(endTime int64, limit int64) (int64, error) {
|
||||
var query string
|
||||
if fs.DriverName() == "postgres" {
|
||||
query = "DELETE from FileInfo WHERE Id = any (array (SELECT Id FROM FileInfo WHERE CreateAt < ? LIMIT ?))"
|
||||
} else {
|
||||
query = "DELETE from FileInfo WHERE CreateAt < ? LIMIT ?"
|
||||
}
|
||||
|
||||
sqlResult, err := fs.GetMasterX().Exec(query, endTime, limit)
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "failed to delete FileInfos in batch")
|
||||
}
|
||||
|
||||
rowsAffected, err := sqlResult.RowsAffected()
|
||||
if err != nil {
|
||||
return 0, errors.Wrapf(err, "unable to retrieve rows affected")
|
||||
}
|
||||
|
||||
return rowsAffected, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) PermanentDeleteByUser(userId string) (int64, error) {
|
||||
query := "DELETE from FileInfo WHERE CreatorId = ?"
|
||||
|
||||
sqlResult, err := fs.GetMasterX().Exec(query, userId)
|
||||
if err != nil {
|
||||
return 0, errors.Wrapf(err, "failed to delete FileInfo with creatorId=%s", userId)
|
||||
}
|
||||
|
||||
rowsAffected, err := sqlResult.RowsAffected()
|
||||
if err != nil {
|
||||
return 0, errors.Wrapf(err, "unable to retrieve rows affected")
|
||||
}
|
||||
|
||||
return rowsAffected, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) Search(paramsList []*model.SearchParams, userId, teamId string, page, perPage int) (*model.FileInfoList, error) {
|
||||
// Since we don't support paging for DB search, we just return nothing for later pages
|
||||
if page > 0 {
|
||||
return model.NewFileInfoList(), nil
|
||||
}
|
||||
if err := model.IsSearchParamsListValid(paramsList); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
query := fs.getQueryBuilder().
|
||||
Select(append(fs.queryFields, "Coalesce(P.ChannelId, '') AS ChannelId")...).
|
||||
From("FileInfo").
|
||||
LeftJoin("Posts as P ON FileInfo.PostId=P.Id").
|
||||
LeftJoin("Channels as C ON C.Id=P.ChannelId").
|
||||
LeftJoin("ChannelMembers as CM ON C.Id=CM.ChannelId").
|
||||
Where(sq.Eq{"FileInfo.DeleteAt": 0}).
|
||||
OrderBy("FileInfo.CreateAt DESC").
|
||||
Limit(100)
|
||||
|
||||
if teamId != "" {
|
||||
query = query.Where(sq.Or{
|
||||
sq.Eq{"C.TeamId": teamId},
|
||||
sq.Eq{"C.TeamId": ""},
|
||||
})
|
||||
}
|
||||
|
||||
now := model.GetMillis()
|
||||
for _, params := range paramsList {
|
||||
if params.Modifier == model.ModifierFiles {
|
||||
// Deliberately keeping non-alphanumeric characters to
|
||||
// prevent surprises in UI.
|
||||
buf, err := json.Marshal(params)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
err = fs.stores.post.LogRecentSearch(userId, buf, now)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
params.Terms = removeNonAlphaNumericUnquotedTerms(params.Terms, " ")
|
||||
|
||||
if !params.IncludeDeletedChannels {
|
||||
query = query.Where(sq.Eq{"C.DeleteAt": 0})
|
||||
}
|
||||
|
||||
if !params.SearchWithoutUserId {
|
||||
query = query.Where(sq.Eq{"CM.UserId": userId})
|
||||
}
|
||||
|
||||
if len(params.InChannels) != 0 {
|
||||
query = query.Where(sq.Eq{"C.Id": params.InChannels})
|
||||
}
|
||||
|
||||
if len(params.Extensions) != 0 {
|
||||
query = query.Where(sq.Eq{"FileInfo.Extension": params.Extensions})
|
||||
}
|
||||
|
||||
if len(params.ExcludedExtensions) != 0 {
|
||||
query = query.Where(sq.NotEq{"FileInfo.Extension": params.ExcludedExtensions})
|
||||
}
|
||||
|
||||
if len(params.ExcludedChannels) != 0 {
|
||||
query = query.Where(sq.NotEq{"C.Id": params.ExcludedChannels})
|
||||
}
|
||||
|
||||
if len(params.FromUsers) != 0 {
|
||||
query = query.Where(sq.Eq{"FileInfo.CreatorId": params.FromUsers})
|
||||
}
|
||||
|
||||
if len(params.ExcludedUsers) != 0 {
|
||||
query = query.Where(sq.NotEq{"FileInfo.CreatorId": params.ExcludedUsers})
|
||||
}
|
||||
|
||||
// handle after: before: on: filters
|
||||
if params.OnDate != "" {
|
||||
onDateStart, onDateEnd := params.GetOnDateMillis()
|
||||
query = query.Where(sq.Expr("FileInfo.CreateAt BETWEEN ? AND ?", strconv.FormatInt(onDateStart, 10), strconv.FormatInt(onDateEnd, 10)))
|
||||
} else {
|
||||
if params.ExcludedDate != "" {
|
||||
excludedDateStart, excludedDateEnd := params.GetExcludedDateMillis()
|
||||
query = query.Where(sq.Expr("FileInfo.CreateAt NOT BETWEEN ? AND ?", strconv.FormatInt(excludedDateStart, 10), strconv.FormatInt(excludedDateEnd, 10)))
|
||||
}
|
||||
|
||||
if params.AfterDate != "" {
|
||||
afterDate := params.GetAfterDateMillis()
|
||||
query = query.Where(sq.GtOrEq{"FileInfo.CreateAt": strconv.FormatInt(afterDate, 10)})
|
||||
}
|
||||
|
||||
if params.BeforeDate != "" {
|
||||
beforeDate := params.GetBeforeDateMillis()
|
||||
query = query.Where(sq.LtOrEq{"FileInfo.CreateAt": strconv.FormatInt(beforeDate, 10)})
|
||||
}
|
||||
|
||||
if params.ExcludedAfterDate != "" {
|
||||
afterDate := params.GetExcludedAfterDateMillis()
|
||||
query = query.Where(sq.Lt{"FileInfo.CreateAt": strconv.FormatInt(afterDate, 10)})
|
||||
}
|
||||
|
||||
if params.ExcludedBeforeDate != "" {
|
||||
beforeDate := params.GetExcludedBeforeDateMillis()
|
||||
query = query.Where(sq.Gt{"FileInfo.CreateAt": strconv.FormatInt(beforeDate, 10)})
|
||||
}
|
||||
}
|
||||
|
||||
terms := params.Terms
|
||||
excludedTerms := params.ExcludedTerms
|
||||
|
||||
for _, c := range fs.specialSearchChars() {
|
||||
terms = strings.Replace(terms, c, " ", -1)
|
||||
excludedTerms = strings.Replace(excludedTerms, c, " ", -1)
|
||||
}
|
||||
|
||||
if terms == "" && excludedTerms == "" {
|
||||
// we've already confirmed that we have a channel or user to search for
|
||||
} else if fs.DriverName() == model.DatabaseDriverPostgres {
|
||||
// Parse text for wildcards
|
||||
if wildcard, err := regexp.Compile(`\*($| )`); err == nil {
|
||||
terms = wildcard.ReplaceAllLiteralString(terms, ":* ")
|
||||
excludedTerms = wildcard.ReplaceAllLiteralString(excludedTerms, ":* ")
|
||||
}
|
||||
|
||||
excludeClause := ""
|
||||
if excludedTerms != "" {
|
||||
excludeClause = " & !(" + strings.Join(strings.Fields(excludedTerms), " | ") + ")"
|
||||
}
|
||||
|
||||
queryTerms := ""
|
||||
if params.OrTerms {
|
||||
queryTerms = "(" + strings.Join(strings.Fields(terms), " | ") + ")" + excludeClause
|
||||
} else {
|
||||
queryTerms = "(" + strings.Join(strings.Fields(terms), " & ") + ")" + excludeClause
|
||||
}
|
||||
|
||||
query = query.Where(sq.Or{
|
||||
sq.Expr(fmt.Sprintf("to_tsvector('%[1]s', FileInfo.Name) @@ to_tsquery('%[1]s', ?)", fs.pgDefaultTextSearchConfig), queryTerms),
|
||||
sq.Expr(fmt.Sprintf("to_tsvector('%[1]s', Translate(FileInfo.Name, '.,-', ' ')) @@ to_tsquery('%[1]s', ?)", fs.pgDefaultTextSearchConfig), queryTerms),
|
||||
sq.Expr(fmt.Sprintf("to_tsvector('%[1]s', FileInfo.Content) @@ to_tsquery('%[1]s', ?)", fs.pgDefaultTextSearchConfig), queryTerms),
|
||||
})
|
||||
} else if fs.DriverName() == model.DatabaseDriverMysql {
|
||||
var err error
|
||||
terms, err = removeMysqlStopWordsFromTerms(terms)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to remove Mysql stop-words from terms")
|
||||
}
|
||||
|
||||
if terms == "" {
|
||||
return model.NewFileInfoList(), nil
|
||||
}
|
||||
|
||||
excludeClause := ""
|
||||
if excludedTerms != "" {
|
||||
excludeClause = " -(" + excludedTerms + ")"
|
||||
}
|
||||
|
||||
queryTerms := ""
|
||||
if params.OrTerms {
|
||||
queryTerms = terms + excludeClause
|
||||
} else {
|
||||
splitTerms := []string{}
|
||||
for _, t := range strings.Fields(terms) {
|
||||
splitTerms = append(splitTerms, "+"+t)
|
||||
}
|
||||
queryTerms = strings.Join(splitTerms, " ") + excludeClause
|
||||
}
|
||||
query = query.Where(sq.Or{
|
||||
sq.Expr("MATCH (FileInfo.Name) AGAINST (? IN BOOLEAN MODE)", queryTerms),
|
||||
sq.Expr("MATCH (FileInfo.Content) AGAINST (? IN BOOLEAN MODE)", queryTerms),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "file_info_tosql")
|
||||
}
|
||||
|
||||
list := model.NewFileInfoList()
|
||||
|
||||
items := []fileInfoWithChannelID{}
|
||||
err = fs.GetSearchReplicaX().Select(&items, queryString, args...)
|
||||
if err != nil {
|
||||
mlog.Warn("Query error searching files.", 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 {
|
||||
for _, item := range items {
|
||||
info := item.ToModel()
|
||||
list.AddFileInfo(info)
|
||||
list.AddOrder(info.Id)
|
||||
}
|
||||
}
|
||||
list.MakeNonNil()
|
||||
return list, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) CountAll() (int64, error) {
|
||||
query := fs.getQueryBuilder().
|
||||
Select("COUNT(*)").
|
||||
From("FileInfo").
|
||||
Where("DeleteAt = 0")
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrap(err, "count_tosql")
|
||||
}
|
||||
|
||||
var count int64
|
||||
err = fs.GetReplicaX().Get(&count, queryString, args...)
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrap(err, "failed to count Files")
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) GetFilesBatchForIndexing(startTime int64, startFileID string, limit int) ([]*model.FileForIndexing, error) {
|
||||
files := []*model.FileForIndexing{}
|
||||
sql, args, _ := fs.getQueryBuilder().
|
||||
Select(append(fs.queryFields, "Coalesce(p.ChannelId, '') AS ChannelId")...).
|
||||
From("FileInfo").
|
||||
LeftJoin("Posts AS p ON FileInfo.PostId = p.Id").
|
||||
Where(sq.Or{
|
||||
sq.Gt{"FileInfo.CreateAt": startTime},
|
||||
sq.And{
|
||||
sq.Eq{"FileInfo.CreateAt": startTime},
|
||||
sq.Gt{"FileInfo.Id": startFileID},
|
||||
},
|
||||
}).
|
||||
OrderBy("FileInfo.CreateAt ASC, FileInfo.Id ASC").
|
||||
Limit(uint64(limit)).
|
||||
ToSql()
|
||||
|
||||
err := fs.GetSearchReplicaX().Select(&files, sql, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Files")
|
||||
}
|
||||
|
||||
return files, nil
|
||||
}
|
||||
|
||||
func (fs SqlFileInfoStore) GetStorageUsage(allowFromCache, includeDeleted bool) (int64, error) {
|
||||
query := fs.getQueryBuilder().
|
||||
Select("COALESCE(SUM(Size), 0)").
|
||||
From("FileInfo")
|
||||
|
||||
if !includeDeleted {
|
||||
query = query.Where("DeleteAt = 0")
|
||||
}
|
||||
|
||||
var size int64
|
||||
err := fs.GetReplicaX().GetBuilder(&size, query)
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrap(err, "failed to get storage usage")
|
||||
}
|
||||
return size, nil
|
||||
}
|
||||
|
||||
// GetUptoNSizeFileTime returns the CreateAt time of the last accessible file with a running-total size upto n bytes.
|
||||
func (fs *SqlFileInfoStore) GetUptoNSizeFileTime(n int64) (int64, error) {
|
||||
if n <= 0 {
|
||||
return 0, errors.New("n can't be less than 1")
|
||||
}
|
||||
|
||||
var sizeSubQuery sq.SelectBuilder
|
||||
// Separate query for MySql, as current min-version 5.x doesn't support window-functions
|
||||
if fs.DriverName() == model.DatabaseDriverMysql {
|
||||
sizeSubQuery = sq.
|
||||
Select("(@runningSum := @runningSum + fi.Size) RunningTotal", "fi.CreateAt").
|
||||
From("FileInfo fi").
|
||||
Join("(SELECT @runningSum := 0) as tmp").
|
||||
Where(sq.Eq{"fi.DeleteAt": 0}).
|
||||
OrderBy("fi.CreateAt DESC, fi.Id")
|
||||
} else {
|
||||
sizeSubQuery = sq.
|
||||
Select("SUM(fi.Size) OVER(ORDER BY CreateAt DESC, fi.Id) RunningTotal", "fi.CreateAt").
|
||||
From("FileInfo fi").
|
||||
Where(sq.Eq{"fi.DeleteAt": 0})
|
||||
}
|
||||
|
||||
builder := fs.getQueryBuilder().
|
||||
Select("fi2.CreateAt").
|
||||
FromSelect(sizeSubQuery, "fi2").
|
||||
Where(sq.LtOrEq{"fi2.RunningTotal": n}).
|
||||
OrderBy("fi2.CreateAt").
|
||||
Limit(1)
|
||||
|
||||
query, queryArgs, err := builder.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "GetUptoNSizeFileTime_tosql")
|
||||
}
|
||||
|
||||
var createAt int64
|
||||
if err := fs.GetReplicaX().Get(&createAt, query, queryArgs...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return 0, store.NewErrNotFound("File", "none")
|
||||
}
|
||||
|
||||
return 0, errors.Wrapf(err, "failed to get the File for size upto=%d", n)
|
||||
}
|
||||
|
||||
return createAt, nil
|
||||
}
|
||||
19
server/channels/store/sqlstore/file_info_store_test.go
Обычный файл
19
server/channels/store/sqlstore/file_info_store_test.go
Обычный файл
@@ -0,0 +1,19 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/searchtest"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestFileInfoStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestFileInfoStore)
|
||||
}
|
||||
|
||||
func TestSearchFileInfoStore(t *testing.T) {
|
||||
StoreTestWithSearchTestEngine(t, searchtest.TestSearchFileInfoStore)
|
||||
}
|
||||
2025
server/channels/store/sqlstore/group_store.go
Обычный файл
2025
server/channels/store/sqlstore/group_store.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
14
server/channels/store/sqlstore/group_store_test.go
Обычный файл
14
server/channels/store/sqlstore/group_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestGroupStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestGroupStore)
|
||||
}
|
||||
12
server/channels/store/sqlstore/init_test.go
Обычный файл
12
server/channels/store/sqlstore/init_test.go
Обычный файл
@@ -0,0 +1,12 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
func InitTest() {
|
||||
initStores()
|
||||
}
|
||||
|
||||
func TearDownTest() {
|
||||
tearDownStores()
|
||||
}
|
||||
536
server/channels/store/sqlstore/integrity.go
Обычный файл
536
server/channels/store/sqlstore/integrity.go
Обычный файл
@@ -0,0 +1,536 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
sq "github.com/mattermost/squirrel"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
)
|
||||
|
||||
type relationalCheckConfig struct {
|
||||
parentName string
|
||||
parentIdAttr string
|
||||
childName string
|
||||
childIdAttr string
|
||||
canParentIdBeEmpty bool
|
||||
sortRecords bool
|
||||
filter any
|
||||
}
|
||||
|
||||
func getOrphanedRecords(ss *SqlStore, cfg relationalCheckConfig) ([]model.OrphanedRecord, error) {
|
||||
records := []model.OrphanedRecord{}
|
||||
|
||||
sub := ss.getQueryBuilder().
|
||||
Select("TRUE").
|
||||
From(cfg.parentName + " AS PT").
|
||||
Prefix("NOT EXISTS (").
|
||||
Suffix(")").
|
||||
Where("PT.id = CT." + cfg.parentIdAttr)
|
||||
|
||||
main := ss.getQueryBuilder().
|
||||
Select().
|
||||
Column("CT." + cfg.parentIdAttr + " AS ParentId").
|
||||
From(cfg.childName + " AS CT").
|
||||
Where(sub)
|
||||
|
||||
if cfg.childIdAttr != "" {
|
||||
main = main.Column("CT." + cfg.childIdAttr + " AS ChildId")
|
||||
}
|
||||
|
||||
if cfg.canParentIdBeEmpty {
|
||||
main = main.Where(sq.NotEq{"CT." + cfg.parentIdAttr: ""})
|
||||
}
|
||||
|
||||
if cfg.filter != nil {
|
||||
main = main.Where(cfg.filter)
|
||||
}
|
||||
|
||||
if cfg.sortRecords {
|
||||
main = main.OrderBy("CT." + cfg.parentIdAttr)
|
||||
}
|
||||
|
||||
query, args, err := main.ToSql()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
err = ss.GetMasterX().Select(&records, query, args...)
|
||||
return records, err
|
||||
}
|
||||
|
||||
func checkParentChildIntegrity(ss *SqlStore, config relationalCheckConfig) model.IntegrityCheckResult {
|
||||
var result model.IntegrityCheckResult
|
||||
var data model.RelationalIntegrityCheckData
|
||||
|
||||
config.sortRecords = true
|
||||
data.Records, result.Err = getOrphanedRecords(ss, config)
|
||||
if result.Err != nil {
|
||||
mlog.Error("Error while getting orphaned records", mlog.Err(result.Err))
|
||||
return result
|
||||
}
|
||||
data.ParentName = config.parentName
|
||||
data.ChildName = config.childName
|
||||
data.ParentIdAttr = config.parentIdAttr
|
||||
data.ChildIdAttr = config.childIdAttr
|
||||
result.Data = data
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func checkChannelsCommandWebhooksIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Channels",
|
||||
parentIdAttr: "ChannelId",
|
||||
childName: "CommandWebhooks",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkChannelsChannelMemberHistoryIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Channels",
|
||||
parentIdAttr: "ChannelId",
|
||||
childName: "ChannelMemberHistory",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkChannelsChannelMembersIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Channels",
|
||||
parentIdAttr: "ChannelId",
|
||||
childName: "ChannelMembers",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkChannelsIncomingWebhooksIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Channels",
|
||||
parentIdAttr: "ChannelId",
|
||||
childName: "IncomingWebhooks",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkChannelsOutgoingWebhooksIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Channels",
|
||||
parentIdAttr: "ChannelId",
|
||||
childName: "OutgoingWebhooks",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkChannelsPostsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Channels",
|
||||
parentIdAttr: "ChannelId",
|
||||
childName: "Posts",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkCommandsCommandWebhooksIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Commands",
|
||||
parentIdAttr: "CommandId",
|
||||
childName: "CommandWebhooks",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkPostsFileInfoIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Posts",
|
||||
parentIdAttr: "PostId",
|
||||
childName: "FileInfo",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkPostsPostsRootIdIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Posts",
|
||||
parentIdAttr: "RootId",
|
||||
childName: "Posts",
|
||||
childIdAttr: "Id",
|
||||
canParentIdBeEmpty: true,
|
||||
})
|
||||
}
|
||||
|
||||
func checkPostsReactionsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Posts",
|
||||
parentIdAttr: "PostId",
|
||||
childName: "Reactions",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkSchemesChannelsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Schemes",
|
||||
parentIdAttr: "SchemeId",
|
||||
childName: "Channels",
|
||||
childIdAttr: "Id",
|
||||
canParentIdBeEmpty: true,
|
||||
})
|
||||
}
|
||||
|
||||
func checkSchemesTeamsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Schemes",
|
||||
parentIdAttr: "SchemeId",
|
||||
childName: "Teams",
|
||||
childIdAttr: "Id",
|
||||
canParentIdBeEmpty: true,
|
||||
})
|
||||
}
|
||||
|
||||
func checkSessionsAuditsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Sessions",
|
||||
parentIdAttr: "SessionId",
|
||||
childName: "Audits",
|
||||
childIdAttr: "Id",
|
||||
canParentIdBeEmpty: true,
|
||||
})
|
||||
}
|
||||
|
||||
func checkTeamsChannelsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
res1 := checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Teams",
|
||||
parentIdAttr: "TeamId",
|
||||
childName: "Channels",
|
||||
childIdAttr: "Id",
|
||||
filter: sq.NotEq{"CT.Type": []model.ChannelType{model.ChannelTypeDirect, model.ChannelTypeGroup}},
|
||||
})
|
||||
res2 := checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Teams",
|
||||
parentIdAttr: "TeamId",
|
||||
childName: "Channels",
|
||||
childIdAttr: "Id",
|
||||
canParentIdBeEmpty: true,
|
||||
filter: sq.Eq{"CT.Type": []model.ChannelType{model.ChannelTypeDirect, model.ChannelTypeGroup}},
|
||||
})
|
||||
data1 := res1.Data.(model.RelationalIntegrityCheckData)
|
||||
data2 := res2.Data.(model.RelationalIntegrityCheckData)
|
||||
data1.Records = append(data1.Records, data2.Records...)
|
||||
res1.Data = data1
|
||||
return res1
|
||||
}
|
||||
|
||||
func checkTeamsCommandsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Teams",
|
||||
parentIdAttr: "TeamId",
|
||||
childName: "Commands",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkTeamsIncomingWebhooksIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Teams",
|
||||
parentIdAttr: "TeamId",
|
||||
childName: "IncomingWebhooks",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkTeamsOutgoingWebhooksIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Teams",
|
||||
parentIdAttr: "TeamId",
|
||||
childName: "OutgoingWebhooks",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkTeamsTeamMembersIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Teams",
|
||||
parentIdAttr: "TeamId",
|
||||
childName: "TeamMembers",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersAuditsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "Audits",
|
||||
childIdAttr: "Id",
|
||||
canParentIdBeEmpty: true,
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersCommandWebhooksIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "CommandWebhooks",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersChannelMemberHistoryIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "ChannelMemberHistory",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersChannelMembersIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "ChannelMembers",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersChannelsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "CreatorId",
|
||||
childName: "Channels",
|
||||
childIdAttr: "Id",
|
||||
canParentIdBeEmpty: true,
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersCommandsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "CreatorId",
|
||||
childName: "Commands",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersCompliancesIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "Compliances",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersEmojiIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "CreatorId",
|
||||
childName: "Emoji",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersFileInfoIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "CreatorId",
|
||||
childName: "FileInfo",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersIncomingWebhooksIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "IncomingWebhooks",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersOAuthAccessDataIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "OAuthAccessData",
|
||||
childIdAttr: "Token",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersOAuthAppsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "CreatorId",
|
||||
childName: "OAuthApps",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersOAuthAuthDataIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "OAuthAuthData",
|
||||
childIdAttr: "Code",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersOutgoingWebhooksIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "CreatorId",
|
||||
childName: "OutgoingWebhooks",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersPostsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "Posts",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersPreferencesIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "Preferences",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersReactionsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "Reactions",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersSessionsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "Sessions",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersStatusIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "Status",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersTeamMembersIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "TeamMembers",
|
||||
childIdAttr: "",
|
||||
})
|
||||
}
|
||||
|
||||
func checkUsersUserAccessTokensIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Users",
|
||||
parentIdAttr: "UserId",
|
||||
childName: "UserAccessTokens",
|
||||
childIdAttr: "Id",
|
||||
})
|
||||
}
|
||||
|
||||
func checkChannelsIntegrity(ss *SqlStore, results chan<- model.IntegrityCheckResult) {
|
||||
results <- checkChannelsCommandWebhooksIntegrity(ss)
|
||||
results <- checkChannelsChannelMemberHistoryIntegrity(ss)
|
||||
results <- checkChannelsChannelMembersIntegrity(ss)
|
||||
results <- checkChannelsIncomingWebhooksIntegrity(ss)
|
||||
results <- checkChannelsOutgoingWebhooksIntegrity(ss)
|
||||
results <- checkChannelsPostsIntegrity(ss)
|
||||
}
|
||||
|
||||
func checkCommandsIntegrity(ss *SqlStore, results chan<- model.IntegrityCheckResult) {
|
||||
results <- checkCommandsCommandWebhooksIntegrity(ss)
|
||||
}
|
||||
|
||||
func checkPostsIntegrity(ss *SqlStore, results chan<- model.IntegrityCheckResult) {
|
||||
results <- checkPostsFileInfoIntegrity(ss)
|
||||
results <- checkPostsPostsRootIdIntegrity(ss)
|
||||
results <- checkPostsReactionsIntegrity(ss)
|
||||
results <- checkThreadsTeamsIntegrity(ss)
|
||||
}
|
||||
|
||||
func checkSchemesIntegrity(ss *SqlStore, results chan<- model.IntegrityCheckResult) {
|
||||
results <- checkSchemesChannelsIntegrity(ss)
|
||||
results <- checkSchemesTeamsIntegrity(ss)
|
||||
}
|
||||
|
||||
func checkSessionsIntegrity(ss *SqlStore, results chan<- model.IntegrityCheckResult) {
|
||||
results <- checkSessionsAuditsIntegrity(ss)
|
||||
}
|
||||
|
||||
func checkTeamsIntegrity(ss *SqlStore, results chan<- model.IntegrityCheckResult) {
|
||||
results <- checkTeamsChannelsIntegrity(ss)
|
||||
results <- checkTeamsCommandsIntegrity(ss)
|
||||
results <- checkTeamsIncomingWebhooksIntegrity(ss)
|
||||
results <- checkTeamsOutgoingWebhooksIntegrity(ss)
|
||||
results <- checkTeamsTeamMembersIntegrity(ss)
|
||||
}
|
||||
|
||||
func checkUsersIntegrity(ss *SqlStore, results chan<- model.IntegrityCheckResult) {
|
||||
results <- checkUsersAuditsIntegrity(ss)
|
||||
results <- checkUsersCommandWebhooksIntegrity(ss)
|
||||
results <- checkUsersChannelMemberHistoryIntegrity(ss)
|
||||
results <- checkUsersChannelMembersIntegrity(ss)
|
||||
results <- checkUsersChannelsIntegrity(ss)
|
||||
results <- checkUsersCommandsIntegrity(ss)
|
||||
results <- checkUsersCompliancesIntegrity(ss)
|
||||
results <- checkUsersEmojiIntegrity(ss)
|
||||
results <- checkUsersFileInfoIntegrity(ss)
|
||||
results <- checkUsersIncomingWebhooksIntegrity(ss)
|
||||
results <- checkUsersOAuthAccessDataIntegrity(ss)
|
||||
results <- checkUsersOAuthAppsIntegrity(ss)
|
||||
results <- checkUsersOAuthAuthDataIntegrity(ss)
|
||||
results <- checkUsersOutgoingWebhooksIntegrity(ss)
|
||||
results <- checkUsersPostsIntegrity(ss)
|
||||
results <- checkUsersPreferencesIntegrity(ss)
|
||||
results <- checkUsersReactionsIntegrity(ss)
|
||||
results <- checkUsersSessionsIntegrity(ss)
|
||||
results <- checkUsersStatusIntegrity(ss)
|
||||
results <- checkUsersTeamMembersIntegrity(ss)
|
||||
results <- checkUsersUserAccessTokensIntegrity(ss)
|
||||
}
|
||||
|
||||
func checkThreadsTeamsIntegrity(ss *SqlStore) model.IntegrityCheckResult {
|
||||
return checkParentChildIntegrity(ss, relationalCheckConfig{
|
||||
parentName: "Teams",
|
||||
parentIdAttr: "ThreadTeamId",
|
||||
childName: "Threads",
|
||||
childIdAttr: "PostId",
|
||||
canParentIdBeEmpty: false,
|
||||
})
|
||||
}
|
||||
|
||||
func CheckRelationalIntegrity(ss *SqlStore, results chan<- model.IntegrityCheckResult) {
|
||||
mlog.Info("Starting relational integrity checks...")
|
||||
checkChannelsIntegrity(ss, results)
|
||||
checkCommandsIntegrity(ss, results)
|
||||
checkPostsIntegrity(ss, results)
|
||||
checkSchemesIntegrity(ss, results)
|
||||
checkSessionsIntegrity(ss, results)
|
||||
checkTeamsIntegrity(ss, results)
|
||||
checkUsersIntegrity(ss, results)
|
||||
mlog.Info("Done with relational integrity checks")
|
||||
close(results)
|
||||
}
|
||||
1643
server/channels/store/sqlstore/integrity_test.go
Обычный файл
1643
server/channels/store/sqlstore/integrity_test.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
344
server/channels/store/sqlstore/job_store.go
Обычный файл
344
server/channels/store/sqlstore/job_store.go
Обычный файл
@@ -0,0 +1,344 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
const (
|
||||
jobsCleanupDelay = 100 * time.Millisecond
|
||||
)
|
||||
|
||||
type SqlJobStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlJobStore(sqlStore *SqlStore) store.JobStore {
|
||||
return &SqlJobStore{sqlStore}
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) Save(job *model.Job) (*model.Job, error) {
|
||||
jsonData, err := json.Marshal(job.Data)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed marshalling job data")
|
||||
}
|
||||
if jss.IsBinaryParamEnabled() {
|
||||
jsonData = AppendBinaryFlag(jsonData)
|
||||
}
|
||||
query := jss.getQueryBuilder().
|
||||
Insert("Jobs").
|
||||
Columns("Id", "Type", "Priority", "CreateAt", "StartAt", "LastActivityAt", "Status", "Progress", "Data").
|
||||
Values(job.Id, job.Type, job.Priority, job.CreateAt, job.StartAt, job.LastActivityAt, job.Status, job.Progress, jsonData)
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to generate sqlquery")
|
||||
}
|
||||
|
||||
if _, err = jss.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save Preference")
|
||||
}
|
||||
|
||||
return job, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) UpdateOptimistically(job *model.Job, currentStatus string) (bool, error) {
|
||||
dataJSON, jsonErr := json.Marshal(job.Data)
|
||||
if jsonErr != nil {
|
||||
return false, errors.Wrap(jsonErr, "failed to encode job's data to JSON")
|
||||
}
|
||||
if jss.IsBinaryParamEnabled() {
|
||||
dataJSON = AppendBinaryFlag(dataJSON)
|
||||
}
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Update("Jobs").
|
||||
Set("LastActivityAt", model.GetMillis()).
|
||||
Set("Status", job.Status).
|
||||
Set("Data", dataJSON).
|
||||
Set("Progress", job.Progress).
|
||||
Where(sq.Eq{"Id": job.Id, "Status": currentStatus}).ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
sqlResult, err := jss.GetMasterX().Exec(query, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to update Job")
|
||||
}
|
||||
|
||||
rows, err := sqlResult.RowsAffected()
|
||||
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "unable to get rows affected")
|
||||
}
|
||||
|
||||
if rows != 1 {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) UpdateStatus(id string, status string) (*model.Job, error) {
|
||||
job := &model.Job{
|
||||
Id: id,
|
||||
Status: status,
|
||||
LastActivityAt: model.GetMillis(),
|
||||
}
|
||||
|
||||
if _, err := jss.GetMasterX().NamedExec(`UPDATE Jobs
|
||||
SET Status=:Status, LastActivityAt=:LastActivityAt
|
||||
WHERE Id=:Id`, job); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update Job with id=%s", id)
|
||||
}
|
||||
|
||||
return job, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) UpdateStatusOptimistically(id string, currentStatus string, newStatus string) (bool, error) {
|
||||
builder := jss.getQueryBuilder().
|
||||
Update("Jobs").
|
||||
Set("LastActivityAt", model.GetMillis()).
|
||||
Set("Status", newStatus).
|
||||
Where(sq.Eq{"Id": id, "Status": currentStatus})
|
||||
|
||||
if newStatus == model.JobStatusInProgress {
|
||||
builder = builder.Set("StartAt", model.GetMillis())
|
||||
}
|
||||
query, args, err := builder.ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
|
||||
sqlResult, err := jss.GetMasterX().Exec(query, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrapf(err, "failed to update Job with id=%s", id)
|
||||
}
|
||||
rows, err := sqlResult.RowsAffected()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "unable to get rows affected")
|
||||
}
|
||||
if rows != 1 {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) Get(id string) (*model.Job, error) {
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Jobs").
|
||||
Where(sq.Eq{"Id": id}).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
var status model.Job
|
||||
if err = jss.GetReplicaX().Get(&status, query, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Job", id)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get Job with id=%s", id)
|
||||
}
|
||||
return &status, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) GetAllPage(offset int, limit int) ([]*model.Job, error) {
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Jobs").
|
||||
OrderBy("CreateAt DESC").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset)).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
|
||||
statuses := []*model.Job{}
|
||||
if err = jss.GetReplicaX().Select(&statuses, query, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Jobs")
|
||||
}
|
||||
return statuses, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) GetAllByTypesPage(jobTypes []string, offset int, limit int) ([]*model.Job, error) {
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Jobs").
|
||||
Where(sq.Eq{"Type": jobTypes}).
|
||||
OrderBy("CreateAt DESC").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset)).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
|
||||
var jobs []*model.Job
|
||||
if err = jss.GetReplicaX().Select(&jobs, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Jobs with types")
|
||||
}
|
||||
return jobs, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) GetAllByType(jobType string) ([]*model.Job, error) {
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Jobs").
|
||||
Where(sq.Eq{"Type": jobType}).
|
||||
OrderBy("CreateAt DESC").ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
statuses := []*model.Job{}
|
||||
if err = jss.GetReplicaX().Select(&statuses, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Jobs with type=%s", jobType)
|
||||
}
|
||||
return statuses, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) GetAllByTypeAndStatus(jobType string, status string) ([]*model.Job, error) {
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Jobs").
|
||||
Where(sq.Eq{"Type": jobType, "Status": status}).
|
||||
OrderBy("CreateAt DESC").ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
jobs := []*model.Job{}
|
||||
if err = jss.GetReplicaX().Select(&jobs, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Jobs with type=%s", jobType)
|
||||
}
|
||||
return jobs, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) GetAllByTypePage(jobType string, offset int, limit int) ([]*model.Job, error) {
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Jobs").
|
||||
Where(sq.Eq{"Type": jobType}).
|
||||
OrderBy("CreateAt DESC").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset)).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
|
||||
statuses := []*model.Job{}
|
||||
if err = jss.GetReplicaX().Select(&statuses, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Jobs with type=%s", jobType)
|
||||
}
|
||||
return statuses, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) GetAllByStatus(status string) ([]*model.Job, error) {
|
||||
statuses := []*model.Job{}
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Jobs").
|
||||
Where(sq.Eq{"Status": status}).
|
||||
OrderBy("CreateAt ASC").ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
|
||||
if err = jss.GetReplicaX().Select(&statuses, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Jobs with status=%s", status)
|
||||
}
|
||||
return statuses, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) GetNewestJobByStatusAndType(status string, jobType string) (*model.Job, error) {
|
||||
return jss.GetNewestJobByStatusesAndType([]string{status}, jobType)
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) GetNewestJobByStatusesAndType(status []string, jobType string) (*model.Job, error) {
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Jobs").
|
||||
Where(sq.Eq{"Status": status, "Type": jobType}).
|
||||
OrderBy("CreateAt DESC").
|
||||
Limit(1).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
|
||||
var job model.Job
|
||||
if err = jss.GetReplicaX().Get(&job, query, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Job", fmt.Sprintf("<status, type>=<%s, %s>", strings.Join(status, ","), jobType))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to find Job with statuses=%s and type=%s", strings.Join(status, ","), jobType)
|
||||
}
|
||||
return &job, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) GetCountByStatusAndType(status string, jobType string) (int64, error) {
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Select("COUNT(*)").
|
||||
From("Jobs").
|
||||
Where(sq.Eq{"Status": status, "Type": jobType}).ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
|
||||
var count int64
|
||||
err = jss.GetReplicaX().Get(&count, query, args...)
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrapf(err, "failed to count Jobs with status=%s and type=%s", status, jobType)
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) Delete(id string) (string, error) {
|
||||
query, args, err := jss.getQueryBuilder().
|
||||
Delete("Jobs").
|
||||
Where(sq.Eq{"Id": id}).ToSql()
|
||||
if err != nil {
|
||||
return "", errors.Wrap(err, "job_tosql")
|
||||
}
|
||||
|
||||
if _, err = jss.GetMasterX().Exec(query, args...); err != nil {
|
||||
return "", errors.Wrapf(err, "failed to delete Job with id=%s", id)
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func (jss SqlJobStore) Cleanup(expiryTime int64, batchSize int) error {
|
||||
var query string
|
||||
if jss.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = "DELETE FROM Jobs WHERE Id IN (SELECT Id FROM Jobs WHERE CreateAt < ? AND (Status != ? AND Status != ?) ORDER BY CreateAt ASC LIMIT ?)"
|
||||
} else {
|
||||
query = "DELETE FROM Jobs WHERE CreateAt < ? AND (Status != ? AND Status != ?) ORDER BY CreateAt ASC LIMIT ?"
|
||||
}
|
||||
|
||||
var rowsAffected int64 = 1
|
||||
|
||||
for rowsAffected > 0 {
|
||||
sqlResult, err := jss.GetMasterX().Exec(query,
|
||||
expiryTime, model.JobStatusInProgress, model.JobStatusPending, batchSize)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "unable to delete jobs")
|
||||
}
|
||||
var rowErr error
|
||||
rowsAffected, rowErr = sqlResult.RowsAffected()
|
||||
if rowErr != nil {
|
||||
return errors.Wrap(err, "unable to delete jobs")
|
||||
}
|
||||
|
||||
time.Sleep(jobsCleanupDelay)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/job_store_test.go
Обычный файл
14
server/channels/store/sqlstore/job_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestJobStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestJobStore)
|
||||
}
|
||||
99
server/channels/store/sqlstore/license_store.go
Обычный файл
99
server/channels/store/sqlstore/license_store.go
Обычный файл
@@ -0,0 +1,99 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
// SqlLicenseStore encapsulates the database writes and reads for
|
||||
// model.LicenseRecord objects.
|
||||
type SqlLicenseStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlLicenseStore(sqlStore *SqlStore) store.LicenseStore {
|
||||
return &SqlLicenseStore{sqlStore}
|
||||
}
|
||||
|
||||
// Save validates and stores the license instance in the database. The Id
|
||||
// and Bytes fields are mandatory. The Bytes field is limited to a maximum
|
||||
// of 10000 bytes. If the license ID matches an existing license in the
|
||||
// database it returns the license stored in the database. If not, it saves the
|
||||
// new database and returns the created license with the CreateAt field
|
||||
// updated.
|
||||
func (ls SqlLicenseStore) Save(license *model.LicenseRecord) (*model.LicenseRecord, error) {
|
||||
license.PreSave()
|
||||
if err := license.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
query := ls.getQueryBuilder().
|
||||
Select("Id, CreateAt, Bytes").
|
||||
From("Licenses").
|
||||
Where(sq.Eq{"Id": license.Id})
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "license_tosql")
|
||||
}
|
||||
var storedLicense model.LicenseRecord
|
||||
if err := ls.GetReplicaX().Get(&storedLicense, queryString, args...); err != nil {
|
||||
// Only insert if not exists
|
||||
query, args, err := ls.getQueryBuilder().
|
||||
Insert("Licenses").
|
||||
Columns("Id", "CreateAt", "Bytes").
|
||||
Values(license.Id, license.CreateAt, license.Bytes).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "license_record_tosql")
|
||||
}
|
||||
if _, err := ls.GetMasterX().Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get License with licenseId=%s", license.Id)
|
||||
}
|
||||
return license, nil
|
||||
}
|
||||
return &storedLicense, nil
|
||||
}
|
||||
|
||||
// Get obtains the license with the provided id parameter from the database.
|
||||
// If the license doesn't exist it returns a model.AppError with
|
||||
// http.StatusNotFound in the StatusCode field.
|
||||
func (ls SqlLicenseStore) Get(id string) (*model.LicenseRecord, error) {
|
||||
query := ls.getQueryBuilder().
|
||||
Select("Id, CreateAt, Bytes").
|
||||
From("Licenses").
|
||||
Where(sq.Eq{"Id": id})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "license_record_tosql")
|
||||
}
|
||||
|
||||
license := &model.LicenseRecord{}
|
||||
if err := ls.GetReplicaX().Get(license, queryString, args...); err != nil {
|
||||
return nil, store.NewErrNotFound("License", id)
|
||||
}
|
||||
return license, nil
|
||||
}
|
||||
|
||||
func (ls SqlLicenseStore) GetAll() ([]*model.LicenseRecord, error) {
|
||||
query := ls.getQueryBuilder().
|
||||
Select("Id, CreateAt, Bytes").
|
||||
From("Licenses")
|
||||
|
||||
queryString, _, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "license_tosql")
|
||||
}
|
||||
|
||||
licenses := []*model.LicenseRecord{}
|
||||
if err := ls.GetReplicaX().Select(&licenses, queryString); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to fetch licenses")
|
||||
}
|
||||
|
||||
return licenses, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/license_store_test.go
Обычный файл
14
server/channels/store/sqlstore/license_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestLicenseStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestLicenseStore)
|
||||
}
|
||||
87
server/channels/store/sqlstore/link_metadata_store.go
Обычный файл
87
server/channels/store/sqlstore/link_metadata_store.go
Обычный файл
@@ -0,0 +1,87 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlLinkMetadataStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlLinkMetadataStore(sqlStore *SqlStore) store.LinkMetadataStore {
|
||||
return &SqlLinkMetadataStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlLinkMetadataStore) Save(metadata *model.LinkMetadata) (*model.LinkMetadata, error) {
|
||||
if err := metadata.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
metadata.PreSave()
|
||||
metadataBytes, err := json.Marshal(metadata.Data)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not serialize metadataBytes to JSON")
|
||||
}
|
||||
if s.IsBinaryParamEnabled() {
|
||||
metadataBytes = AppendBinaryFlag(metadataBytes)
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Insert("LinkMetadata").
|
||||
Columns("Hash", "URL", "Timestamp", "Type", "Data").
|
||||
Values(metadata.Hash, metadata.URL, metadata.Timestamp, metadata.Type, metadataBytes)
|
||||
|
||||
if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE URL = ?, Timestamp = ?, Type = ?, Data = ?", metadata.URL, metadata.Timestamp, metadata.Type, metadataBytes))
|
||||
} else {
|
||||
query = query.SuffixExpr(sq.Expr("ON CONFLICT (hash) DO UPDATE SET URL = ?, Timestamp = ?, Type = ?, Data = ?", metadata.URL, metadata.Timestamp, metadata.Type, metadataBytes))
|
||||
}
|
||||
|
||||
q, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "metadata_tosql")
|
||||
}
|
||||
|
||||
_, err = s.GetMasterX().Exec(q, args...)
|
||||
if err != nil && !IsUniqueConstraintError(err, []string{"PRIMARY", "linkmetadata_pkey"}) {
|
||||
return nil, errors.Wrap(err, "could not save link metadata")
|
||||
}
|
||||
|
||||
return metadata, nil
|
||||
}
|
||||
|
||||
func (s SqlLinkMetadataStore) Get(url string, timestamp int64) (*model.LinkMetadata, error) {
|
||||
var metadata model.LinkMetadata
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("LinkMetadata").
|
||||
Where(sq.Eq{"URL": url, "Timestamp": timestamp}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not create query with querybuilder")
|
||||
}
|
||||
err = s.GetReplicaX().Get(&metadata, query, args...)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("LinkMetadata", "url="+url)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "could not get metadata with selectone: url=%s", url)
|
||||
}
|
||||
|
||||
err = metadata.DeserializeDataToConcreteType()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "could not deserialize metadata to concrete type for url=%s", url)
|
||||
}
|
||||
|
||||
return &metadata, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/link_metadata_store_test.go
Обычный файл
14
server/channels/store/sqlstore/link_metadata_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestLinkMetadataStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestLinkMetadataStore)
|
||||
}
|
||||
23
server/channels/store/sqlstore/main_test.go
Обычный файл
23
server/channels/store/sqlstore/main_test.go
Обычный файл
@@ -0,0 +1,23 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/sqlstore"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/testlib"
|
||||
)
|
||||
|
||||
var mainHelper *testlib.MainHelper
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
mainHelper = testlib.NewMainHelperWithOptions(nil)
|
||||
defer mainHelper.Close()
|
||||
|
||||
sqlstore.InitTest()
|
||||
|
||||
mainHelper.Main(m)
|
||||
sqlstore.TearDownTest()
|
||||
}
|
||||
96
server/channels/store/sqlstore/notify_admin_store.go
Обычный файл
96
server/channels/store/sqlstore/notify_admin_store.go
Обычный файл
@@ -0,0 +1,96 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlNotifyAdminStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlNotifyAdminStore(sqlStore *SqlStore) store.NotifyAdminStore {
|
||||
return &SqlNotifyAdminStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlNotifyAdminStore) insert(data *model.NotifyAdminData) (sql.Result, error) {
|
||||
query := `INSERT INTO NotifyAdmin (UserId, CreateAt, RequiredPlan, RequiredFeature, Trial) VALUES (:UserId, :CreateAt, :RequiredPlan, :RequiredFeature, :Trial)`
|
||||
return s.GetMasterX().NamedExec(query, data)
|
||||
}
|
||||
|
||||
func (s SqlNotifyAdminStore) Save(data *model.NotifyAdminData) (*model.NotifyAdminData, error) {
|
||||
if err := data.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
data.PreSave()
|
||||
|
||||
_, err := s.insert(data)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save Notify Admin data")
|
||||
}
|
||||
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func (s SqlNotifyAdminStore) GetDataByUserIdAndFeature(userId string, feature model.MattermostFeature) ([]*model.NotifyAdminData, error) {
|
||||
data := []*model.NotifyAdminData{}
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("NotifyAdmin").
|
||||
Where(sq.Eq{"UserId": userId, "RequiredFeature": feature}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not build sql query to get all notification data by user id and required feature")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&data, query, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("NotifyAdmin", fmt.Sprintf("user id: %s and required feature: %s", userId, feature))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "notifcation data by user id: %s and required feature: %s", userId, feature)
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func (s SqlNotifyAdminStore) Get(trial bool) ([]*model.NotifyAdminData, error) {
|
||||
data := []*model.NotifyAdminData{}
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("NotifyAdmin").
|
||||
Where(sq.Eq{"Trial": trial}).
|
||||
Where("(SentAt IS NULL)").
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not build sql query to get all notifcation data")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&data, query, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "notifcation data")
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func (s SqlNotifyAdminStore) DeleteBefore(trial bool, now int64) error {
|
||||
if _, err := s.GetMasterX().Exec("DELETE FROM NotifyAdmin WHERE Trial = ? AND CreateAt < ? AND SentAt IS NULL", trial, now); err != nil {
|
||||
return errors.Wrapf(err, "failed to remove all notification data with trial=%t", trial)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlNotifyAdminStore) Update(userId string, requiredPlan string, requiredFeature model.MattermostFeature, now int64) error {
|
||||
if _, err := s.GetMasterX().Exec("UPDATE NotifyAdmin SET SentAt = ? WHERE UserId = ? AND RequiredPlan = ? AND RequiredFeature = ?", now, userId, requiredPlan, requiredFeature); err != nil {
|
||||
return errors.Wrapf(err, "failed to update SentAt for userId=%s and requiredPlan=%s", userId, requiredPlan)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/notify_admin_store_test.go
Обычный файл
14
server/channels/store/sqlstore/notify_admin_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestNotifyAdminStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestNotifyAdminStore)
|
||||
}
|
||||
322
server/channels/store/sqlstore/oauth_store.go
Обычный файл
322
server/channels/store/sqlstore/oauth_store.go
Обычный файл
@@ -0,0 +1,322 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlOAuthStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlOAuthStore(sqlStore *SqlStore) store.OAuthStore {
|
||||
return &SqlOAuthStore{sqlStore}
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) SaveApp(app *model.OAuthApp) (*model.OAuthApp, error) {
|
||||
if app.Id != "" {
|
||||
return nil, store.NewErrInvalidInput("OAuthApp", "Id", app.Id)
|
||||
}
|
||||
|
||||
app.PreSave()
|
||||
if err := app.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := as.GetMasterX().NamedExec(`INSERT INTO OAuthApps
|
||||
(Id, CreatorId, CreateAt, UpdateAt, ClientSecret, Name, Description, IconURL, CallbackUrls, Homepage, IsTrusted, MattermostAppID)
|
||||
VALUES
|
||||
(:Id, :CreatorId, :CreateAt, :UpdateAt, :ClientSecret, :Name, :Description, :IconURL, :CallbackUrls, :Homepage, :IsTrusted, :MattermostAppID)`, app); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save OAuthApp")
|
||||
}
|
||||
return app, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) UpdateApp(app *model.OAuthApp) (*model.OAuthApp, error) {
|
||||
app.PreUpdate()
|
||||
|
||||
if err := app.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var oldApp model.OAuthApp
|
||||
err := as.GetMasterX().Get(&oldApp, `SELECT * FROM OAuthApps
|
||||
WHERE id=?`, app.Id)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get OAuthApp with id=%s", app.Id)
|
||||
}
|
||||
if oldApp.Id == "" {
|
||||
return nil, store.NewErrInvalidInput("OAuthApp", "Id", app.Id)
|
||||
}
|
||||
|
||||
app.CreateAt = oldApp.CreateAt
|
||||
app.CreatorId = oldApp.CreatorId
|
||||
|
||||
res, err := as.GetMasterX().NamedExec(`UPDATE OAuthApps
|
||||
SET UpdateAt=:UpdateAt, ClientSecret=:ClientSecret, Name=:Name,
|
||||
Description=:Description, IconURL=:IconURL, CallbackUrls=:CallbackUrls,
|
||||
Homepage=:Homepage, IsTrusted=:IsTrusted, MattermostAppID=:MattermostAppID
|
||||
WHERE Id=:Id`, app)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update OAuthApp with id=%s", app.Id)
|
||||
}
|
||||
count, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "error while getting rows_affected")
|
||||
}
|
||||
if count > 1 {
|
||||
return nil, store.NewErrInvalidInput("OAuthApp", "Id", app.Id)
|
||||
}
|
||||
return app, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) GetApp(id string) (*model.OAuthApp, error) {
|
||||
var app model.OAuthApp
|
||||
if err := as.GetReplicaX().Get(&app, `SELECT * FROM OAuthApps WHERE Id=?`, id); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("OAuthApp", id)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get OAuthApp with id=%s", id)
|
||||
}
|
||||
if app.Id == "" {
|
||||
return nil, store.NewErrNotFound("OAuthApp", id)
|
||||
}
|
||||
return &app, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) GetAppByUser(userId string, offset, limit int) ([]*model.OAuthApp, error) {
|
||||
apps := []*model.OAuthApp{}
|
||||
|
||||
if err := as.GetReplicaX().Select(&apps, "SELECT * FROM OAuthApps WHERE CreatorId = ? LIMIT ? OFFSET ?", userId, limit, offset); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find OAuthApps with userId=%s", userId)
|
||||
}
|
||||
|
||||
return apps, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) GetApps(offset, limit int) ([]*model.OAuthApp, error) {
|
||||
apps := []*model.OAuthApp{}
|
||||
|
||||
if err := as.GetReplicaX().Select(&apps, "SELECT * FROM OAuthApps LIMIT ? OFFSET ?", limit, offset); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find OAuthApps")
|
||||
}
|
||||
|
||||
return apps, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) GetAuthorizedApps(userId string, offset, limit int) ([]*model.OAuthApp, error) {
|
||||
apps := []*model.OAuthApp{}
|
||||
|
||||
if err := as.GetReplicaX().Select(&apps,
|
||||
`SELECT o.* FROM OAuthApps AS o INNER JOIN
|
||||
Preferences AS p ON p.Name=o.Id AND p.UserId=? LIMIT ? OFFSET ?`, userId, limit, offset); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find OAuthApps with userId=%s", userId)
|
||||
}
|
||||
|
||||
return apps, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) DeleteApp(id string) (err error) {
|
||||
// wrap in a transaction so that if one fails, everything fails
|
||||
transaction, err := as.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
|
||||
if err := as.deleteApp(transaction, id); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := transaction.Commit(); err != nil {
|
||||
// don't need to rollback here since the transaction is already closed
|
||||
return errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) SaveAccessData(accessData *model.AccessData) (*model.AccessData, error) {
|
||||
if err := accessData.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := as.GetMasterX().NamedExec(`INSERT INTO OAuthAccessData
|
||||
(ClientId, UserId, Token, RefreshToken, RedirectUri, ExpiresAt, Scope)
|
||||
VALUES
|
||||
(:ClientId, :UserId, :Token, :RefreshToken, :RedirectUri, :ExpiresAt, :Scope)`, accessData); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save AccessData")
|
||||
}
|
||||
return accessData, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) GetAccessData(token string) (*model.AccessData, error) {
|
||||
accessData := model.AccessData{}
|
||||
|
||||
if err := as.GetReplicaX().Get(&accessData, "SELECT * FROM OAuthAccessData WHERE Token = ?", token); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get OAuthAccessData with token=%s", token)
|
||||
}
|
||||
return &accessData, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) GetAccessDataByUserForApp(userID, clientID string) ([]*model.AccessData, error) {
|
||||
accessData := []*model.AccessData{}
|
||||
|
||||
if err := as.GetReplicaX().Select(&accessData,
|
||||
"SELECT * FROM OAuthAccessData WHERE UserId = ? AND ClientId = ?", userID, clientID); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to delete OAuthAccessData with userId=%s and clientId=%s", userID, clientID)
|
||||
}
|
||||
return accessData, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) GetAccessDataByRefreshToken(token string) (*model.AccessData, error) {
|
||||
accessData := model.AccessData{}
|
||||
|
||||
if err := as.GetReplicaX().Get(&accessData, "SELECT * FROM OAuthAccessData WHERE RefreshToken = ?", token); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find OAuthAccessData with refreshToken=%s", token)
|
||||
}
|
||||
return &accessData, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) GetPreviousAccessData(userID, clientID string) (*model.AccessData, error) {
|
||||
accessData := model.AccessData{}
|
||||
|
||||
if err := as.GetReplicaX().Get(&accessData, "SELECT * FROM OAuthAccessData WHERE ClientId = ? AND UserId = ?", clientID, userID); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
return nil, errors.Wrapf(err, "failed to get AccessData with clientId=%s and userId=%s", clientID, userID)
|
||||
}
|
||||
return &accessData, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) UpdateAccessData(accessData *model.AccessData) (*model.AccessData, error) {
|
||||
if err := accessData.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := as.GetMasterX().NamedExec("UPDATE OAuthAccessData SET Token = :Token, ExpiresAt = :ExpiresAt, RefreshToken = :RefreshToken WHERE ClientId = :ClientId AND UserID = :UserId", accessData); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update OAuthAccessData with userId=%s and clientId=%s", accessData.UserId, accessData.ClientId)
|
||||
}
|
||||
return accessData, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) RemoveAccessData(token string) error {
|
||||
if _, err := as.GetMasterX().Exec("DELETE FROM OAuthAccessData WHERE Token = ?", token); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete OAuthAccessData with token=%s", token)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) RemoveAllAccessData() error {
|
||||
if _, err := as.GetMasterX().Exec("DELETE FROM OAuthAccessData"); err != nil {
|
||||
return errors.Wrap(err, "failed to delete OAuthAccessData")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) SaveAuthData(authData *model.AuthData) (*model.AuthData, error) {
|
||||
authData.PreSave()
|
||||
if err := authData.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := as.GetMasterX().NamedExec(`INSERT INTO OAuthAuthData
|
||||
(ClientId, UserId, Code, ExpiresIn, CreateAt, RedirectUri, State, Scope)
|
||||
VALUES
|
||||
(:ClientId, :UserId, :Code, :ExpiresIn, :CreateAt, :RedirectUri, :State, :Scope)`, authData); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save AuthData")
|
||||
}
|
||||
return authData, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) GetAuthData(code string) (*model.AuthData, error) {
|
||||
var authData model.AuthData
|
||||
err := as.GetReplicaX().Get(&authData, `SELECT * FROM OAuthAuthData WHERE Code=?`, code)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("AuthData", fmt.Sprintf("code=%s", code))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get AuthData with code=%s", code)
|
||||
}
|
||||
if authData.Code == "" {
|
||||
return nil, store.NewErrNotFound("AuthData", fmt.Sprintf("code=%s", code))
|
||||
}
|
||||
return &authData, nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) RemoveAuthData(code string) error {
|
||||
_, err := as.GetMasterX().Exec("DELETE FROM OAuthAuthData WHERE Code = ?", code)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete AuthData with code=%s", code)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) RemoveAuthDataByClientId(clientId string, userId string) error {
|
||||
_, err := as.GetMasterX().Exec("DELETE FROM OAuthAuthData WHERE ClientId = ? and UserId = ?", clientId, userId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete AuthData with clientId=%s and userId=%s", clientId, userId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) PermanentDeleteAuthDataByUser(userId string) error {
|
||||
_, err := as.GetMasterX().Exec("DELETE FROM OAuthAccessData WHERE UserId = ?", userId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete OAuthAccessData with userId=%s", userId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) deleteApp(transaction *sqlxTxWrapper, clientId string) error {
|
||||
if _, err := transaction.Exec("DELETE FROM OAuthApps WHERE Id = ?", clientId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete OAuthApp with id=%s", clientId)
|
||||
}
|
||||
|
||||
return as.deleteOAuthAppSessions(transaction, clientId)
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) deleteOAuthAppSessions(transaction *sqlxTxWrapper, clientId string) error {
|
||||
query := ""
|
||||
if as.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = "DELETE FROM Sessions s USING OAuthAccessData o WHERE o.Token = s.Token AND o.ClientId = ?"
|
||||
} else if as.DriverName() == model.DatabaseDriverMysql {
|
||||
query = "DELETE s.* FROM Sessions s INNER JOIN OAuthAccessData o ON o.Token = s.Token WHERE o.ClientId = ?"
|
||||
}
|
||||
|
||||
if _, err := transaction.Exec(query, clientId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Session with OAuthAccessData.Id=%s", clientId)
|
||||
}
|
||||
|
||||
return as.deleteOAuthTokens(transaction, clientId)
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) deleteOAuthTokens(transaction *sqlxTxWrapper, clientId string) error {
|
||||
if _, err := transaction.Exec("DELETE FROM OAuthAccessData WHERE ClientId = ?", clientId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete OAuthAccessData with id=%s", clientId)
|
||||
}
|
||||
|
||||
return as.deleteAppExtras(transaction, clientId)
|
||||
}
|
||||
|
||||
func (as SqlOAuthStore) deleteAppExtras(transaction *sqlxTxWrapper, clientId string) error {
|
||||
if _, err := transaction.Exec(
|
||||
`DELETE FROM
|
||||
Preferences
|
||||
WHERE
|
||||
Category = ?
|
||||
AND Name = ?`, model.PreferenceCategoryAuthorizedOAuthApp, clientId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Preferences with name=%s", clientId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/oauth_store_test.go
Обычный файл
14
server/channels/store/sqlstore/oauth_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestOAuthStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestOAuthStore)
|
||||
}
|
||||
358
server/channels/store/sqlstore/plugin_store.go
Обычный файл
358
server/channels/store/sqlstore/plugin_store.go
Обычный файл
@@ -0,0 +1,358 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultPluginKeyFetchLimit = 10
|
||||
)
|
||||
|
||||
type SqlPluginStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlPluginStore(sqlStore *SqlStore) store.PluginStore {
|
||||
return &SqlPluginStore{sqlStore}
|
||||
}
|
||||
|
||||
func (ps SqlPluginStore) SaveOrUpdate(kv *model.PluginKeyValue) (*model.PluginKeyValue, error) {
|
||||
if err := kv.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if kv.Value == nil {
|
||||
// Setting a key to nil is the same as removing it
|
||||
err := ps.Delete(kv.PluginId, kv.Key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return kv, nil
|
||||
}
|
||||
|
||||
query := ps.getQueryBuilder().
|
||||
Insert("PluginKeyValueStore").
|
||||
Columns("PluginId", "PKey", "PValue", "ExpireAt").
|
||||
Values(kv.PluginId, kv.Key, kv.Value, kv.ExpireAt)
|
||||
if ps.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = query.SuffixExpr(sq.Expr("ON CONFLICT (pluginid, pkey) DO UPDATE SET PValue = ?, ExpireAt = ?", kv.Value, kv.ExpireAt))
|
||||
} else if ps.DriverName() == model.DatabaseDriverMysql {
|
||||
query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE PValue = ?, ExpireAt = ?", kv.Value, kv.ExpireAt))
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
if _, err := ps.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to upsert PluginKeyValue")
|
||||
}
|
||||
|
||||
return kv, nil
|
||||
}
|
||||
|
||||
func (ps SqlPluginStore) CompareAndSet(kv *model.PluginKeyValue, oldValue []byte) (bool, error) {
|
||||
if err := kv.IsValid(); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
if kv.Value == nil {
|
||||
// Setting a key to nil is the same as removing it
|
||||
return ps.CompareAndDelete(kv, oldValue)
|
||||
}
|
||||
|
||||
if oldValue == nil {
|
||||
// Delete any existing, expired value.
|
||||
query := ps.getQueryBuilder().
|
||||
Delete("PluginKeyValueStore").
|
||||
Where(sq.Eq{"PluginId": kv.PluginId}).
|
||||
Where(sq.Eq{"PKey": kv.Key}).
|
||||
Where(sq.NotEq{"ExpireAt": int(0)}).
|
||||
Where(sq.Lt{"ExpireAt": model.GetMillis()})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
if _, err = ps.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return false, errors.Wrap(err, "failed to delete PluginKeyValue")
|
||||
}
|
||||
|
||||
// Insert if oldValue is nil
|
||||
queryString, args, err = ps.getQueryBuilder().
|
||||
Insert("PluginKeyValueStore").
|
||||
Columns("PluginId", "PKey", "PValue", "ExpireAt").
|
||||
Values(kv.PluginId, kv.Key, kv.Value, kv.ExpireAt).ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
if _, err := ps.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
// If the error is from unique constraints violation, it's the result of a
|
||||
// race condition, return false and no error. Otherwise we have a real error and
|
||||
// need to return it.
|
||||
if IsUniqueConstraintError(err, []string{"PRIMARY", "PluginId", "Key", "PKey", "pkey"}) {
|
||||
return false, nil
|
||||
}
|
||||
return false, errors.Wrap(err, "failed to insert PluginKeyValue")
|
||||
}
|
||||
} else {
|
||||
currentTime := model.GetMillis()
|
||||
|
||||
// Update if oldValue is not nil
|
||||
query := ps.getQueryBuilder().
|
||||
Update("PluginKeyValueStore").
|
||||
Set("PValue", kv.Value).
|
||||
Set("ExpireAt", kv.ExpireAt).
|
||||
Where(sq.Eq{"PluginId": kv.PluginId}).
|
||||
Where(sq.Eq{"PKey": kv.Key}).
|
||||
Where(sq.Eq{"PValue": oldValue}).
|
||||
Where(sq.Or{
|
||||
sq.Eq{"ExpireAt": int(0)},
|
||||
sq.Gt{"ExpireAt": currentTime},
|
||||
})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
updateResult, err := ps.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to update PluginKeyValue")
|
||||
}
|
||||
|
||||
if rowsAffected, err := updateResult.RowsAffected(); err != nil {
|
||||
// Failed to update
|
||||
return false, errors.Wrap(err, "unable to get rows affected")
|
||||
} else if rowsAffected == 0 {
|
||||
if ps.DriverName() == model.DatabaseDriverMysql && bytes.Equal(oldValue, kv.Value) {
|
||||
// ROW_COUNT on MySQL is zero even if the row existed but no changes to the row were required.
|
||||
// Check if the row exists with the required value to distinguish this case. Strictly speaking,
|
||||
// this isn't a good use of CompareAndSet anyway, since there's no corresponding guarantee of
|
||||
// atomicity. Nevertheless, let's return results consistent with Postgres and with what might
|
||||
// be expected in this case.
|
||||
query := ps.getQueryBuilder().
|
||||
Select("COUNT(*)").
|
||||
From("PluginKeyValueStore").
|
||||
Where(sq.Eq{"PluginId": kv.PluginId}).
|
||||
Where(sq.Eq{"PKey": kv.Key}).
|
||||
Where(sq.Eq{"PValue": kv.Value}).
|
||||
Where(sq.Or{
|
||||
sq.Eq{"ExpireAt": int(0)},
|
||||
sq.Gt{"ExpireAt": currentTime},
|
||||
})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
var count int64
|
||||
err = ps.GetReplicaX().Get(&count, queryString, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrapf(err, "failed to count PluginKeyValue with pluginId=%s and key=%s", kv.PluginId, kv.Key)
|
||||
}
|
||||
|
||||
if count == 0 {
|
||||
return false, nil
|
||||
} else if count == 1 {
|
||||
return true, nil
|
||||
} else {
|
||||
return false, errors.Wrapf(err, "got too many rows when counting PluginKeyValue with pluginId=%s, key=%s, rows=%d", kv.PluginId, kv.Key, count)
|
||||
}
|
||||
}
|
||||
|
||||
// No rows were affected by the update, where condition was not satisfied,
|
||||
// return false, but no error.
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (ps SqlPluginStore) CompareAndDelete(kv *model.PluginKeyValue, oldValue []byte) (bool, error) {
|
||||
if err := kv.IsValid(); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
if oldValue == nil {
|
||||
// nil can't be stored. Return showing that we didn't do anything
|
||||
return false, nil
|
||||
}
|
||||
|
||||
query := ps.getQueryBuilder().
|
||||
Delete("PluginKeyValueStore").
|
||||
Where(sq.Eq{"PluginId": kv.PluginId}).
|
||||
Where(sq.Eq{"PKey": kv.Key}).
|
||||
Where(sq.Eq{"PValue": oldValue}).
|
||||
Where(sq.Or{
|
||||
sq.Eq{"ExpireAt": int(0)},
|
||||
sq.Gt{"ExpireAt": model.GetMillis()},
|
||||
})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
deleteResult, err := ps.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to delete PluginKeyValue")
|
||||
}
|
||||
|
||||
if rowsAffected, err := deleteResult.RowsAffected(); err != nil {
|
||||
return false, errors.Wrap(err, "unable to get rows affected")
|
||||
} else if rowsAffected == 0 {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (ps SqlPluginStore) SetWithOptions(pluginId string, key string, value []byte, opt model.PluginKVSetOptions) (bool, error) {
|
||||
if err := opt.IsValid(); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
kv, err := model.NewPluginKeyValueFromOptions(pluginId, key, value, opt)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
if opt.Atomic {
|
||||
return ps.CompareAndSet(kv, opt.OldValue)
|
||||
}
|
||||
|
||||
savedKv, nErr := ps.SaveOrUpdate(kv)
|
||||
if nErr != nil {
|
||||
return false, nErr
|
||||
}
|
||||
|
||||
return savedKv != nil, nil
|
||||
}
|
||||
|
||||
func (ps SqlPluginStore) Get(pluginId, key string) (*model.PluginKeyValue, error) {
|
||||
currentTime := model.GetMillis()
|
||||
query := ps.getQueryBuilder().Select("PluginId, PKey, PValue, ExpireAt").
|
||||
From("PluginKeyValueStore").
|
||||
Where(sq.Eq{"PluginId": pluginId}).
|
||||
Where(sq.Eq{"PKey": key}).
|
||||
Where(sq.Or{sq.Eq{"ExpireAt": 0}, sq.Gt{"ExpireAt": currentTime}})
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
row := ps.GetReplicaX().QueryRowx(queryString, args...)
|
||||
var kv model.PluginKeyValue
|
||||
if err := row.Scan(&kv.PluginId, &kv.Key, &kv.Value, &kv.ExpireAt); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("PluginKeyValue", fmt.Sprintf("pluginId=%s, key=%s", pluginId, key))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get PluginKeyValue with pluginId=%s and key=%s", pluginId, key)
|
||||
}
|
||||
|
||||
return &kv, nil
|
||||
}
|
||||
|
||||
func (ps SqlPluginStore) Delete(pluginId, key string) error {
|
||||
query := ps.getQueryBuilder().
|
||||
Delete("PluginKeyValueStore").
|
||||
Where(sq.Eq{"PluginId": pluginId}).
|
||||
Where(sq.Eq{"Pkey": key})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
if _, err := ps.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete PluginKeyValue with pluginId=%s and key=%s", pluginId, key)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ps SqlPluginStore) DeleteAllForPlugin(pluginId string) error {
|
||||
query := ps.getQueryBuilder().
|
||||
Delete("PluginKeyValueStore").
|
||||
Where(sq.Eq{"PluginId": pluginId})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
if _, err := ps.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrapf(err, "failed to get all PluginKeyValues with pluginId=%s ", pluginId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ps SqlPluginStore) DeleteAllExpired() error {
|
||||
currentTime := model.GetMillis()
|
||||
query := ps.getQueryBuilder().
|
||||
Delete("PluginKeyValueStore").
|
||||
Where(sq.NotEq{"ExpireAt": 0}).
|
||||
Where(sq.Lt{"ExpireAt": currentTime})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
if _, err := ps.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to delete all expired PluginKeyValues")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ps SqlPluginStore) List(pluginId string, offset int, limit int) ([]string, error) {
|
||||
if limit <= 0 {
|
||||
limit = defaultPluginKeyFetchLimit
|
||||
}
|
||||
|
||||
if offset <= 0 {
|
||||
offset = 0
|
||||
}
|
||||
|
||||
query := ps.getQueryBuilder().
|
||||
Select("Pkey").
|
||||
From("PluginKeyValueStore").
|
||||
Where(sq.Eq{"PluginId": pluginId}).
|
||||
Where(sq.Or{
|
||||
sq.Eq{"ExpireAt": int(0)},
|
||||
sq.Gt{"ExpireAt": model.GetMillis()},
|
||||
}).
|
||||
OrderBy("PKey").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset))
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "plugin_tosql")
|
||||
}
|
||||
|
||||
keys := []string{}
|
||||
err = ps.GetReplicaX().Select(&keys, queryString, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get PluginKeyValues with pluginId=%s", pluginId)
|
||||
}
|
||||
|
||||
return keys, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/plugin_store_test.go
Обычный файл
14
server/channels/store/sqlstore/plugin_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestPluginStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestPluginStore)
|
||||
}
|
||||
192
server/channels/store/sqlstore/post_acknowledgements_store.go
Обычный файл
192
server/channels/store/sqlstore/post_acknowledgements_store.go
Обычный файл
@@ -0,0 +1,192 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlPostAcknowledgementStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlPostAcknowledgementStore(sqlStore *SqlStore) store.PostAcknowledgementStore {
|
||||
return &SqlPostAcknowledgementStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s *SqlPostAcknowledgementStore) Get(postID, userID string) (*model.PostAcknowledgement, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("PostId", "UserId", "AcknowledgedAt").
|
||||
From("PostAcknowledgements").
|
||||
Where(sq.And{
|
||||
sq.Eq{"PostId": postID},
|
||||
sq.Eq{"UserId": userID},
|
||||
sq.NotEq{"AcknowledgedAt": 0},
|
||||
})
|
||||
|
||||
var acknowledgement model.PostAcknowledgement
|
||||
err := s.GetReplicaX().GetBuilder(&acknowledgement, query)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("PostAcknowledgement", postID)
|
||||
}
|
||||
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &acknowledgement, nil
|
||||
}
|
||||
|
||||
func (s *SqlPostAcknowledgementStore) Save(postID, userID string, acknowledgedAt int64) (*model.PostAcknowledgement, error) {
|
||||
if acknowledgedAt == 0 {
|
||||
acknowledgedAt = model.GetMillis()
|
||||
}
|
||||
|
||||
acknowledgement := &model.PostAcknowledgement{
|
||||
UserId: userID,
|
||||
PostId: postID,
|
||||
AcknowledgedAt: acknowledgedAt,
|
||||
}
|
||||
|
||||
if err := acknowledgement.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Insert("PostAcknowledgements").
|
||||
Columns("PostId", "UserId", "AcknowledgedAt").
|
||||
Values(acknowledgement.PostId, acknowledgement.UserId, acknowledgement.AcknowledgedAt)
|
||||
|
||||
if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE AcknowledgedAt = ?", acknowledgement.AcknowledgedAt))
|
||||
} else {
|
||||
query = query.SuffixExpr(sq.Expr("ON CONFLICT (postid, userid) DO UPDATE SET AcknowledgedAt = ?", acknowledgement.AcknowledgedAt))
|
||||
}
|
||||
|
||||
_, err = transaction.ExecBuilder(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
err = updatePost(transaction, acknowledgement.PostId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
err = transaction.Commit()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
|
||||
return acknowledgement, nil
|
||||
}
|
||||
|
||||
func (s *SqlPostAcknowledgementStore) Delete(acknowledgement *model.PostAcknowledgement) error {
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Update("PostAcknowledgements").
|
||||
Set("AcknowledgedAt", 0).
|
||||
Where(sq.And{
|
||||
sq.Eq{"PostId": acknowledgement.PostId},
|
||||
sq.Eq{"UserId": acknowledgement.UserId},
|
||||
})
|
||||
|
||||
_, err = transaction.ExecBuilder(query)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = updatePost(transaction, acknowledgement.PostId)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = transaction.Commit()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SqlPostAcknowledgementStore) GetForPost(postID string) ([]*model.PostAcknowledgement, error) {
|
||||
var acknowledgements []*model.PostAcknowledgement
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("PostId", "UserId", "AcknowledgedAt").
|
||||
From("PostAcknowledgements").
|
||||
Where(sq.And{
|
||||
sq.NotEq{"AcknowledgedAt": 0},
|
||||
sq.Eq{"PostId": postID},
|
||||
})
|
||||
|
||||
err := s.GetReplicaX().SelectBuilder(&acknowledgements, query)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get PostAcknowledgements for postID=%s", postID)
|
||||
}
|
||||
|
||||
return acknowledgements, nil
|
||||
}
|
||||
|
||||
func (s *SqlPostAcknowledgementStore) GetForPosts(postIds []string) ([]*model.PostAcknowledgement, error) {
|
||||
var acknowledgements []*model.PostAcknowledgement
|
||||
|
||||
perPage := 200
|
||||
for i := 0; i < len(postIds); i += perPage {
|
||||
j := i + perPage
|
||||
if len(postIds) < j {
|
||||
j = len(postIds)
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("PostId", "UserId", "AcknowledgedAt").
|
||||
From("PostAcknowledgements").
|
||||
Where(sq.And{
|
||||
sq.Eq{"PostId": postIds[i:j]},
|
||||
sq.NotEq{"AcknowledgedAt": 0},
|
||||
})
|
||||
|
||||
var acknowledgementsBatch []*model.PostAcknowledgement
|
||||
err := s.GetReplicaX().SelectBuilder(&acknowledgementsBatch, query)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get PostAcknowledgements for post list")
|
||||
}
|
||||
|
||||
acknowledgements = append(acknowledgements, acknowledgementsBatch...)
|
||||
}
|
||||
|
||||
return acknowledgements, nil
|
||||
}
|
||||
|
||||
func updatePost(transaction *sqlxTxWrapper, postId string) error {
|
||||
_, err := transaction.Exec(
|
||||
`UPDATE
|
||||
Posts
|
||||
SET
|
||||
UpdateAt = ?
|
||||
WHERE
|
||||
Id = ?`,
|
||||
model.GetMillis(),
|
||||
postId,
|
||||
)
|
||||
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestPostAcknowledgementsStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestPostAcknowledgementsStore)
|
||||
}
|
||||
64
server/channels/store/sqlstore/post_priority_store.go
Обычный файл
64
server/channels/store/sqlstore/post_priority_store.go
Обычный файл
@@ -0,0 +1,64 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
sq "github.com/mattermost/squirrel"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlPostPriorityStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlPostPriorityStore(sqlStore *SqlStore) store.PostPriorityStore {
|
||||
return &SqlPostPriorityStore{
|
||||
SqlStore: sqlStore,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *SqlPostPriorityStore) GetForPost(postId string) (*model.PostPriority, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("Priority", "RequestedAck", "PersistentNotifications").
|
||||
From("PostsPriority").
|
||||
Where(sq.Eq{"PostId": postId})
|
||||
|
||||
var postPriority model.PostPriority
|
||||
err := s.GetReplicaX().GetBuilder(&postPriority, query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &postPriority, nil
|
||||
}
|
||||
|
||||
func (s *SqlPostPriorityStore) GetForPosts(postIds []string) ([]*model.PostPriority, error) {
|
||||
var priority []*model.PostPriority
|
||||
|
||||
perPage := 200
|
||||
for i := 0; i < len(postIds); i += perPage {
|
||||
j := i + perPage
|
||||
if len(postIds) < j {
|
||||
j = len(postIds)
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("PostId", "Priority", "RequestedAck", "PersistentNotifications").
|
||||
From("PostsPriority").
|
||||
Where(sq.Eq{"PostId": postIds[i:j]})
|
||||
|
||||
var priorityBatch []*model.PostPriority
|
||||
err := s.GetReplicaX().SelectBuilder(&priority, query)
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
priority = append(priority, priorityBatch...)
|
||||
}
|
||||
|
||||
return priority, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/post_priority_store_test.go
Обычный файл
14
server/channels/store/sqlstore/post_priority_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestPostPriorityStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestPostPriorityStore)
|
||||
}
|
||||
3379
server/channels/store/sqlstore/post_store.go
Обычный файл
3379
server/channels/store/sqlstore/post_store.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
65
server/channels/store/sqlstore/post_store_test.go
Обычный файл
65
server/channels/store/sqlstore/post_store_test.go
Обычный файл
@@ -0,0 +1,65 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/searchtest"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestPostStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestPostStore)
|
||||
}
|
||||
|
||||
func TestSearchPostStore(t *testing.T) {
|
||||
StoreTestWithSearchTestEngine(t, searchtest.TestSearchPostStore)
|
||||
}
|
||||
|
||||
func TestMysqlStopWords(t *testing.T) {
|
||||
mysqlStopWordsTests := []struct {
|
||||
Name string
|
||||
Args []string
|
||||
Expected []string
|
||||
Empty bool
|
||||
}{
|
||||
{
|
||||
Name: "Should remove only the stop words",
|
||||
Args: []string{"where is my car", "so this is real", "test this-and-that is awesome"},
|
||||
Expected: []string{"my car", "so real", "test this-and-that awesome"},
|
||||
},
|
||||
{
|
||||
Name: "Should not remove part of a word containing stop words",
|
||||
Args: []string{"whereabouts", "wherein", "tothis", "thisorthat", "waswhen", "whowas", "inthe", "whowill", "thewww"},
|
||||
Expected: []string{"whereabouts", "wherein", "tothis", "thisorthat", "waswhen", "whowas", "inthe", "whowill", "thewww"},
|
||||
},
|
||||
{
|
||||
Name: "Should remove all words from terms",
|
||||
Args: []string{"where about", "where in", "to this", "this or that", "was when", "who was", "in the", "who will", "the www"},
|
||||
Empty: true,
|
||||
},
|
||||
{
|
||||
Name: "Should not remove part of a word containing stop words separated by hyphens",
|
||||
Args: []string{"where-about", "where-in", "to-this", "this-or-that", "was-when", "who-was", "in-the", "who-will", "the-www"},
|
||||
Expected: []string{"where-about", "where-in", "to-this", "this-or-that", "was-when", "who-was", "in-the", "who-will", "the-www"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range mysqlStopWordsTests {
|
||||
t.Run(tc.Name, func(t *testing.T) {
|
||||
for i, term := range tc.Args {
|
||||
got, err := removeMysqlStopWordsFromTerms(term)
|
||||
require.NoError(t, err)
|
||||
if tc.Empty {
|
||||
require.Empty(t, got)
|
||||
} else {
|
||||
require.Equal(t, tc.Expected[i], got)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
321
server/channels/store/sqlstore/preference_store.go
Обычный файл
321
server/channels/store/sqlstore/preference_store.go
Обычный файл
@@ -0,0 +1,321 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
)
|
||||
|
||||
type SqlPreferenceStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlPreferenceStore(sqlStore *SqlStore) store.PreferenceStore {
|
||||
s := &SqlPreferenceStore{sqlStore}
|
||||
return s
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) deleteUnusedFeatures() {
|
||||
mlog.Debug("Deleting any unused pre-release features")
|
||||
sql, args, err := s.getQueryBuilder().
|
||||
Delete("Preferences").
|
||||
Where(sq.Eq{"Category": model.PreferenceCategoryAdvancedSettings}).
|
||||
Where(sq.Eq{"Value": "false"}).
|
||||
Where(sq.Like{"Name": store.FeatureTogglePrefix + "%"}).ToSql()
|
||||
if err != nil {
|
||||
mlog.Warn("Could not build sql query to delete unused features", mlog.Err(err))
|
||||
}
|
||||
if _, err = s.GetMasterX().Exec(sql, args...); err != nil {
|
||||
mlog.Warn("Failed to delete unused features", mlog.Err(err))
|
||||
}
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) Save(preferences model.Preferences) (err error) {
|
||||
// wrap in a transaction so that if one fails, everything fails
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
for _, preference := range preferences {
|
||||
preference := preference
|
||||
if upsertErr := s.saveTx(transaction, &preference); upsertErr != nil {
|
||||
return upsertErr
|
||||
}
|
||||
}
|
||||
|
||||
if err := transaction.Commit(); err != nil {
|
||||
// don't need to rollback here since the transaction is already closed
|
||||
return errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) save(transaction *sqlxTxWrapper, preference *model.Preference) error {
|
||||
preference.PreUpdate()
|
||||
|
||||
if err := preference.IsValid(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Insert("Preferences").
|
||||
Columns("UserId", "Category", "Name", "Value").
|
||||
Values(preference.UserId, preference.Category, preference.Name, preference.Value)
|
||||
|
||||
if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE Value = ?", preference.Value))
|
||||
} else if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = query.SuffixExpr(sq.Expr("ON CONFLICT (userid, category, name) DO UPDATE SET Value = ?", preference.Value))
|
||||
} else {
|
||||
return store.NewErrNotImplemented("failed to update preference because of missing driver")
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to generate sqlquery")
|
||||
}
|
||||
|
||||
if _, err = transaction.Exec(queryString, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to save Preference")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) saveTx(transaction *sqlxTxWrapper, preference *model.Preference) error {
|
||||
preference.PreUpdate()
|
||||
|
||||
if err := preference.IsValid(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Insert("Preferences").
|
||||
Columns("UserId", "Category", "Name", "Value").
|
||||
Values(preference.UserId, preference.Category, preference.Name, preference.Value)
|
||||
|
||||
if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE Value = ?", preference.Value))
|
||||
} else if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = query.SuffixExpr(sq.Expr("ON CONFLICT (userid, category, name) DO UPDATE SET Value = ?", preference.Value))
|
||||
} else {
|
||||
return store.NewErrNotImplemented("failed to update preference because of missing driver")
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to generate sqlquery")
|
||||
}
|
||||
|
||||
if _, err = transaction.Exec(queryString, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to save Preference")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) Get(userId string, category string, name string) (*model.Preference, error) {
|
||||
var preference model.Preference
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Preferences").
|
||||
Where(sq.Eq{"UserId": userId}).
|
||||
Where(sq.Eq{"Category": category}).
|
||||
Where(sq.Eq{"Name": name}).
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not build sql query to get preference")
|
||||
}
|
||||
if err = s.GetReplicaX().Get(&preference, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Preference with userId=%s, category=%s, name=%s", userId, category, name)
|
||||
}
|
||||
|
||||
return &preference, nil
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) GetCategoryAndName(category string, name string) (model.Preferences, error) {
|
||||
var preferences model.Preferences
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Preferences").
|
||||
Where(sq.Eq{"Category": category}).
|
||||
Where(sq.Eq{"Name": name}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not build sql query to get preference")
|
||||
}
|
||||
if err = s.GetReplicaX().Select(&preferences, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Preference with category=%s, name=%s", category, name)
|
||||
}
|
||||
return preferences, nil
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) GetCategory(userId string, category string) (model.Preferences, error) {
|
||||
var preferences model.Preferences
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Preferences").
|
||||
Where(sq.Eq{"UserId": userId}).
|
||||
Where(sq.Eq{"Category": category}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not build sql query to get preference")
|
||||
}
|
||||
if err = s.GetReplicaX().Select(&preferences, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Preference with userId=%s, category=%s", userId, category)
|
||||
}
|
||||
return preferences, nil
|
||||
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) GetAll(userId string) (model.Preferences, error) {
|
||||
var preferences model.Preferences
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Preferences").
|
||||
Where(sq.Eq{"UserId": userId}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not build sql query to get preference")
|
||||
}
|
||||
if err = s.GetReplicaX().Select(&preferences, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Preference with userId=%s", userId)
|
||||
}
|
||||
return preferences, nil
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) PermanentDeleteByUser(userId string) error {
|
||||
sql, args, err := s.getQueryBuilder().
|
||||
Delete("Preferences").
|
||||
Where(sq.Eq{"UserId": userId}).ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "could not build sql query to get delete preference by user")
|
||||
}
|
||||
if _, err := s.GetMasterX().Exec(sql, args...); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Preference with userId=%s", userId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) Delete(userId, category, name string) error {
|
||||
|
||||
sql, args, err := s.getQueryBuilder().
|
||||
Delete("Preferences").
|
||||
Where(sq.Eq{"UserId": userId}).
|
||||
Where(sq.Eq{"Category": category}).
|
||||
Where(sq.Eq{"Name": name}).ToSql()
|
||||
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "could not build sql query to get delete preference")
|
||||
}
|
||||
|
||||
if _, err = s.GetMasterX().Exec(sql, args...); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Preference with userId=%s, category=%s and name=%s", userId, category, name)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) DeleteCategory(userId string, category string) error {
|
||||
|
||||
sql, args, err := s.getQueryBuilder().
|
||||
Delete("Preferences").
|
||||
Where(sq.Eq{"UserId": userId}).
|
||||
Where(sq.Eq{"Category": category}).ToSql()
|
||||
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "could not build sql query to get delete preference by category")
|
||||
}
|
||||
|
||||
if _, err = s.GetMasterX().Exec(sql, args...); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Preference with userId=%s and category=%s", userId, category)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) DeleteCategoryAndName(category string, name string) error {
|
||||
sql, args, err := s.getQueryBuilder().
|
||||
Delete("Preferences").
|
||||
Where(sq.Eq{"Name": name}).
|
||||
Where(sq.Eq{"Category": category}).ToSql()
|
||||
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "could not build sql query to get delete preference by category and name")
|
||||
}
|
||||
|
||||
if _, err = s.GetMasterX().Exec(sql, args...); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Preference with category=%s and name=%s", category, name)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteOrphanedRows removes entries from Preferences (flagged post) when a
|
||||
// corresponding post no longer exists.
|
||||
func (s *SqlPreferenceStore) DeleteOrphanedRows(limit int) (deleted int64, err error) {
|
||||
// We need the extra level of nesting to deal with MySQL's locking
|
||||
const query = `
|
||||
DELETE FROM Preferences WHERE Name IN (
|
||||
SELECT * FROM (
|
||||
SELECT Preferences.Name FROM Preferences
|
||||
LEFT JOIN Posts ON Preferences.Name = Posts.Id
|
||||
WHERE Posts.Id IS NULL AND Category = ?
|
||||
LIMIT ?
|
||||
) AS A
|
||||
)`
|
||||
|
||||
result, err := s.GetMasterX().Exec(query, model.PreferenceCategoryFlaggedPost, limit)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
deleted, err = result.RowsAffected()
|
||||
return
|
||||
}
|
||||
|
||||
func (s SqlPreferenceStore) CleanupFlagsBatch(limit int64) (int64, error) {
|
||||
if limit < 0 {
|
||||
// uint64 does not throw an error, it overflows if it is negative.
|
||||
// it is better to manually check here, or change the function type to uint64
|
||||
return int64(0), errors.Errorf("Received a negative limit")
|
||||
}
|
||||
nameInQ, nameInArgs, err := sq.Select("*").
|
||||
FromSelect(
|
||||
sq.Select("Preferences.Name").
|
||||
From("Preferences").
|
||||
LeftJoin("Posts ON Preferences.Name = Posts.Id").
|
||||
Where(sq.Eq{"Preferences.Category": model.PreferenceCategoryFlaggedPost}).
|
||||
Where(sq.Eq{"Posts.Id": nil}).
|
||||
Limit(uint64(limit)),
|
||||
"t").
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrap(err, "could not build nested sql query to delete preference")
|
||||
}
|
||||
query, args, err := s.getQueryBuilder().Delete("Preferences").
|
||||
Where(sq.Eq{"Category": model.PreferenceCategoryFlaggedPost}).
|
||||
Where(sq.Expr("name IN ("+nameInQ+")", nameInArgs...)).
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrap(err, "could not build sql query to delete preference")
|
||||
}
|
||||
|
||||
sqlResult, err := s.GetMasterX().Exec(query, args...)
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrap(err, "failed to delete Preference")
|
||||
}
|
||||
rowsAffected, err := sqlResult.RowsAffected()
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrap(err, "unable to get rows affected")
|
||||
}
|
||||
|
||||
return rowsAffected, nil
|
||||
}
|
||||
83
server/channels/store/sqlstore/preference_store_test.go
Обычный файл
83
server/channels/store/sqlstore/preference_store_test.go
Обычный файл
@@ -0,0 +1,83 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestPreferenceStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestPreferenceStore)
|
||||
}
|
||||
|
||||
func TestDeleteUnusedFeatures(t *testing.T) {
|
||||
StoreTest(t, func(t *testing.T, ss store.Store) {
|
||||
userId1 := model.NewId()
|
||||
userId2 := model.NewId()
|
||||
category := model.PreferenceCategoryAdvancedSettings
|
||||
feature1 := "feature1"
|
||||
feature2 := "feature2"
|
||||
|
||||
features := model.Preferences{
|
||||
{
|
||||
UserId: userId1,
|
||||
Category: category,
|
||||
Name: store.FeatureTogglePrefix + feature1,
|
||||
Value: "true",
|
||||
},
|
||||
{
|
||||
UserId: userId2,
|
||||
Category: category,
|
||||
Name: store.FeatureTogglePrefix + feature1,
|
||||
Value: "false",
|
||||
},
|
||||
{
|
||||
UserId: userId1,
|
||||
Category: category,
|
||||
Name: store.FeatureTogglePrefix + feature2,
|
||||
Value: "false",
|
||||
},
|
||||
{
|
||||
UserId: userId2,
|
||||
Category: category,
|
||||
Name: store.FeatureTogglePrefix + feature2,
|
||||
Value: "true",
|
||||
},
|
||||
}
|
||||
|
||||
err := ss.Preference().Save(features)
|
||||
require.NoError(t, err)
|
||||
|
||||
ss.Preference().(*SqlPreferenceStore).deleteUnusedFeatures()
|
||||
|
||||
//make sure features with value "false" have actually been deleted from the database
|
||||
var val int64
|
||||
if err := ss.Preference().(*SqlPreferenceStore).GetReplicaX().Get(&val, `SELECT COUNT(*)
|
||||
FROM Preferences
|
||||
WHERE Category = ?
|
||||
AND Value = ?
|
||||
AND Name LIKE '`+store.FeatureTogglePrefix+`%'`, model.PreferenceCategoryAdvancedSettings, "false"); err != nil {
|
||||
require.NoError(t, err)
|
||||
} else if val != 0 {
|
||||
require.Fail(t, "Found %d features with value 'false', expected all to be deleted", val)
|
||||
}
|
||||
//
|
||||
// make sure features with value "true" remain saved
|
||||
if err := ss.Preference().(*SqlPreferenceStore).GetReplicaX().Get(&val, `SELECT COUNT(*)
|
||||
FROM Preferences
|
||||
WHERE Category = ?
|
||||
AND Value = ?
|
||||
AND Name LIKE '`+store.FeatureTogglePrefix+`%'`, model.PreferenceCategoryAdvancedSettings, "true"); err != nil {
|
||||
require.NoError(t, err)
|
||||
} else if val == 0 {
|
||||
require.Fail(t, "Found %d features with value 'true', expected to find at least %d features", val, 2)
|
||||
}
|
||||
})
|
||||
}
|
||||
126
server/channels/store/sqlstore/product_notices_store.go
Обычный файл
126
server/channels/store/sqlstore/product_notices_store.go
Обычный файл
@@ -0,0 +1,126 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlProductNoticesStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlProductNoticesStore(sqlStore *SqlStore) store.ProductNoticesStore {
|
||||
return &SqlProductNoticesStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlProductNoticesStore) Clear(notices []string) error {
|
||||
sql, args, err := s.getQueryBuilder().Delete("ProductNoticeViewState").Where(sq.Eq{"NoticeId": notices}).ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "product_notice_view_state_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(sql, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to delete records from ProductNoticeViewState")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlProductNoticesStore) ClearOldNotices(currentNotices model.ProductNotices) error {
|
||||
var notices []string
|
||||
for _, currentNotice := range currentNotices {
|
||||
notices = append(notices, currentNotice.ID)
|
||||
}
|
||||
sql, args, err := s.getQueryBuilder().Delete("ProductNoticeViewState").Where(sq.NotEq{"NoticeId": notices}).ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "product_notice_view_state_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(sql, args...); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete records from ProductNoticeViewState")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlProductNoticesStore) View(userId string, notices []string) (err error) {
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
|
||||
noticeStates := []model.ProductNoticeViewState{}
|
||||
sql, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("ProductNoticeViewState").
|
||||
Where(sq.And{sq.Eq{"UserId": userId}, sq.Eq{"NoticeId": notices}}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "View_ToSql")
|
||||
}
|
||||
if err := transaction.Select(¬iceStates, sql, args...); err != nil {
|
||||
return errors.Wrapf(err, "failed to get ProductNoticeViewState with userId=%s", userId)
|
||||
}
|
||||
|
||||
now := time.Now().UTC().Unix()
|
||||
|
||||
// update existing records
|
||||
for i := range noticeStates {
|
||||
noticeStates[i].Viewed += 1
|
||||
noticeStates[i].Timestamp = now
|
||||
if _, err := transaction.NamedExec(`UPDATE ProductNoticeViewState
|
||||
SET Viewed=:Viewed, Timestamp=:Timestamp WHERE UserId=:UserId AND NoticeId=:NoticeId`, ¬iceStates[i]); err != nil {
|
||||
return errors.Wrapf(err, "failed to update ProductNoticeViewState")
|
||||
}
|
||||
}
|
||||
|
||||
// add new ones
|
||||
haveNoticeState := func(n string) bool {
|
||||
for _, ns := range noticeStates {
|
||||
if ns.NoticeId == n {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
for _, noticeId := range notices {
|
||||
if !haveNoticeState(noticeId) {
|
||||
productNoticeViewState := &model.ProductNoticeViewState{
|
||||
UserId: userId,
|
||||
NoticeId: noticeId,
|
||||
Viewed: 1,
|
||||
Timestamp: now,
|
||||
}
|
||||
if _, err := transaction.NamedExec(`INSERT INTO ProductNoticeViewState (UserId, NoticeId, Viewed, Timestamp)
|
||||
VALUES (:UserId, :NoticeId, :Viewed, :Timestamp)`, productNoticeViewState); err != nil {
|
||||
return errors.Wrapf(err, "failed to insert ProductNoticeViewState")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if err := transaction.Commit(); err != nil {
|
||||
return errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlProductNoticesStore) GetViews(userId string) ([]model.ProductNoticeViewState, error) {
|
||||
noticeStates := []model.ProductNoticeViewState{}
|
||||
sql, args, err := s.getQueryBuilder().Select("*").From("ProductNoticeViewState").Where(sq.Eq{"UserId": userId}).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "product_notice_view_state_tosql")
|
||||
}
|
||||
if err := s.GetReplicaX().Select(¬iceStates, sql, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get ProductNoticeViewState with userId=%s", userId)
|
||||
}
|
||||
return noticeStates, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/product_notices_store_test.go
Обычный файл
14
server/channels/store/sqlstore/product_notices_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestProductNoticesStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestProductNoticesStore)
|
||||
}
|
||||
436
server/channels/store/sqlstore/reaction_store.go
Обычный файл
436
server/channels/store/sqlstore/reaction_store.go
Обычный файл
@@ -0,0 +1,436 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
sq "github.com/mattermost/squirrel"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
type SqlReactionStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlReactionStore(sqlStore *SqlStore) store.ReactionStore {
|
||||
return &SqlReactionStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s *SqlReactionStore) Save(reaction *model.Reaction) (re *model.Reaction, err error) {
|
||||
reaction.PreSave()
|
||||
if err := reaction.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
if reaction.ChannelId == "" {
|
||||
// get channelId, if not already populated
|
||||
var channelIds []string
|
||||
var args []interface{}
|
||||
query := "SELECT ChannelId from Posts where Id = ?"
|
||||
args = append(args, reaction.PostId)
|
||||
err = transaction.Select(&channelIds, query, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed while getting channelId from Posts")
|
||||
}
|
||||
reaction.ChannelId = channelIds[0]
|
||||
}
|
||||
err = s.saveReactionAndUpdatePost(transaction, reaction)
|
||||
if err != nil {
|
||||
// We don't consider duplicated save calls as an error
|
||||
if !IsUniqueConstraintError(err, []string{"reactions_pkey", "PRIMARY"}) {
|
||||
return nil, errors.Wrap(err, "failed while saving reaction or updating post")
|
||||
}
|
||||
} else {
|
||||
if err := transaction.Commit(); err != nil {
|
||||
return nil, errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
}
|
||||
|
||||
return reaction, nil
|
||||
}
|
||||
|
||||
func (s *SqlReactionStore) Delete(reaction *model.Reaction) (re *model.Reaction, err error) {
|
||||
reaction.PreUpdate()
|
||||
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
|
||||
if err := deleteReactionAndUpdatePost(transaction, reaction); err != nil {
|
||||
return nil, errors.Wrap(err, "deleteReactionAndUpdatePost")
|
||||
}
|
||||
|
||||
if err := transaction.Commit(); err != nil {
|
||||
return nil, errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
|
||||
return reaction, nil
|
||||
}
|
||||
|
||||
// GetForPost returns all reactions associated with `postId` that are not deleted.
|
||||
func (s *SqlReactionStore) GetForPost(postId string, allowFromCache bool) ([]*model.Reaction, error) {
|
||||
queryString, args, err := s.getQueryBuilder().
|
||||
Select("UserId", "PostId", "EmojiName", "CreateAt", "COALESCE(UpdateAt, CreateAt) As UpdateAt",
|
||||
"COALESCE(DeleteAt, 0) As DeleteAt", "RemoteId", "ChannelId").
|
||||
From("Reactions").
|
||||
Where(sq.Eq{"PostId": postId}).
|
||||
Where(sq.Eq{"COALESCE(DeleteAt, 0)": 0}).
|
||||
OrderBy("CreateAt").
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "reactions_getforpost_tosql")
|
||||
}
|
||||
|
||||
var reactions []*model.Reaction
|
||||
if err := s.GetReplicaX().Select(&reactions, queryString, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get Reactions with postId=%s", postId)
|
||||
}
|
||||
return reactions, nil
|
||||
}
|
||||
|
||||
// GetForPostSince returns all reactions associated with `postId` updated after `since`.
|
||||
func (s *SqlReactionStore) GetForPostSince(postId string, since int64, excludeRemoteId string, inclDeleted bool) ([]*model.Reaction, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("UserId", "PostId", "EmojiName", "CreateAt", "COALESCE(UpdateAt, CreateAt) As UpdateAt",
|
||||
"COALESCE(DeleteAt, 0) As DeleteAt", "RemoteId").
|
||||
From("Reactions").
|
||||
Where(sq.Eq{"PostId": postId}).
|
||||
Where(sq.Gt{"UpdateAt": since})
|
||||
|
||||
if excludeRemoteId != "" {
|
||||
query = query.Where(sq.NotEq{"COALESCE(RemoteId, '')": excludeRemoteId})
|
||||
}
|
||||
|
||||
if !inclDeleted {
|
||||
query = query.Where(sq.Eq{"COALESCE(DeleteAt, 0)": 0})
|
||||
}
|
||||
|
||||
query.OrderBy("CreateAt")
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "reactions_getforpostsince_tosql")
|
||||
}
|
||||
|
||||
var reactions []*model.Reaction
|
||||
if err := s.GetReplicaX().Select(&reactions, queryString, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find reactions")
|
||||
}
|
||||
return reactions, nil
|
||||
}
|
||||
|
||||
func (s *SqlReactionStore) BulkGetForPosts(postIds []string) ([]*model.Reaction, error) {
|
||||
placeholder, values := constructArrayArgs(postIds)
|
||||
var reactions []*model.Reaction
|
||||
|
||||
if err := s.GetReplicaX().Select(&reactions,
|
||||
`SELECT
|
||||
UserId,
|
||||
PostId,
|
||||
EmojiName,
|
||||
CreateAt,
|
||||
COALESCE(UpdateAt, CreateAt) As UpdateAt,
|
||||
COALESCE(DeleteAt, 0) As DeleteAt,
|
||||
RemoteId,
|
||||
ChannelId
|
||||
FROM
|
||||
Reactions
|
||||
WHERE
|
||||
PostId IN `+placeholder+` AND COALESCE(DeleteAt, 0) = 0
|
||||
ORDER BY
|
||||
CreateAt`, values...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get Reactions")
|
||||
}
|
||||
return reactions, nil
|
||||
}
|
||||
|
||||
func (s *SqlReactionStore) DeleteAllWithEmojiName(emojiName string) error {
|
||||
var reactions []*model.Reaction
|
||||
now := model.GetMillis()
|
||||
|
||||
if err := s.GetReplicaX().Select(&reactions,
|
||||
`SELECT
|
||||
UserId,
|
||||
PostId,
|
||||
EmojiName,
|
||||
CreateAt,
|
||||
COALESCE(UpdateAt, CreateAt) As UpdateAt,
|
||||
COALESCE(DeleteAt, 0) As DeleteAt,
|
||||
RemoteId
|
||||
FROM
|
||||
Reactions
|
||||
WHERE
|
||||
EmojiName = ? AND COALESCE(DeleteAt, 0) = 0`, emojiName); err != nil {
|
||||
return errors.Wrapf(err, "failed to get Reactions with emojiName=%s", emojiName)
|
||||
}
|
||||
|
||||
_, err := s.GetMasterX().Exec(
|
||||
`UPDATE
|
||||
Reactions
|
||||
SET
|
||||
UpdateAt = ?, DeleteAt = ?
|
||||
WHERE
|
||||
EmojiName = ? AND COALESCE(DeleteAt, 0) = 0`, now, now, emojiName)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Reactions with emojiName=%s", emojiName)
|
||||
}
|
||||
|
||||
for _, reaction := range reactions {
|
||||
reaction := reaction
|
||||
_, err := s.GetMasterX().Exec(UpdatePostHasReactionsOnDeleteQuery, now, reaction.PostId, reaction.PostId)
|
||||
if err != nil {
|
||||
mlog.Warn("Unable to update Post.HasReactions while removing reactions",
|
||||
mlog.String("post_id", reaction.PostId),
|
||||
mlog.Err(err))
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteOrphanedRows removes entries from Reactions when a corresponding post no longer exists.
|
||||
func (s *SqlReactionStore) DeleteOrphanedRows(limit int) (deleted int64, err error) {
|
||||
// We need the extra level of nesting to deal with MySQL's locking
|
||||
const query = `
|
||||
DELETE FROM Reactions WHERE PostId IN (
|
||||
SELECT * FROM (
|
||||
SELECT PostId FROM Reactions
|
||||
LEFT JOIN Posts ON Reactions.PostId = Posts.Id
|
||||
WHERE Posts.Id IS NULL
|
||||
LIMIT ?
|
||||
) AS A
|
||||
)`
|
||||
result, err := s.GetMasterX().Exec(query, limit)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
deleted, err = result.RowsAffected()
|
||||
return
|
||||
}
|
||||
|
||||
func (s *SqlReactionStore) PermanentDeleteBatch(endTime int64, limit int64) (int64, error) {
|
||||
var query string
|
||||
if s.DriverName() == "postgres" {
|
||||
query = "DELETE from Reactions WHERE CreateAt = any (array (SELECT CreateAt FROM Reactions WHERE CreateAt < ? LIMIT ?))"
|
||||
} else {
|
||||
query = "DELETE from Reactions WHERE CreateAt < ? LIMIT ?"
|
||||
}
|
||||
|
||||
sqlResult, err := s.GetMasterX().Exec(query, endTime, limit)
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "failed to delete Reactions")
|
||||
}
|
||||
|
||||
rowsAffected, err := sqlResult.RowsAffected()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "unable to get rows affected for deleted Reactions")
|
||||
}
|
||||
return rowsAffected, nil
|
||||
}
|
||||
|
||||
// GetTopForTeamSince returns the instance counts of the following Reactions sets:
|
||||
// a) those created by anyone in private channels in the given user's membership graph on the given team, and
|
||||
// b) those created by anyone in public channels on the given team.
|
||||
func (s *SqlReactionStore) GetTopForTeamSince(teamID string, userID string, since int64, offset int, limit int) (*model.TopReactionList, error) {
|
||||
reactions := make([]*model.TopReaction, 0)
|
||||
|
||||
query := `
|
||||
SELECT
|
||||
EmojiName,
|
||||
sum(EmojiCount) AS Count
|
||||
FROM ((
|
||||
SELECT
|
||||
EmojiName,
|
||||
count(EmojiName) AS EmojiCount,
|
||||
Reactions.DeleteAt AS DeleteAt,
|
||||
Reactions.CreateAt AS CreateAt
|
||||
FROM
|
||||
ChannelMembers
|
||||
INNER JOIN Channels ON ChannelMembers.ChannelId = Channels.Id
|
||||
INNER JOIN Reactions ON Channels.Id = Reactions.ChannelId
|
||||
WHERE
|
||||
ChannelMembers.UserId = ?
|
||||
AND Channels.Type = 'P'
|
||||
AND Channels.TeamId = ?
|
||||
GROUP BY
|
||||
Reactions.EmojiName,
|
||||
Reactions.DeleteAt,
|
||||
Reactions.CreateAt)
|
||||
UNION ALL (
|
||||
SELECT
|
||||
EmojiName,
|
||||
count(EmojiName) AS EmojiCount,
|
||||
Reactions.DeleteAt AS DeleteAt,
|
||||
Reactions.CreateAt AS CreateAt
|
||||
FROM
|
||||
Reactions
|
||||
INNER JOIN PublicChannels ON Reactions.ChannelId = PublicChannels.Id
|
||||
WHERE
|
||||
PublicChannels.TeamId = ?
|
||||
GROUP BY
|
||||
Reactions.EmojiName,
|
||||
Reactions.DeleteAt,
|
||||
Reactions.CreateAt)) AS A
|
||||
WHERE
|
||||
DeleteAt = 0
|
||||
AND CreateAt > ?
|
||||
GROUP BY
|
||||
EmojiName
|
||||
ORDER BY
|
||||
Count DESC,
|
||||
EmojiName ASC
|
||||
LIMIT ?
|
||||
OFFSET ?`
|
||||
|
||||
if err := s.GetReplicaX().Select(&reactions, query, userID, teamID, teamID, since, limit+1, offset); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get top Reactions")
|
||||
}
|
||||
|
||||
return model.GetTopReactionListWithPagination(reactions, limit), nil
|
||||
}
|
||||
|
||||
// GetTopForUserSince returns the instance counts of the following Reactions sets:
|
||||
// a) those created by the given user in any channel type on the given team (across the workspace if no team is given), and
|
||||
// b) those created by the given user in DM or group channels.
|
||||
func (s *SqlReactionStore) GetTopForUserSince(userID string, teamID string, since int64, offset int, limit int) (*model.TopReactionList, error) {
|
||||
reactions := make([]*model.TopReaction, 0)
|
||||
var args []any
|
||||
var query string
|
||||
|
||||
if teamID != "" {
|
||||
query = `
|
||||
SELECT
|
||||
EmojiName,
|
||||
count(EmojiName) AS Count
|
||||
FROM
|
||||
Reactions
|
||||
INNER JOIN Channels ON Channels.Id = Reactions.ChannelId
|
||||
WHERE
|
||||
Reactions.DeleteAt = 0
|
||||
AND Reactions.UserId = ?
|
||||
AND (Channels.TeamId = ? OR Channels.Type = 'D' OR Channels.Type = 'G')
|
||||
AND Reactions.CreateAt > ?
|
||||
GROUP BY
|
||||
EmojiName
|
||||
ORDER BY
|
||||
Count DESC,
|
||||
EmojiName ASC
|
||||
LIMIT ?
|
||||
OFFSET ?`
|
||||
args = []any{userID, teamID, since, limit + 1, offset}
|
||||
} else {
|
||||
query = `
|
||||
SELECT
|
||||
EmojiName,
|
||||
count(EmojiName) AS Count
|
||||
FROM
|
||||
Reactions
|
||||
WHERE
|
||||
Reactions.DeleteAt = 0
|
||||
AND Reactions.UserId = ?
|
||||
AND Reactions.CreateAt > ?
|
||||
GROUP BY
|
||||
Reactions.EmojiName
|
||||
ORDER BY
|
||||
Count DESC,
|
||||
EmojiName ASC
|
||||
LIMIT ?
|
||||
OFFSET ?`
|
||||
args = []any{userID, since, limit + 1, offset}
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&reactions, query, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get top Reactions")
|
||||
}
|
||||
|
||||
return model.GetTopReactionListWithPagination(reactions, limit), nil
|
||||
}
|
||||
|
||||
func (s *SqlReactionStore) saveReactionAndUpdatePost(transaction *sqlxTxWrapper, reaction *model.Reaction) error {
|
||||
reaction.DeleteAt = 0
|
||||
|
||||
if s.DriverName() == model.DatabaseDriverMysql {
|
||||
if _, err := transaction.NamedExec(
|
||||
`INSERT INTO
|
||||
Reactions
|
||||
(UserId, PostId, EmojiName, CreateAt, UpdateAt, DeleteAt, RemoteId, ChannelId)
|
||||
VALUES
|
||||
(:UserId, :PostId, :EmojiName, :CreateAt, :UpdateAt, :DeleteAt, :RemoteId, :ChannelId)
|
||||
ON DUPLICATE KEY UPDATE
|
||||
UpdateAt = :UpdateAt, DeleteAt = :DeleteAt, RemoteId = :RemoteId, ChannelId = :ChannelId`, reaction); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
if _, err := transaction.NamedExec(
|
||||
`INSERT INTO
|
||||
Reactions
|
||||
(UserId, PostId, EmojiName, CreateAt, UpdateAt, DeleteAt, RemoteId, ChannelId)
|
||||
VALUES
|
||||
(:UserId, :PostId, :EmojiName, :CreateAt, :UpdateAt, :DeleteAt, :RemoteId, :ChannelId)
|
||||
ON CONFLICT (UserId, PostId, EmojiName)
|
||||
DO UPDATE SET UpdateAt = :UpdateAt, DeleteAt = :DeleteAt, RemoteId = :RemoteId, ChannelId = :ChannelId`, reaction); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return updatePostForReactionsOnInsert(transaction, reaction.PostId)
|
||||
}
|
||||
|
||||
func deleteReactionAndUpdatePost(transaction *sqlxTxWrapper, reaction *model.Reaction) error {
|
||||
if _, err := transaction.Exec(
|
||||
`UPDATE
|
||||
Reactions
|
||||
SET
|
||||
UpdateAt = ?, DeleteAt = ?, RemoteId = ?
|
||||
WHERE
|
||||
PostId = ? AND
|
||||
UserId = ? AND
|
||||
EmojiName = ?`, reaction.UpdateAt, reaction.UpdateAt, reaction.RemoteId, reaction.PostId, reaction.UserId, reaction.EmojiName); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return updatePostForReactionsOnDelete(transaction, reaction.PostId)
|
||||
}
|
||||
|
||||
const (
|
||||
UpdatePostHasReactionsOnDeleteQuery = `UPDATE
|
||||
Posts
|
||||
SET
|
||||
UpdateAt = ?,
|
||||
HasReactions = (SELECT count(0) > 0 FROM Reactions WHERE PostId = ? AND COALESCE(DeleteAt, 0) = 0)
|
||||
WHERE
|
||||
Id = ?`
|
||||
)
|
||||
|
||||
func updatePostForReactionsOnDelete(transaction *sqlxTxWrapper, postId string) error {
|
||||
updateAt := model.GetMillis()
|
||||
_, err := transaction.Exec(UpdatePostHasReactionsOnDeleteQuery, updateAt, postId, postId)
|
||||
return err
|
||||
}
|
||||
|
||||
func updatePostForReactionsOnInsert(transaction *sqlxTxWrapper, postId string) error {
|
||||
_, err := transaction.Exec(
|
||||
`UPDATE
|
||||
Posts
|
||||
SET
|
||||
HasReactions = True,
|
||||
UpdateAt = ?
|
||||
WHERE
|
||||
Id = ?`,
|
||||
model.GetMillis(),
|
||||
postId,
|
||||
)
|
||||
|
||||
return err
|
||||
}
|
||||
14
server/channels/store/sqlstore/reaction_store_test.go
Обычный файл
14
server/channels/store/sqlstore/reaction_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestReactionStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestReactionStore)
|
||||
}
|
||||
188
server/channels/store/sqlstore/remote_cluster_store.go
Обычный файл
188
server/channels/store/sqlstore/remote_cluster_store.go
Обычный файл
@@ -0,0 +1,188 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type sqlRemoteClusterStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlRemoteClusterStore(sqlStore *SqlStore) store.RemoteClusterStore {
|
||||
return &sqlRemoteClusterStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s sqlRemoteClusterStore) Save(remoteCluster *model.RemoteCluster) (*model.RemoteCluster, error) {
|
||||
remoteCluster.PreSave()
|
||||
if err := remoteCluster.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query := `INSERT INTO RemoteClusters
|
||||
(RemoteId, RemoteTeamId, Name, DisplayName, SiteURL, CreateAt,
|
||||
LastPingAt, Token, RemoteToken, Topics, CreatorId)
|
||||
VALUES
|
||||
(:RemoteId, :RemoteTeamId, :Name, :DisplayName, :SiteURL, :CreateAt,
|
||||
:LastPingAt, :Token, :RemoteToken, :Topics, :CreatorId)`
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(query, remoteCluster); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save RemoteCluster")
|
||||
}
|
||||
return remoteCluster, nil
|
||||
}
|
||||
|
||||
func (s sqlRemoteClusterStore) Update(remoteCluster *model.RemoteCluster) (*model.RemoteCluster, error) {
|
||||
remoteCluster.PreUpdate()
|
||||
if err := remoteCluster.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query := `UPDATE RemoteClusters
|
||||
SET Token = :Token,
|
||||
RemoteTeamId = :RemoteTeamId,
|
||||
CreateAt = :CreateAt,
|
||||
LastPingAt = :LastPingAt,
|
||||
RemoteToken = :RemoteToken,
|
||||
CreatorId = :CreatorId,
|
||||
DisplayName = :DisplayName,
|
||||
SiteURL = :SiteURL,
|
||||
Topics = :Topics
|
||||
WHERE RemoteId = :RemoteId AND Name = :Name`
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(query, remoteCluster); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to update RemoteCluster")
|
||||
}
|
||||
return remoteCluster, nil
|
||||
}
|
||||
|
||||
func (s sqlRemoteClusterStore) Delete(remoteId string) (bool, error) {
|
||||
squery, args, err := s.getQueryBuilder().
|
||||
Delete("RemoteClusters").
|
||||
Where(sq.Eq{"RemoteId": remoteId}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "delete_remote_cluster_tosql")
|
||||
}
|
||||
|
||||
result, err := s.GetMasterX().Exec(squery, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to delete RemoteCluster")
|
||||
}
|
||||
|
||||
count, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to determine rows affected")
|
||||
}
|
||||
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (s sqlRemoteClusterStore) Get(remoteId string) (*model.RemoteCluster, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("RemoteClusters").
|
||||
Where(sq.Eq{"RemoteId": remoteId})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "remote_cluster_get_tosql")
|
||||
}
|
||||
|
||||
var rc model.RemoteCluster
|
||||
if err := s.GetReplicaX().Get(&rc, queryString, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find RemoteCluster")
|
||||
}
|
||||
return &rc, nil
|
||||
}
|
||||
|
||||
func (s sqlRemoteClusterStore) GetAll(filter model.RemoteClusterQueryFilter) ([]*model.RemoteCluster, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("rc.*").
|
||||
From("RemoteClusters rc")
|
||||
|
||||
if filter.InChannel != "" {
|
||||
query = query.Where("rc.RemoteId IN (SELECT scr.RemoteId FROM SharedChannelRemotes scr WHERE scr.ChannelId = ?)", filter.InChannel)
|
||||
}
|
||||
|
||||
if filter.NotInChannel != "" {
|
||||
query = query.Where("rc.RemoteId NOT IN (SELECT scr.RemoteId FROM SharedChannelRemotes scr WHERE scr.ChannelId = ?)", filter.NotInChannel)
|
||||
}
|
||||
|
||||
if filter.ExcludeOffline {
|
||||
query = query.Where(sq.Gt{"rc.LastPingAt": model.GetMillis() - model.RemoteOfflineAfterMillis})
|
||||
}
|
||||
|
||||
if filter.CreatorId != "" {
|
||||
query = query.Where(sq.Eq{"rc.CreatorId": filter.CreatorId})
|
||||
}
|
||||
|
||||
if filter.OnlyConfirmed {
|
||||
query = query.Where(sq.NotEq{"rc.SiteURL": ""})
|
||||
}
|
||||
|
||||
if filter.Topic != "" {
|
||||
trimmed := strings.TrimSpace(filter.Topic)
|
||||
if trimmed == "" || trimmed == "*" {
|
||||
return nil, errors.New("invalid topic")
|
||||
}
|
||||
queryTopic := fmt.Sprintf("%% %s %%", trimmed)
|
||||
query = query.Where(sq.Or{sq.Like{"rc.Topics": queryTopic}, sq.Eq{"rc.Topics": "*"}})
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "remote_cluster_getall_tosql")
|
||||
}
|
||||
|
||||
list := []*model.RemoteCluster{}
|
||||
if err := s.GetReplicaX().Select(&list, queryString, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find RemoteClusters")
|
||||
}
|
||||
return list, nil
|
||||
}
|
||||
|
||||
func (s sqlRemoteClusterStore) UpdateTopics(remoteClusterid string, topics string) (*model.RemoteCluster, error) {
|
||||
rc, err := s.Get(remoteClusterid)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rc.Topics = topics
|
||||
|
||||
rc.PreUpdate()
|
||||
|
||||
query := `UPDATE RemoteClusters
|
||||
SET Topics = :Topics
|
||||
WHERE RemoteId = :RemoteId`
|
||||
|
||||
if _, err = s.GetMasterX().NamedExec(query, rc); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return rc, nil
|
||||
}
|
||||
|
||||
func (s sqlRemoteClusterStore) SetLastPingAt(remoteClusterId string) error {
|
||||
query := s.getQueryBuilder().
|
||||
Update("RemoteClusters").
|
||||
Set("LastPingAt", model.GetMillis()).
|
||||
Where(sq.Eq{"RemoteId": remoteClusterId})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "remote_cluster_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to update RemoteCluster")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/remote_cluster_store_test.go
Обычный файл
14
server/channels/store/sqlstore/remote_cluster_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestRemoteClusterStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestRemoteClusterStore)
|
||||
}
|
||||
988
server/channels/store/sqlstore/retention_policy_store.go
Обычный файл
988
server/channels/store/sqlstore/retention_policy_store.go
Обычный файл
@@ -0,0 +1,988 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/go-sql-driver/mysql"
|
||||
"github.com/lib/pq"
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/einterfaces"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlRetentionPolicyStore struct {
|
||||
*SqlStore
|
||||
metrics einterfaces.MetricsInterface
|
||||
}
|
||||
|
||||
func newSqlRetentionPolicyStore(sqlStore *SqlStore, metrics einterfaces.MetricsInterface) store.RetentionPolicyStore {
|
||||
return &SqlRetentionPolicyStore{
|
||||
SqlStore: sqlStore,
|
||||
metrics: metrics,
|
||||
}
|
||||
}
|
||||
|
||||
// executePossiblyEmptyQuery only executes the query if it is non-empty. This helps avoid
|
||||
// having to check for MySQL, which, unlike Postgres, does not allow empty queries.
|
||||
func executePossiblyEmptyQuery(txn *sqlxTxWrapper, query string, args ...any) (sql.Result, error) {
|
||||
if query == "" {
|
||||
return nil, nil
|
||||
}
|
||||
return txn.Exec(query, args...)
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) Save(policy *model.RetentionPolicyWithTeamAndChannelIDs) (_ *model.RetentionPolicyWithTeamAndChannelCounts, err error) {
|
||||
// Strategy:
|
||||
// 1. Insert new policy
|
||||
// 2. Insert new channels into policy
|
||||
// 3. Insert new teams into policy
|
||||
|
||||
if err = s.checkTeamsExist(policy.TeamIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = s.checkChannelsExist(policy.ChannelIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
policy.ID = model.NewId()
|
||||
|
||||
policyInsertQuery, policyInsertArgs, err := s.getQueryBuilder().
|
||||
Insert("RetentionPolicies").
|
||||
Columns("Id", "DisplayName", "PostDuration").
|
||||
Values(policy.ID, policy.DisplayName, policy.PostDurationDays).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
channelsInsertQuery, channelsInsertArgs, err := s.buildInsertRetentionPoliciesChannelsQuery(policy.ID, policy.ChannelIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
teamsInsertQuery, teamsInsertArgs, err := s.buildInsertRetentionPoliciesTeamsQuery(policy.ID, policy.TeamIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
queryString, args, err := s.buildGetPolicyQuery(policy.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
txn, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer finalizeTransactionX(txn, &err)
|
||||
|
||||
// Create a new policy in RetentionPolicies
|
||||
if _, err = txn.Exec(policyInsertQuery, policyInsertArgs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Insert the channel IDs into RetentionPoliciesChannels
|
||||
if _, err = executePossiblyEmptyQuery(txn, channelsInsertQuery, channelsInsertArgs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Insert the team IDs into RetentionPoliciesTeams
|
||||
if _, err = executePossiblyEmptyQuery(txn, teamsInsertQuery, teamsInsertArgs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Select the new policy (with team/channel counts) which we just created
|
||||
var newPolicy model.RetentionPolicyWithTeamAndChannelCounts
|
||||
|
||||
if err = txn.Get(&newPolicy, queryString, args...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = txn.Commit(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &newPolicy, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) checkTeamsExist(teamIDs []string) error {
|
||||
if len(teamIDs) > 0 {
|
||||
teamsSelectQuery, teamsSelectArgs, err := s.getQueryBuilder().
|
||||
Select("Id").
|
||||
From("Teams").
|
||||
Where(sq.Eq{"Id": teamIDs}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rows := []*string{}
|
||||
err = s.GetReplicaX().Select(&rows, teamsSelectQuery, teamsSelectArgs...)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(rows) == len(teamIDs) {
|
||||
return nil
|
||||
}
|
||||
retrievedIDs := make(map[string]bool)
|
||||
for _, teamID := range rows {
|
||||
retrievedIDs[*teamID] = true
|
||||
}
|
||||
for _, teamID := range teamIDs {
|
||||
if _, ok := retrievedIDs[teamID]; !ok {
|
||||
return store.NewErrNotFound("Team", teamID)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) checkChannelsExist(channelIDs []string) error {
|
||||
if len(channelIDs) > 0 {
|
||||
channelsSelectQuery, channelsSelectArgs, err := s.getQueryBuilder().
|
||||
Select("Id").
|
||||
From("Channels").
|
||||
Where(sq.Eq{"Id": channelIDs}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rows := []*string{}
|
||||
err = s.GetReplicaX().Select(&rows, channelsSelectQuery, channelsSelectArgs...)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(rows) == len(channelIDs) {
|
||||
return nil
|
||||
}
|
||||
retrievedIDs := make(map[string]bool)
|
||||
for _, channelID := range rows {
|
||||
retrievedIDs[*channelID] = true
|
||||
}
|
||||
for _, channelID := range channelIDs {
|
||||
if _, ok := retrievedIDs[channelID]; !ok {
|
||||
return store.NewErrNotFound("Channel", channelID)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) buildInsertRetentionPoliciesChannelsQuery(policyID string, channelIDs []string) (query string, args []any, err error) {
|
||||
if len(channelIDs) > 0 {
|
||||
builder := s.getQueryBuilder().
|
||||
Insert("RetentionPoliciesChannels").
|
||||
Columns("PolicyId", "ChannelId")
|
||||
for _, channelID := range channelIDs {
|
||||
builder = builder.Values(policyID, channelID)
|
||||
}
|
||||
query, args, err = builder.ToSql()
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) buildInsertRetentionPoliciesTeamsQuery(policyID string, teamIDs []string) (query string, args []any, err error) {
|
||||
if len(teamIDs) > 0 {
|
||||
builder := s.getQueryBuilder().
|
||||
Insert("RetentionPoliciesTeams").
|
||||
Columns("PolicyId", "TeamId")
|
||||
for _, teamID := range teamIDs {
|
||||
builder = builder.Values(policyID, teamID)
|
||||
}
|
||||
query, args, err = builder.ToSql()
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) Patch(patch *model.RetentionPolicyWithTeamAndChannelIDs) (_ *model.RetentionPolicyWithTeamAndChannelCounts, err error) {
|
||||
// Strategy:
|
||||
// 1. Update policy attributes
|
||||
// 2. Delete existing channels from policy
|
||||
// 3. Insert new channels into policy
|
||||
// 4. Delete existing teams from policy
|
||||
// 5. Insert new teams into policy
|
||||
// 6. Read new policy
|
||||
|
||||
if err = s.checkTeamsExist(patch.TeamIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = s.checkChannelsExist(patch.ChannelIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
policyUpdateQuery := ""
|
||||
policyUpdateArgs := []any{}
|
||||
if patch.DisplayName != "" || patch.PostDurationDays != nil {
|
||||
builder := s.getQueryBuilder().Update("RetentionPolicies")
|
||||
if patch.DisplayName != "" {
|
||||
builder = builder.Set("DisplayName", patch.DisplayName)
|
||||
}
|
||||
if patch.PostDurationDays != nil {
|
||||
builder = builder.Set("PostDuration", *patch.PostDurationDays)
|
||||
}
|
||||
policyUpdateQuery, policyUpdateArgs, err = builder.
|
||||
Where(sq.Eq{"Id": patch.ID}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
channelsDeleteQuery := ""
|
||||
channelsDeleteArgs := []any{}
|
||||
channelsInsertQuery := ""
|
||||
channelsInsertArgs := []any{}
|
||||
if patch.ChannelIDs != nil {
|
||||
channelsDeleteQuery, channelsDeleteArgs, err = s.getQueryBuilder().
|
||||
Delete("RetentionPoliciesChannels").
|
||||
Where(sq.Eq{"PolicyId": patch.ID}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
channelsInsertQuery, channelsInsertArgs, err = s.buildInsertRetentionPoliciesChannelsQuery(patch.ID, patch.ChannelIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
teamsDeleteQuery := ""
|
||||
teamsDeleteArgs := []any{}
|
||||
teamsInsertQuery := ""
|
||||
teamsInsertArgs := []any{}
|
||||
if patch.TeamIDs != nil {
|
||||
teamsDeleteQuery, teamsDeleteArgs, err = s.getQueryBuilder().
|
||||
Delete("RetentionPoliciesTeams").
|
||||
Where(sq.Eq{"PolicyId": patch.ID}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
teamsInsertQuery, teamsInsertArgs, err = s.buildInsertRetentionPoliciesTeamsQuery(patch.ID, patch.TeamIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
queryString, args, err := s.buildGetPolicyQuery(patch.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
txn, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer finalizeTransactionX(txn, &err)
|
||||
|
||||
// Update the fields of the policy in RetentionPolicies
|
||||
if _, err = executePossiblyEmptyQuery(txn, policyUpdateQuery, policyUpdateArgs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Remove all channels from the policy in RetentionPoliciesChannels
|
||||
if _, err = executePossiblyEmptyQuery(txn, channelsDeleteQuery, channelsDeleteArgs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Insert the new channels for the policy in RetentionPoliciesChannels
|
||||
if _, err = executePossiblyEmptyQuery(txn, channelsInsertQuery, channelsInsertArgs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Remove all teams from the policy in RetentionPoliciesTeams
|
||||
if _, err = executePossiblyEmptyQuery(txn, teamsDeleteQuery, teamsDeleteArgs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Insert the new teams for the policy in RetentionPoliciesTeams
|
||||
if _, err = executePossiblyEmptyQuery(txn, teamsInsertQuery, teamsInsertArgs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Select the policy which we just updated
|
||||
var newPolicy model.RetentionPolicyWithTeamAndChannelCounts
|
||||
if err = txn.Get(&newPolicy, queryString, args...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err = txn.Commit(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &newPolicy, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) buildGetPolicyQuery(id string) (string, []any, error) {
|
||||
return s.buildGetPoliciesQuery(id, 0, 1)
|
||||
}
|
||||
|
||||
// buildGetPoliciesQuery builds a query to select information for the policy with the specified
|
||||
// ID, or, if `id` is the empty string, from all policies. The results returned will be sorted by
|
||||
// policy display name and ID.
|
||||
func (s *SqlRetentionPolicyStore) buildGetPoliciesQuery(id string, offset, limit int) (string, []any, error) {
|
||||
rpcSubQuery := s.getQueryBuilder().
|
||||
Select("RetentionPolicies.Id, COUNT(RetentionPoliciesChannels.ChannelId) AS Count").
|
||||
From("RetentionPolicies").
|
||||
LeftJoin("RetentionPoliciesChannels ON RetentionPolicies.Id = RetentionPoliciesChannels.PolicyId").
|
||||
GroupBy("RetentionPolicies.Id").
|
||||
OrderBy("RetentionPolicies.DisplayName, RetentionPolicies.Id").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset))
|
||||
|
||||
if id != "" {
|
||||
rpcSubQuery = rpcSubQuery.Where(sq.Eq{"RetentionPolicies.Id": id})
|
||||
}
|
||||
|
||||
rpcSubQueryString, args, err := rpcSubQuery.ToSql()
|
||||
if err != nil {
|
||||
return "", nil, errors.Wrap(err, "retention_policies_tosql")
|
||||
}
|
||||
|
||||
rptSubQuery := s.getQueryBuilder().
|
||||
Select("RetentionPolicies.Id, COUNT(RetentionPoliciesTeams.TeamId) AS Count").
|
||||
From("RetentionPolicies").
|
||||
LeftJoin("RetentionPoliciesTeams ON RetentionPolicies.Id = RetentionPoliciesTeams.PolicyId").
|
||||
GroupBy("RetentionPolicies.Id").
|
||||
OrderBy("RetentionPolicies.DisplayName, RetentionPolicies.Id").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset))
|
||||
|
||||
if id != "" {
|
||||
rptSubQuery = rptSubQuery.Where(sq.Eq{"RetentionPolicies.Id": id})
|
||||
}
|
||||
|
||||
rptSubQueryString, _, err := rptSubQuery.ToSql()
|
||||
if err != nil {
|
||||
return "", nil, errors.Wrap(err, "retention_policies_tosql")
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select(`
|
||||
RetentionPolicies.Id as "Id",
|
||||
RetentionPolicies.DisplayName,
|
||||
RetentionPolicies.PostDuration as "PostDuration",
|
||||
A.Count AS ChannelCount,
|
||||
B.Count AS TeamCount
|
||||
`).
|
||||
From("RetentionPolicies").
|
||||
InnerJoin(`(` + rpcSubQueryString + `) AS A ON RetentionPolicies.Id = A.Id`).
|
||||
InnerJoin(`(` + rptSubQueryString + `) AS B ON RetentionPolicies.Id = B.Id`).
|
||||
OrderBy("RetentionPolicies.DisplayName, RetentionPolicies.Id")
|
||||
|
||||
queryString, _, err := query.ToSql()
|
||||
if err != nil {
|
||||
return "", nil, errors.Wrap(err, "retention_policies_tosql")
|
||||
}
|
||||
|
||||
// MySQL does not support positional params, so we add one param for each WHERE clause.
|
||||
if s.DriverName() == model.DatabaseDriverMysql {
|
||||
args = append(args, args...)
|
||||
}
|
||||
|
||||
return queryString, args, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) Get(id string) (*model.RetentionPolicyWithTeamAndChannelCounts, error) {
|
||||
queryString, args, err := s.buildGetPolicyQuery(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var policy model.RetentionPolicyWithTeamAndChannelCounts
|
||||
if err := s.GetReplicaX().Get(&policy, queryString, args...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &policy, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetAll(offset, limit int) ([]*model.RetentionPolicyWithTeamAndChannelCounts, error) {
|
||||
policies := []*model.RetentionPolicyWithTeamAndChannelCounts{}
|
||||
queryString, args, err := s.buildGetPoliciesQuery("", offset, limit)
|
||||
if err != nil {
|
||||
return policies, err
|
||||
}
|
||||
err = s.GetReplicaX().Select(&policies, queryString, args...)
|
||||
return policies, err
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetCount() (int64, error) {
|
||||
var count int64
|
||||
err := s.GetReplicaX().Get(&count, "SELECT COUNT(*) FROM RetentionPolicies")
|
||||
if err != nil {
|
||||
return count, err
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) Delete(id string) error {
|
||||
query := s.getQueryBuilder().
|
||||
Delete("RetentionPolicies").
|
||||
Where(sq.Eq{"Id": id})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "retention_policies_tosql")
|
||||
}
|
||||
|
||||
sqlResult, err := s.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to permanent delete retention policy with id=%s", id)
|
||||
}
|
||||
|
||||
numRowsAffected, err := sqlResult.RowsAffected()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "unable to get rows affected")
|
||||
} else if numRowsAffected == 0 {
|
||||
return errors.New("policy not found")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetChannels(policyId string, offset, limit int) (model.ChannelListWithTeamData, error) {
|
||||
query := s.getQueryBuilder().Select(`Channels.*, Teams.DisplayName AS TeamDisplayName,
|
||||
Teams.Name AS TeamName,Teams.UpdateAt AS TeamUpdateAt`).
|
||||
From("RetentionPoliciesChannels").
|
||||
InnerJoin("Channels ON RetentionPoliciesChannels.ChannelId = Channels.Id").
|
||||
InnerJoin("Teams ON Channels.TeamId = Teams.Id").
|
||||
Where(sq.Eq{"RetentionPoliciesChannels.PolicyId": policyId}).
|
||||
OrderBy("Channels.DisplayName, Channels.Id").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset))
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "retention_policies_channels_tosql")
|
||||
}
|
||||
|
||||
channels := model.ChannelListWithTeamData{}
|
||||
if err := s.GetReplicaX().Select(&channels, queryString, args...); err != nil {
|
||||
return channels, errors.Wrap(err, "failed to find RetentionPoliciesChannels")
|
||||
}
|
||||
|
||||
for _, channel := range channels {
|
||||
channel.PolicyID = model.NewString(policyId)
|
||||
}
|
||||
|
||||
return channels, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetChannelsCount(policyId string) (int64, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("Count(*)").
|
||||
From("RetentionPolicies").
|
||||
InnerJoin("RetentionPoliciesChannels ON RetentionPolicies.Id = RetentionPoliciesChannels.PolicyId").
|
||||
Where(sq.Eq{"RetentionPolicies.Id": policyId})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "retention_policies_tosql")
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := s.GetReplicaX().Get(&count, queryString, args...); err != nil {
|
||||
return 0, errors.Wrap(err, "failed to count RetentionPolicies")
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) AddChannels(policyId string, channelIds []string) error {
|
||||
if len(channelIds) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := s.checkChannelsExist(channelIds); err != nil {
|
||||
return err
|
||||
}
|
||||
query := s.getQueryBuilder().
|
||||
Insert("RetentionPoliciesChannels").
|
||||
Columns("policyId", "channelId")
|
||||
|
||||
for _, channelId := range channelIds {
|
||||
query = query.Values(policyId, channelId)
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "retention_policies_channels_tosql")
|
||||
}
|
||||
|
||||
_, err = s.GetMasterX().Exec(queryString, args...)
|
||||
if err != nil {
|
||||
switch dbErr := err.(type) {
|
||||
case *pq.Error:
|
||||
if dbErr.Code == PGForeignKeyViolationErrorCode {
|
||||
return store.NewErrNotFound("RetentionPolicy", policyId)
|
||||
}
|
||||
case *mysql.MySQLError:
|
||||
if dbErr.Number == MySQLForeignKeyViolationErrorCode {
|
||||
return store.NewErrNotFound("RetentionPolicy", policyId)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) RemoveChannels(policyId string, channelIds []string) error {
|
||||
if len(channelIds) == 0 {
|
||||
return nil
|
||||
}
|
||||
query := s.getQueryBuilder().
|
||||
Delete("RetentionPoliciesChannels").
|
||||
Where(sq.And{
|
||||
sq.Eq{"PolicyId": policyId},
|
||||
sq.Eq{"ChannelId": channelIds},
|
||||
})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "retention_policies_channels_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrapf(err, "failed to permanent delete retention policy channels with policyid=%s", policyId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetTeams(policyId string, offset, limit int) ([]*model.Team, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("Teams.*").
|
||||
From("RetentionPoliciesTeams").
|
||||
InnerJoin("Teams ON RetentionPoliciesTeams.TeamId = Teams.Id").
|
||||
Where(sq.Eq{"RetentionPoliciesTeams.PolicyId": policyId}).
|
||||
OrderBy("Teams.DisplayName, Teams.Id").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset))
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "retention_policies_teams_tosql")
|
||||
}
|
||||
|
||||
teams := []*model.Team{}
|
||||
if err = s.GetReplicaX().Select(&teams, queryString, args...); err != nil {
|
||||
return teams, errors.Wrap(err, "failed to find Teams")
|
||||
}
|
||||
|
||||
return teams, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetTeamsCount(policyId string) (int64, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("Count(*)").
|
||||
From("RetentionPolicies").
|
||||
InnerJoin("RetentionPoliciesTeams ON RetentionPolicies.Id = RetentionPoliciesTeams.PolicyId").
|
||||
Where(sq.Eq{"RetentionPolicies.Id": policyId})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "retention_policies_tosql")
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := s.GetReplicaX().Get(&count, queryString, args...); err != nil {
|
||||
return 0, errors.Wrap(err, "failed to count RetentionPolicies")
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) AddTeams(policyId string, teamIds []string) error {
|
||||
if len(teamIds) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := s.checkTeamsExist(teamIds); err != nil {
|
||||
return err
|
||||
}
|
||||
query := s.getQueryBuilder().
|
||||
Insert("RetentionPoliciesTeams").
|
||||
Columns("PolicyId", "TeamId")
|
||||
for _, teamId := range teamIds {
|
||||
query = query.Values(policyId, teamId)
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "retention_policies_teams_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to insert retention policies teams")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) RemoveTeams(policyId string, teamIds []string) error {
|
||||
if len(teamIds) == 0 {
|
||||
return nil
|
||||
}
|
||||
query := s.getQueryBuilder().
|
||||
Delete("RetentionPoliciesTeams").
|
||||
Where(sq.And{
|
||||
sq.Eq{"PolicyId": policyId},
|
||||
sq.Eq{"TeamId": teamIds},
|
||||
})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "retention_policies_teams_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrapf(err, "unable to permanent delete retention policies teams with policyid=%s", policyId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func subQueryIN(property string, query sq.SelectBuilder) sq.Sqlizer {
|
||||
queryString, args := query.MustSql()
|
||||
subQuery := fmt.Sprintf("%s IN (SELECT * FROM (%s) AS A)", property, queryString)
|
||||
return sq.Expr(subQuery, args...)
|
||||
}
|
||||
|
||||
// DeleteOrphanedRows removes entries from RetentionPoliciesChannels and RetentionPoliciesTeams
|
||||
// where a channel or team no longer exists.
|
||||
func (s *SqlRetentionPolicyStore) DeleteOrphanedRows(limit int) (deleted int64, err error) {
|
||||
// We need the extra level of nesting to deal with MySQL's locking
|
||||
rpcSubQuery := sq.Select("ChannelId").
|
||||
From("RetentionPoliciesChannels").
|
||||
LeftJoin("Channels ON RetentionPoliciesChannels.ChannelId = Channels.Id").
|
||||
Where("Channels.Id IS NULL").
|
||||
Limit(uint64(limit))
|
||||
|
||||
rpcDeleteQuery, rpcArgs, err := s.getQueryBuilder().
|
||||
Delete("RetentionPoliciesChannels").
|
||||
Where(subQueryIN("ChannelId", rpcSubQuery)).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrap(err, "retention_policies_channels_tosql")
|
||||
}
|
||||
|
||||
rptSubQuery := sq.Select("TeamId").
|
||||
From("RetentionPoliciesTeams").
|
||||
LeftJoin("Teams ON RetentionPoliciesTeams.TeamId = Teams.Id").
|
||||
Where("Teams.Id IS NULL").
|
||||
Limit(uint64(limit))
|
||||
|
||||
rptDeleteQuery, rptArgs, err := s.getQueryBuilder().
|
||||
Delete("RetentionPoliciesTeams").
|
||||
Where(subQueryIN("TeamId", rptSubQuery)).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return int64(0), errors.Wrap(err, "retention_policies_teams_tosql")
|
||||
}
|
||||
|
||||
result, err := s.GetMasterX().Exec(rpcDeleteQuery, rpcArgs...)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
rpcDeleted, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
result, err = s.GetMasterX().Exec(rptDeleteQuery, rptArgs...)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
rptDeleted, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
deleted = rpcDeleted + rptDeleted
|
||||
return
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetTeamPoliciesForUser(userID string, offset, limit int) ([]*model.RetentionPolicyForTeam, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select(`Teams.Id AS "Id", RetentionPolicies.PostDuration AS "PostDuration"`).
|
||||
From("Users").
|
||||
InnerJoin("TeamMembers ON Users.Id = TeamMembers.UserId").
|
||||
InnerJoin("Teams ON TeamMembers.TeamId = Teams.Id").
|
||||
InnerJoin("RetentionPoliciesTeams ON Teams.Id = RetentionPoliciesTeams.TeamId").
|
||||
InnerJoin("RetentionPolicies ON RetentionPoliciesTeams.PolicyId = RetentionPolicies.Id").
|
||||
Where(
|
||||
sq.And{
|
||||
sq.Eq{"Users.Id": userID},
|
||||
sq.Eq{"TeamMembers.DeleteAt": 0},
|
||||
sq.Eq{"Teams.DeleteAt": 0},
|
||||
},
|
||||
).
|
||||
OrderBy("Teams.Id").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset))
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "team_policies_for_user_tosql")
|
||||
}
|
||||
|
||||
policies := []*model.RetentionPolicyForTeam{}
|
||||
if err := s.GetReplicaX().Select(&policies, queryString, args...); err != nil {
|
||||
return policies, errors.Wrap(err, "failed to find Users")
|
||||
}
|
||||
|
||||
return policies, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetTeamPoliciesCountForUser(userID string) (int64, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("Count(*)").
|
||||
From("Users").
|
||||
InnerJoin("TeamMembers ON Users.Id = TeamMembers.UserId").
|
||||
InnerJoin("Teams ON TeamMembers.TeamId = Teams.Id").
|
||||
InnerJoin("RetentionPoliciesTeams ON Teams.Id = RetentionPoliciesTeams.TeamId").
|
||||
InnerJoin("RetentionPolicies ON RetentionPoliciesTeams.PolicyId = RetentionPolicies.Id").
|
||||
Where(
|
||||
sq.And{
|
||||
sq.Eq{"Users.Id": userID},
|
||||
sq.Eq{"TeamMembers.DeleteAt": 0},
|
||||
sq.Eq{"Teams.DeleteAt": 0},
|
||||
},
|
||||
)
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "team_policies_count_for_user_tosql")
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := s.GetReplicaX().Get(&count, queryString, args...); err != nil {
|
||||
return 0, errors.Wrap(err, "failed to count TeamPoliciesCountForUser")
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetChannelPoliciesForUser(userID string, offset, limit int) ([]*model.RetentionPolicyForChannel, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select(`Channels.Id as "Id", RetentionPolicies.PostDuration as "PostDuration"`).
|
||||
From("Users").
|
||||
InnerJoin("ChannelMembers ON Users.Id = ChannelMembers.UserId").
|
||||
InnerJoin("Channels ON ChannelMembers.ChannelId = Channels.Id").
|
||||
InnerJoin("RetentionPoliciesChannels ON Channels.Id = RetentionPoliciesChannels.ChannelId").
|
||||
InnerJoin("RetentionPolicies ON RetentionPoliciesChannels.PolicyId = RetentionPolicies.Id").
|
||||
Where(
|
||||
sq.And{
|
||||
sq.Eq{"Users.Id": userID},
|
||||
sq.Eq{"Channels.DeleteAt": 0},
|
||||
},
|
||||
).
|
||||
OrderBy("Channels.Id").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset))
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "channel_policies_for_user_tosql")
|
||||
}
|
||||
|
||||
policies := []*model.RetentionPolicyForChannel{}
|
||||
if err := s.GetReplicaX().Select(&policies, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Users")
|
||||
}
|
||||
|
||||
return policies, nil
|
||||
}
|
||||
|
||||
func (s *SqlRetentionPolicyStore) GetChannelPoliciesCountForUser(userID string) (int64, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("Count(*)").
|
||||
From("Users").
|
||||
InnerJoin("ChannelMembers ON Users.Id = ChannelMembers.UserId").
|
||||
InnerJoin("Channels ON ChannelMembers.ChannelId = Channels.Id").
|
||||
InnerJoin("RetentionPoliciesChannels ON Channels.Id = RetentionPoliciesChannels.ChannelId").
|
||||
InnerJoin("RetentionPolicies ON RetentionPoliciesChannels.PolicyId = RetentionPolicies.Id").
|
||||
Where(
|
||||
sq.And{
|
||||
sq.Eq{"Users.Id": userID},
|
||||
sq.Eq{"Channels.DeleteAt": 0},
|
||||
},
|
||||
)
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "channel_policies_count_users_tosql")
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := s.GetReplicaX().Get(&count, queryString, args...); err != nil {
|
||||
return 0, errors.Wrap(err, "failed to count ChannelPoliciesCountForUser")
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// RetentionPolicyBatchDeletionInfo gives information on how to delete records
|
||||
// under a retention policy; see `genericPermanentDeleteBatchForRetentionPolicies`.
|
||||
//
|
||||
// `BaseBuilder` should already have selected the primary key(s) for the main table
|
||||
// and should be joined to a table with a ChannelId column, which will be used to join
|
||||
// on the Channels table.
|
||||
// `Table` is the name of the table from which records are being deleted.
|
||||
// `TimeColumn` is the name of the column which contains the timestamp of the record.
|
||||
// `PrimaryKeys` contains the primary keys of `table`. It should be the same as the
|
||||
// `From` clause in `baseBuilder`.
|
||||
// `ChannelIDTable` is the table which contains the ChannelId column, it may be the
|
||||
// same as `table`, or will be different if a join was used.
|
||||
// `NowMillis` must be a Unix timestamp in milliseconds and is used by the granular
|
||||
// policies; if `nowMillis - timestamp(record)` is greater than
|
||||
// the post duration of a granular policy, than the record will be deleted.
|
||||
// `GlobalPolicyEndTime` is used by the global policy; any record older than this time
|
||||
// will be deleted by the global policy if it does not fall under a granular policy.
|
||||
// To disable the granular policies, set `NowMillis` to 0.
|
||||
// To disable the global policy, set `GlobalPolicyEndTime` to 0.
|
||||
type RetentionPolicyBatchDeletionInfo struct {
|
||||
BaseBuilder sq.SelectBuilder
|
||||
Table string
|
||||
TimeColumn string
|
||||
PrimaryKeys []string
|
||||
ChannelIDTable string
|
||||
NowMillis int64
|
||||
GlobalPolicyEndTime int64
|
||||
Limit int64
|
||||
}
|
||||
|
||||
// genericPermanentDeleteBatchForRetentionPolicies is a helper function for tables
|
||||
// which need to delete records for granular and global policies.
|
||||
func genericPermanentDeleteBatchForRetentionPolicies(
|
||||
r RetentionPolicyBatchDeletionInfo,
|
||||
s *SqlStore,
|
||||
cursor model.RetentionPolicyCursor,
|
||||
) (int64, model.RetentionPolicyCursor, error) {
|
||||
baseBuilder := r.BaseBuilder.InnerJoin("Channels ON " + r.ChannelIDTable + ".ChannelId = Channels.Id")
|
||||
|
||||
scopedTimeColumn := r.Table + "." + r.TimeColumn
|
||||
nowStr := strconv.FormatInt(r.NowMillis, 10)
|
||||
// A record falls under the scope of a granular retention policy if:
|
||||
// 1. The policy's post duration is >= 0
|
||||
// 2. The record's lifespan has not exceeded the policy's post duration
|
||||
const millisecondsInADay = 24 * 60 * 60 * 1000
|
||||
fallsUnderGranularPolicy := sq.And{
|
||||
sq.GtOrEq{"RetentionPolicies.PostDuration": 0},
|
||||
sq.Expr(nowStr + " - " + scopedTimeColumn + " > RetentionPolicies.PostDuration * " + strconv.FormatInt(millisecondsInADay, 10)),
|
||||
}
|
||||
|
||||
// If the caller wants to disable the global policy from running
|
||||
if r.GlobalPolicyEndTime <= 0 {
|
||||
cursor.GlobalPoliciesDone = true
|
||||
}
|
||||
// If the caller wants to disable the granular policies from running
|
||||
if r.NowMillis <= 0 {
|
||||
cursor.ChannelPoliciesDone = true
|
||||
cursor.TeamPoliciesDone = true
|
||||
}
|
||||
|
||||
var totalRowsAffected int64
|
||||
|
||||
// First, delete all of the records which fall under the scope of a channel-specific policy
|
||||
if !cursor.ChannelPoliciesDone {
|
||||
channelPoliciesBuilder := baseBuilder.
|
||||
InnerJoin("RetentionPoliciesChannels ON " + r.ChannelIDTable + ".ChannelId = RetentionPoliciesChannels.ChannelId").
|
||||
InnerJoin("RetentionPolicies ON RetentionPoliciesChannels.PolicyId = RetentionPolicies.Id").
|
||||
Where(fallsUnderGranularPolicy).
|
||||
Limit(uint64(r.Limit))
|
||||
rowsAffected, err := genericRetentionPoliciesDeletion(channelPoliciesBuilder, r, s)
|
||||
if err != nil {
|
||||
return 0, cursor, err
|
||||
}
|
||||
if rowsAffected < r.Limit {
|
||||
cursor.ChannelPoliciesDone = true
|
||||
}
|
||||
totalRowsAffected += rowsAffected
|
||||
r.Limit -= rowsAffected
|
||||
}
|
||||
|
||||
// Next, delete all of the records which fall under the scope of a team-specific policy
|
||||
if cursor.ChannelPoliciesDone && !cursor.TeamPoliciesDone {
|
||||
// Channel-specific policies override team-specific policies.
|
||||
teamPoliciesBuilder := baseBuilder.
|
||||
LeftJoin("RetentionPoliciesChannels ON " + r.ChannelIDTable + ".ChannelId = RetentionPoliciesChannels.ChannelId").
|
||||
InnerJoin("RetentionPoliciesTeams ON Channels.TeamId = RetentionPoliciesTeams.TeamId").
|
||||
InnerJoin("RetentionPolicies ON RetentionPoliciesTeams.PolicyId = RetentionPolicies.Id").
|
||||
Where(sq.And{
|
||||
sq.Eq{"RetentionPoliciesChannels.PolicyId": nil},
|
||||
sq.Expr("RetentionPoliciesTeams.PolicyId = RetentionPolicies.Id"),
|
||||
}).
|
||||
Where(fallsUnderGranularPolicy).
|
||||
Limit(uint64(r.Limit))
|
||||
rowsAffected, err := genericRetentionPoliciesDeletion(teamPoliciesBuilder, r, s)
|
||||
if err != nil {
|
||||
return 0, cursor, err
|
||||
}
|
||||
if rowsAffected < r.Limit {
|
||||
cursor.TeamPoliciesDone = true
|
||||
}
|
||||
totalRowsAffected += rowsAffected
|
||||
r.Limit -= rowsAffected
|
||||
}
|
||||
|
||||
// Finally, delete all of the records which fall under the scope of the global policy
|
||||
if cursor.ChannelPoliciesDone && cursor.TeamPoliciesDone && !cursor.GlobalPoliciesDone {
|
||||
// Granular policies override the global policy.
|
||||
globalPolicyBuilder := baseBuilder.
|
||||
LeftJoin("RetentionPoliciesChannels ON " + r.ChannelIDTable + ".ChannelId = RetentionPoliciesChannels.ChannelId").
|
||||
LeftJoin("RetentionPoliciesTeams ON Channels.TeamId = RetentionPoliciesTeams.TeamId").
|
||||
LeftJoin("RetentionPolicies ON RetentionPoliciesChannels.PolicyId = RetentionPolicies.Id").
|
||||
Where(sq.And{
|
||||
sq.Eq{"RetentionPoliciesChannels.PolicyId": nil},
|
||||
sq.Eq{"RetentionPoliciesTeams.PolicyId": nil},
|
||||
}).
|
||||
Where(sq.Lt{scopedTimeColumn: r.GlobalPolicyEndTime}).
|
||||
Limit(uint64(r.Limit))
|
||||
rowsAffected, err := genericRetentionPoliciesDeletion(globalPolicyBuilder, r, s)
|
||||
if err != nil {
|
||||
return 0, cursor, err
|
||||
}
|
||||
if rowsAffected < r.Limit {
|
||||
cursor.GlobalPoliciesDone = true
|
||||
}
|
||||
totalRowsAffected += rowsAffected
|
||||
}
|
||||
|
||||
return totalRowsAffected, cursor, nil
|
||||
}
|
||||
|
||||
// genericRetentionPoliciesDeletion actually executes the DELETE query using a sq.SelectBuilder
|
||||
// which selects the rows to delete.
|
||||
func genericRetentionPoliciesDeletion(
|
||||
builder sq.SelectBuilder,
|
||||
r RetentionPolicyBatchDeletionInfo,
|
||||
s *SqlStore,
|
||||
) (rowsAffected int64, err error) {
|
||||
query, args, err := builder.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, r.Table+"_tosql")
|
||||
}
|
||||
if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
primaryKeysStr := "(" + strings.Join(r.PrimaryKeys, ",") + ")"
|
||||
query = `
|
||||
DELETE FROM ` + r.Table + ` WHERE ` + primaryKeysStr + ` IN (
|
||||
` + query + `
|
||||
)`
|
||||
} else {
|
||||
// MySQL does not support the LIMIT clause in a subquery with IN
|
||||
clauses := make([]string, len(r.PrimaryKeys))
|
||||
for i, key := range r.PrimaryKeys {
|
||||
clauses[i] = r.Table + "." + key + " = A." + key
|
||||
}
|
||||
joinClause := strings.Join(clauses, " AND ")
|
||||
query = `
|
||||
DELETE ` + r.Table + ` FROM ` + r.Table + ` INNER JOIN (
|
||||
` + query + `
|
||||
) AS A ON ` + joinClause
|
||||
}
|
||||
result, err := s.GetMasterX().Exec(query, args...)
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "failed to delete "+r.Table)
|
||||
}
|
||||
rowsAffected, err = result.RowsAffected()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "failed to get rows affected for "+r.Table)
|
||||
}
|
||||
return
|
||||
}
|
||||
14
server/channels/store/sqlstore/retention_policy_store_test.go
Обычный файл
14
server/channels/store/sqlstore/retention_policy_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestRetentionPolicyStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestRetentionPolicyStore)
|
||||
}
|
||||
433
server/channels/store/sqlstore/role_store.go
Обычный файл
433
server/channels/store/sqlstore/role_store.go
Обычный файл
@@ -0,0 +1,433 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlRoleStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
type Role struct {
|
||||
Id string
|
||||
Name string
|
||||
DisplayName string
|
||||
Description string
|
||||
CreateAt int64
|
||||
UpdateAt int64
|
||||
DeleteAt int64
|
||||
Permissions string
|
||||
SchemeManaged bool
|
||||
BuiltIn bool
|
||||
}
|
||||
|
||||
type channelRolesPermissions struct {
|
||||
GuestRoleName string
|
||||
UserRoleName string
|
||||
AdminRoleName string
|
||||
HigherScopedGuestPermissions string
|
||||
HigherScopedUserPermissions string
|
||||
HigherScopedAdminPermissions string
|
||||
}
|
||||
|
||||
func NewRoleFromModel(role *model.Role) *Role {
|
||||
permissionsMap := make(map[string]bool)
|
||||
permissions := ""
|
||||
|
||||
for _, permission := range role.Permissions {
|
||||
if !permissionsMap[permission] {
|
||||
permissions += fmt.Sprintf(" %v", permission)
|
||||
permissionsMap[permission] = true
|
||||
}
|
||||
}
|
||||
|
||||
return &Role{
|
||||
Id: role.Id,
|
||||
Name: role.Name,
|
||||
DisplayName: role.DisplayName,
|
||||
Description: role.Description,
|
||||
CreateAt: role.CreateAt,
|
||||
UpdateAt: role.UpdateAt,
|
||||
DeleteAt: role.DeleteAt,
|
||||
Permissions: permissions,
|
||||
SchemeManaged: role.SchemeManaged,
|
||||
BuiltIn: role.BuiltIn,
|
||||
}
|
||||
}
|
||||
|
||||
func (role Role) ToModel() *model.Role {
|
||||
return &model.Role{
|
||||
Id: role.Id,
|
||||
Name: role.Name,
|
||||
DisplayName: role.DisplayName,
|
||||
Description: role.Description,
|
||||
CreateAt: role.CreateAt,
|
||||
UpdateAt: role.UpdateAt,
|
||||
DeleteAt: role.DeleteAt,
|
||||
Permissions: strings.Fields(role.Permissions),
|
||||
SchemeManaged: role.SchemeManaged,
|
||||
BuiltIn: role.BuiltIn,
|
||||
}
|
||||
}
|
||||
|
||||
func newSqlRoleStore(sqlStore *SqlStore) store.RoleStore {
|
||||
return &SqlRoleStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) Save(role *model.Role) (_ *model.Role, err error) {
|
||||
// Check the role is valid before proceeding.
|
||||
if !role.IsValidWithoutId() {
|
||||
return nil, store.NewErrInvalidInput("Role", "<any>", fmt.Sprintf("%v", role))
|
||||
}
|
||||
|
||||
if role.Id == "" {
|
||||
transaction, terr := s.GetMasterX().Beginx()
|
||||
if terr != nil {
|
||||
return nil, errors.Wrap(terr, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &terr)
|
||||
|
||||
createdRole, terr := s.createRole(role, transaction)
|
||||
if terr != nil {
|
||||
return nil, errors.Wrap(terr, "unable to create Role")
|
||||
} else if terr = transaction.Commit(); terr != nil {
|
||||
return nil, errors.Wrap(terr, "commit_transaction")
|
||||
}
|
||||
return createdRole, nil
|
||||
}
|
||||
|
||||
dbRole := NewRoleFromModel(role)
|
||||
dbRole.UpdateAt = model.GetMillis()
|
||||
|
||||
res, err := s.GetMasterX().NamedExec(`UPDATE Roles
|
||||
SET UpdateAt=:UpdateAt, DeleteAt=:DeleteAt, CreateAt=:CreateAt, Name=:Name, DisplayName=:DisplayName,
|
||||
Description=:Description, Permissions=:Permissions, SchemeManaged=:SchemeManaged, BuiltIn=:BuiltIn
|
||||
WHERE Id=:Id`, &dbRole)
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to update Role")
|
||||
}
|
||||
|
||||
rowsChanged, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "error while getting rows_affected")
|
||||
}
|
||||
|
||||
if rowsChanged != 1 {
|
||||
return nil, fmt.Errorf("invalid number of updated rows, expected 1 but got %d", rowsChanged)
|
||||
}
|
||||
|
||||
return dbRole.ToModel(), nil
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) createRole(role *model.Role, transaction *sqlxTxWrapper) (*model.Role, error) {
|
||||
// Check the role is valid before proceeding.
|
||||
if !role.IsValidWithoutId() {
|
||||
return nil, store.NewErrInvalidInput("Role", "<any>", fmt.Sprintf("%v", role))
|
||||
}
|
||||
|
||||
dbRole := NewRoleFromModel(role)
|
||||
|
||||
dbRole.Id = model.NewId()
|
||||
dbRole.CreateAt = model.GetMillis()
|
||||
dbRole.UpdateAt = dbRole.CreateAt
|
||||
|
||||
if _, err := transaction.NamedExec(`INSERT INTO Roles
|
||||
(Id, Name, DisplayName, Description, Permissions, CreateAt, UpdateAt, DeleteAt, SchemeManaged, BuiltIn)
|
||||
VALUES
|
||||
(:Id, :Name, :DisplayName, :Description, :Permissions, :CreateAt, :UpdateAt, :DeleteAt, :SchemeManaged, :BuiltIn)`, dbRole); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save Role")
|
||||
}
|
||||
|
||||
return dbRole.ToModel(), nil
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) Get(roleId string) (*model.Role, error) {
|
||||
dbRole := Role{}
|
||||
|
||||
if err := s.GetReplicaX().Get(&dbRole, "SELECT * from Roles WHERE Id = ?", roleId); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Role", roleId)
|
||||
}
|
||||
return nil, errors.Wrap(err, "failed to get Role")
|
||||
}
|
||||
|
||||
return dbRole.ToModel(), nil
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) GetAll() ([]*model.Role, error) {
|
||||
dbRoles := []Role{}
|
||||
|
||||
if err := s.GetReplicaX().Select(&dbRoles, "SELECT * from Roles"); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Roles")
|
||||
}
|
||||
|
||||
roles := []*model.Role{}
|
||||
for _, dbRole := range dbRoles {
|
||||
roles = append(roles, dbRole.ToModel())
|
||||
}
|
||||
return roles, nil
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) GetByName(ctx context.Context, name string) (*model.Role, error) {
|
||||
dbRole := Role{}
|
||||
if err := s.DBXFromContext(ctx).Get(&dbRole, "SELECT * from Roles WHERE Name = ?", name); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Role", fmt.Sprintf("name=%s", name))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to find Roles with name=%s", name)
|
||||
}
|
||||
|
||||
return dbRole.ToModel(), nil
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) GetByNames(names []string) ([]*model.Role, error) {
|
||||
if len(names) == 0 {
|
||||
return []*model.Role{}, nil
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("Id, Name, DisplayName, Description, CreateAt, UpdateAt, DeleteAt, Permissions, SchemeManaged, BuiltIn").
|
||||
From("Roles").
|
||||
Where(sq.Eq{"Name": names})
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "role_tosql")
|
||||
}
|
||||
|
||||
rows, err := s.GetReplicaX().DB.Query(queryString, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Roles")
|
||||
}
|
||||
|
||||
roles := []*model.Role{}
|
||||
defer rows.Close()
|
||||
for rows.Next() {
|
||||
var role Role
|
||||
err = rows.Scan(
|
||||
&role.Id, &role.Name, &role.DisplayName, &role.Description,
|
||||
&role.CreateAt, &role.UpdateAt, &role.DeleteAt, &role.Permissions,
|
||||
&role.SchemeManaged, &role.BuiltIn)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to scan values")
|
||||
}
|
||||
roles = append(roles, role.ToModel())
|
||||
}
|
||||
if err = rows.Err(); err != nil {
|
||||
return nil, errors.Wrap(err, "unable to iterate over rows")
|
||||
}
|
||||
|
||||
return roles, nil
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) Delete(roleId string) (*model.Role, error) {
|
||||
// Get the role.
|
||||
var role Role
|
||||
if err := s.GetReplicaX().Get(&role, "SELECT * from Roles WHERE Id = ?", roleId); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Role", roleId)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get Role with id=%s", roleId)
|
||||
}
|
||||
|
||||
time := model.GetMillis()
|
||||
role.DeleteAt = time
|
||||
role.UpdateAt = time
|
||||
|
||||
res, err := s.GetMasterX().NamedExec(`UPDATE Roles
|
||||
SET UpdateAt=:UpdateAt, DeleteAt=:DeleteAt, CreateAt=:CreateAt, Name=:Name, DisplayName=:DisplayName,
|
||||
Description=:Description, Permissions=:Permissions, SchemeManaged=:SchemeManaged, BuiltIn=:BuiltIn
|
||||
WHERE Id=:Id`, &role)
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to update Role")
|
||||
}
|
||||
|
||||
rowsChanged, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "error while getting rows_affected")
|
||||
}
|
||||
|
||||
if rowsChanged != 1 {
|
||||
return nil, fmt.Errorf("invalid number of updated rows, expected 1 but got %d", rowsChanged)
|
||||
}
|
||||
|
||||
return role.ToModel(), nil
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) PermanentDeleteAll() error {
|
||||
if _, err := s.GetMasterX().Exec("DELETE FROM Roles"); err != nil {
|
||||
return errors.Wrap(err, "failed to delete Roles")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) channelHigherScopedPermissionsQuery(roleNames []string) string {
|
||||
sqlTmpl := `
|
||||
SELECT
|
||||
'' AS GuestRoleName,
|
||||
RoleSchemes.DefaultChannelUserRole AS UserRoleName,
|
||||
RoleSchemes.DefaultChannelAdminRole AS AdminRoleName,
|
||||
'' AS HigherScopedGuestPermissions,
|
||||
UserRoles.Permissions AS HigherScopedUserPermissions,
|
||||
AdminRoles.Permissions AS HigherScopedAdminPermissions
|
||||
FROM
|
||||
Schemes AS RoleSchemes
|
||||
JOIN Channels ON Channels.SchemeId = RoleSchemes.Id
|
||||
JOIN Teams ON Teams.Id = Channels.TeamId
|
||||
JOIN Schemes ON Schemes.Id = Teams.SchemeId
|
||||
RIGHT JOIN Roles AS UserRoles ON UserRoles.Name = Schemes.DefaultChannelUserRole
|
||||
RIGHT JOIN Roles AS AdminRoles ON AdminRoles.Name = Schemes.DefaultChannelAdminRole
|
||||
WHERE
|
||||
RoleSchemes.DefaultChannelUserRole IN ('%[1]s')
|
||||
OR RoleSchemes.DefaultChannelAdminRole IN ('%[1]s')
|
||||
|
||||
UNION
|
||||
|
||||
SELECT
|
||||
RoleSchemes.DefaultChannelGuestRole AS GuestRoleName,
|
||||
'' AS UserRoleName,
|
||||
'' AS AdminRoleName,
|
||||
GuestRoles.Permissions AS HigherScopedGuestPermissions,
|
||||
'' AS HigherScopedUserPermissions,
|
||||
'' AS HigherScopedAdminPermissions
|
||||
FROM
|
||||
Schemes AS RoleSchemes
|
||||
JOIN Channels ON Channels.SchemeId = RoleSchemes.Id
|
||||
JOIN Teams ON Teams.Id = Channels.TeamId
|
||||
JOIN Schemes ON Schemes.Id = Teams.SchemeId
|
||||
RIGHT JOIN Roles AS GuestRoles ON GuestRoles.Name = Schemes.DefaultChannelGuestRole
|
||||
WHERE
|
||||
RoleSchemes.DefaultChannelGuestRole IN ('%[1]s')
|
||||
|
||||
UNION
|
||||
|
||||
SELECT
|
||||
Schemes.DefaultChannelGuestRole AS GuestRoleName,
|
||||
Schemes.DefaultChannelUserRole AS UserRoleName,
|
||||
Schemes.DefaultChannelAdminRole AS AdminRoleName,
|
||||
GuestRoles.Permissions AS HigherScopedGuestPermissions,
|
||||
UserRoles.Permissions AS HigherScopedUserPermissions,
|
||||
AdminRoles.Permissions AS HigherScopedAdminPermissions
|
||||
FROM
|
||||
Schemes
|
||||
JOIN Channels ON Channels.SchemeId = Schemes.Id
|
||||
JOIN Teams ON Teams.Id = Channels.TeamId
|
||||
JOIN Roles AS GuestRoles ON GuestRoles.Name = '%[2]s'
|
||||
JOIN Roles AS UserRoles ON UserRoles.Name = '%[3]s'
|
||||
JOIN Roles AS AdminRoles ON AdminRoles.Name = '%[4]s'
|
||||
WHERE
|
||||
(Schemes.DefaultChannelGuestRole IN ('%[1]s')
|
||||
OR Schemes.DefaultChannelUserRole IN ('%[1]s')
|
||||
OR Schemes.DefaultChannelAdminRole IN ('%[1]s'))
|
||||
AND (Teams.SchemeId = ''
|
||||
OR Teams.SchemeId IS NULL)
|
||||
`
|
||||
|
||||
// The below three channel role names are referenced by their name value because there is no system scheme
|
||||
// record that ships with Mattermost, otherwise the system scheme would be referenced by name and the channel
|
||||
// roles would be referenced by their column names.
|
||||
return fmt.Sprintf(
|
||||
sqlTmpl,
|
||||
strings.Join(roleNames, "', '"),
|
||||
model.ChannelGuestRoleId,
|
||||
model.ChannelUserRoleId,
|
||||
model.ChannelAdminRoleId,
|
||||
)
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) ChannelHigherScopedPermissions(roleNames []string) (map[string]*model.RolePermissions, error) {
|
||||
query := s.channelHigherScopedPermissionsQuery(roleNames)
|
||||
|
||||
rolesPermissions := []*channelRolesPermissions{}
|
||||
if err := s.GetReplicaX().Select(&rolesPermissions, query); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find RolePermissions")
|
||||
}
|
||||
|
||||
roleNameHigherScopedPermissions := map[string]*model.RolePermissions{}
|
||||
|
||||
for _, rp := range rolesPermissions {
|
||||
roleNameHigherScopedPermissions[rp.GuestRoleName] = &model.RolePermissions{RoleID: model.ChannelGuestRoleId, Permissions: strings.Split(rp.HigherScopedGuestPermissions, " ")}
|
||||
roleNameHigherScopedPermissions[rp.UserRoleName] = &model.RolePermissions{RoleID: model.ChannelUserRoleId, Permissions: strings.Split(rp.HigherScopedUserPermissions, " ")}
|
||||
roleNameHigherScopedPermissions[rp.AdminRoleName] = &model.RolePermissions{RoleID: model.ChannelAdminRoleId, Permissions: strings.Split(rp.HigherScopedAdminPermissions, " ")}
|
||||
}
|
||||
|
||||
return roleNameHigherScopedPermissions, nil
|
||||
}
|
||||
|
||||
func (s *SqlRoleStore) AllChannelSchemeRoles() ([]*model.Role, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("Roles.*").
|
||||
From("Schemes").
|
||||
Join("Roles ON Schemes.DefaultChannelGuestRole = Roles.Name OR Schemes.DefaultChannelUserRole = Roles.Name OR Schemes.DefaultChannelAdminRole = Roles.Name").
|
||||
Where(sq.Eq{"Schemes.Scope": model.SchemeScopeChannel}).
|
||||
Where(sq.Eq{"Roles.DeleteAt": 0}).
|
||||
Where(sq.Eq{"Schemes.DeleteAt": 0})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "role_tosql")
|
||||
}
|
||||
|
||||
dbRoles := []*Role{}
|
||||
if err = s.GetReplicaX().Select(&dbRoles, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Roles")
|
||||
}
|
||||
|
||||
roles := []*model.Role{}
|
||||
for _, dbRole := range dbRoles {
|
||||
roles = append(roles, dbRole.ToModel())
|
||||
}
|
||||
|
||||
return roles, nil
|
||||
}
|
||||
|
||||
// ChannelRolesUnderTeamRole finds all of the channel-scheme roles under the team of the given team-scheme role.
|
||||
func (s *SqlRoleStore) ChannelRolesUnderTeamRole(roleName string) ([]*model.Role, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("ChannelSchemeRoles.*").
|
||||
From("Roles AS HigherScopedRoles").
|
||||
Join("Schemes AS HigherScopedSchemes ON (HigherScopedRoles.Name = HigherScopedSchemes.DefaultChannelGuestRole OR HigherScopedRoles.Name = HigherScopedSchemes.DefaultChannelUserRole OR HigherScopedRoles.Name = HigherScopedSchemes.DefaultChannelAdminRole)").
|
||||
Join("Teams ON Teams.SchemeId = HigherScopedSchemes.Id").
|
||||
Join("Channels ON Channels.TeamId = Teams.Id").
|
||||
Join("Schemes AS ChannelSchemes ON Channels.SchemeId = ChannelSchemes.Id").
|
||||
Join("Roles AS ChannelSchemeRoles ON (ChannelSchemeRoles.Name = ChannelSchemes.DefaultChannelGuestRole OR ChannelSchemeRoles.Name = ChannelSchemes.DefaultChannelUserRole OR ChannelSchemeRoles.Name = ChannelSchemes.DefaultChannelAdminRole)").
|
||||
Where(sq.Eq{"HigherScopedSchemes.Scope": model.SchemeScopeTeam}).
|
||||
Where(sq.Eq{"HigherScopedRoles.Name": roleName}).
|
||||
Where(sq.Eq{"HigherScopedRoles.DeleteAt": 0}).
|
||||
Where(sq.Eq{"HigherScopedSchemes.DeleteAt": 0}).
|
||||
Where(sq.Eq{"Teams.DeleteAt": 0}).
|
||||
Where(sq.Eq{"Channels.DeleteAt": 0}).
|
||||
Where(sq.Eq{"ChannelSchemes.DeleteAt": 0}).
|
||||
Where(sq.Eq{"ChannelSchemeRoles.DeleteAt": 0})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "role_tosql")
|
||||
}
|
||||
|
||||
dbRoles := []*Role{}
|
||||
if err = s.GetReplicaX().Select(&dbRoles, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Roles")
|
||||
}
|
||||
|
||||
roles := []*model.Role{}
|
||||
for _, dbRole := range dbRoles {
|
||||
roles = append(roles, dbRole.ToModel())
|
||||
}
|
||||
|
||||
return roles, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/role_store_test.go
Обычный файл
14
server/channels/store/sqlstore/role_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestRoleStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestRoleStore)
|
||||
}
|
||||
470
server/channels/store/sqlstore/scheme_store.go
Обычный файл
470
server/channels/store/sqlstore/scheme_store.go
Обычный файл
@@ -0,0 +1,470 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
const (
|
||||
SchemeRoleDisplayNameTeamAdmin = "Team Admin Role for Scheme"
|
||||
SchemeRoleDisplayNameTeamUser = "Team User Role for Scheme"
|
||||
SchemeRoleDisplayNameTeamGuest = "Team Guest Role for Scheme"
|
||||
|
||||
SchemeRoleDisplayNameChannelAdmin = "Channel Admin Role for Scheme"
|
||||
SchemeRoleDisplayNameChannelUser = "Channel User Role for Scheme"
|
||||
SchemeRoleDisplayNameChannelGuest = "Channel Guest Role for Scheme"
|
||||
|
||||
SchemeRoleDisplayNamePlaybookAdmin = "Playbook Admin Role for Scheme"
|
||||
SchemeRoleDisplayNamePlaybookMember = "Playbook Member Role for Scheme"
|
||||
|
||||
SchemeRoleDisplayNameRunAdmin = "Run Admin Role for Scheme"
|
||||
SchemeRoleDisplayNameRunMember = "Run Member Role for Scheme"
|
||||
)
|
||||
|
||||
type SqlSchemeStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlSchemeStore(sqlStore *SqlStore) store.SchemeStore {
|
||||
return &SqlSchemeStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s *SqlSchemeStore) Save(scheme *model.Scheme) (_ *model.Scheme, err error) {
|
||||
if scheme.Id == "" {
|
||||
transaction, terr := s.GetMasterX().Beginx()
|
||||
if terr != nil {
|
||||
return nil, errors.Wrap(terr, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &terr)
|
||||
|
||||
newScheme, terr := s.createScheme(scheme, transaction)
|
||||
if terr != nil {
|
||||
return nil, terr
|
||||
}
|
||||
if terr = transaction.Commit(); terr != nil {
|
||||
return nil, errors.Wrap(terr, "commit_transaction")
|
||||
}
|
||||
return newScheme, nil
|
||||
}
|
||||
|
||||
if !scheme.IsValid() {
|
||||
return nil, store.NewErrInvalidInput("Scheme", "<any>", fmt.Sprintf("%v", scheme))
|
||||
}
|
||||
|
||||
scheme.UpdateAt = model.GetMillis()
|
||||
|
||||
res, err := s.GetMasterX().NamedExec(`UPDATE Schemes
|
||||
SET UpdateAt=:UpdateAt, CreateAt=:CreateAt, DeleteAt=:DeleteAt, Name=:Name, DisplayName=:DisplayName, Description=:Description, Scope=:Scope,
|
||||
DefaultTeamAdminRole=:DefaultTeamAdminRole, DefaultTeamUserRole=:DefaultTeamUserRole, DefaultTeamGuestRole=:DefaultTeamGuestRole,
|
||||
DefaultChannelAdminRole=:DefaultChannelAdminRole, DefaultChannelUserRole=:DefaultChannelUserRole, DefaultChannelGuestRole=:DefaultChannelGuestRole,
|
||||
DefaultPlaybookMemberRole=:DefaultPlaybookMemberRole, DefaultPlaybookAdminRole=:DefaultPlaybookAdminRole, DefaultRunMemberRole=:DefaultRunMemberRole, DefaultRunAdminRole=:DefaultRunAdminRole
|
||||
WHERE Id=:Id`, scheme)
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to update Scheme")
|
||||
}
|
||||
|
||||
rowsChanged, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "error while getting rows_affected")
|
||||
}
|
||||
if rowsChanged != 1 {
|
||||
return nil, errors.New("no record to update")
|
||||
}
|
||||
|
||||
return scheme, nil
|
||||
}
|
||||
|
||||
func (s *SqlSchemeStore) createScheme(scheme *model.Scheme, transaction *sqlxTxWrapper) (*model.Scheme, error) {
|
||||
// Fetch the default system scheme roles to populate default permissions.
|
||||
defaultRoleNames := []string{
|
||||
model.TeamAdminRoleId,
|
||||
model.TeamUserRoleId,
|
||||
model.TeamGuestRoleId,
|
||||
model.ChannelAdminRoleId,
|
||||
model.ChannelUserRoleId,
|
||||
model.ChannelGuestRoleId,
|
||||
model.PlaybookAdminRoleId,
|
||||
model.PlaybookMemberRoleId,
|
||||
model.RunAdminRoleId,
|
||||
model.RunMemberRoleId,
|
||||
}
|
||||
defaultRoles := make(map[string]*model.Role)
|
||||
roles, err := s.SqlStore.Role().GetByNames(defaultRoleNames)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, role := range roles {
|
||||
defaultRoles[role.Name] = role
|
||||
}
|
||||
|
||||
if len(defaultRoles) != len(defaultRoleNames) {
|
||||
return nil, errors.New("createScheme: unable to retrieve default scheme roles")
|
||||
}
|
||||
|
||||
// Create the appropriate default roles for the scheme.
|
||||
if scheme.Scope == model.SchemeScopeTeam {
|
||||
// Team Admin Role
|
||||
teamAdminRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("%s %s", SchemeRoleDisplayNameTeamAdmin, scheme.Name),
|
||||
Permissions: defaultRoles[model.TeamAdminRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
|
||||
savedRole, err := s.SqlStore.Role().(*SqlRoleStore).createRole(teamAdminRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultTeamAdminRole = savedRole.Name
|
||||
|
||||
// Team User Role
|
||||
teamUserRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("%s %s", SchemeRoleDisplayNameTeamUser, scheme.Name),
|
||||
Permissions: defaultRoles[model.TeamUserRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
|
||||
savedRole, err = s.SqlStore.Role().(*SqlRoleStore).createRole(teamUserRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultTeamUserRole = savedRole.Name
|
||||
|
||||
// Team Guest Role
|
||||
teamGuestRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("%s %s", SchemeRoleDisplayNameTeamGuest, scheme.Name),
|
||||
Permissions: defaultRoles[model.TeamGuestRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
|
||||
savedRole, err = s.SqlStore.Role().(*SqlRoleStore).createRole(teamGuestRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultTeamGuestRole = savedRole.Name
|
||||
|
||||
// playbook admin role
|
||||
playbookAdminRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("%s %s", SchemeRoleDisplayNamePlaybookAdmin, scheme.Name),
|
||||
Permissions: defaultRoles[model.PlaybookAdminRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
savedRole, err = s.SqlStore.Role().(*SqlRoleStore).createRole(playbookAdminRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultPlaybookAdminRole = savedRole.Name
|
||||
|
||||
// playbook member role
|
||||
playbookMemberRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("%s %s", SchemeRoleDisplayNamePlaybookMember, scheme.Name),
|
||||
Permissions: defaultRoles[model.PlaybookMemberRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
savedRole, err = s.SqlStore.Role().(*SqlRoleStore).createRole(playbookMemberRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultPlaybookMemberRole = savedRole.Name
|
||||
|
||||
// run admin role
|
||||
runAdminRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("%s %s", SchemeRoleDisplayNameRunAdmin, scheme.Name),
|
||||
Permissions: defaultRoles[model.RunAdminRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
savedRole, err = s.SqlStore.Role().(*SqlRoleStore).createRole(runAdminRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultRunAdminRole = savedRole.Name
|
||||
|
||||
// run member role
|
||||
runMemberRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("%s %s", SchemeRoleDisplayNameRunMember, scheme.Name),
|
||||
Permissions: defaultRoles[model.RunMemberRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
savedRole, err = s.SqlStore.Role().(*SqlRoleStore).createRole(runMemberRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultRunMemberRole = savedRole.Name
|
||||
}
|
||||
|
||||
if scheme.Scope == model.SchemeScopeTeam || scheme.Scope == model.SchemeScopeChannel {
|
||||
// Channel Admin Role
|
||||
channelAdminRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("Channel Admin Role for Scheme %s", scheme.Name),
|
||||
Permissions: defaultRoles[model.ChannelAdminRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
|
||||
if scheme.Scope == model.SchemeScopeChannel {
|
||||
channelAdminRole.Permissions = []string{}
|
||||
}
|
||||
|
||||
savedRole, err := s.SqlStore.Role().(*SqlRoleStore).createRole(channelAdminRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultChannelAdminRole = savedRole.Name
|
||||
|
||||
// Channel User Role
|
||||
channelUserRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("Channel User Role for Scheme %s", scheme.Name),
|
||||
Permissions: defaultRoles[model.ChannelUserRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
|
||||
if scheme.Scope == model.SchemeScopeChannel {
|
||||
channelUserRole.Permissions = filterModerated(channelUserRole.Permissions)
|
||||
}
|
||||
|
||||
savedRole, err = s.SqlStore.Role().(*SqlRoleStore).createRole(channelUserRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultChannelUserRole = savedRole.Name
|
||||
|
||||
// Channel Guest Role
|
||||
channelGuestRole := &model.Role{
|
||||
Name: model.NewId(),
|
||||
DisplayName: fmt.Sprintf("Channel Guest Role for Scheme %s", scheme.Name),
|
||||
Permissions: defaultRoles[model.ChannelGuestRoleId].Permissions,
|
||||
SchemeManaged: true,
|
||||
}
|
||||
|
||||
if scheme.Scope == model.SchemeScopeChannel {
|
||||
channelGuestRole.Permissions = filterModerated(channelGuestRole.Permissions)
|
||||
}
|
||||
|
||||
savedRole, err = s.SqlStore.Role().(*SqlRoleStore).createRole(channelGuestRole, transaction)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
scheme.DefaultChannelGuestRole = savedRole.Name
|
||||
}
|
||||
|
||||
scheme.Id = model.NewId()
|
||||
if scheme.Name == "" {
|
||||
scheme.Name = model.NewId()
|
||||
}
|
||||
scheme.CreateAt = model.GetMillis()
|
||||
scheme.UpdateAt = scheme.CreateAt
|
||||
|
||||
// Validate the scheme
|
||||
if !scheme.IsValidForCreate() {
|
||||
return nil, store.NewErrInvalidInput("Scheme", "<any>", fmt.Sprintf("%v", scheme))
|
||||
}
|
||||
|
||||
if _, err := transaction.NamedExec(`INSERT INTO Schemes
|
||||
(Id, Name, DisplayName, Description, Scope, DefaultTeamAdminRole, DefaultTeamUserRole, DefaultTeamGuestRole, DefaultChannelAdminRole, DefaultChannelUserRole, DefaultChannelGuestRole, CreateAt, UpdateAt, DeleteAt, DefaultPlaybookAdminRole, DefaultPlaybookMemberRole, DefaultRunAdminRole, DefaultRunMemberRole)
|
||||
VALUES
|
||||
(:Id, :Name, :DisplayName, :Description, :Scope, :DefaultTeamAdminRole, :DefaultTeamUserRole, :DefaultTeamGuestRole, :DefaultChannelAdminRole, :DefaultChannelUserRole, :DefaultChannelGuestRole, :CreateAt, :UpdateAt, :DeleteAt, :DefaultPlaybookAdminRole, :DefaultPlaybookMemberRole, :DefaultRunAdminRole, :DefaultRunMemberRole)`, scheme); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save Scheme")
|
||||
}
|
||||
|
||||
return scheme, nil
|
||||
}
|
||||
|
||||
func filterModerated(permissions []string) []string {
|
||||
filteredPermissions := []string{}
|
||||
for _, perm := range permissions {
|
||||
if _, ok := model.ChannelModeratedPermissionsMap[perm]; ok {
|
||||
filteredPermissions = append(filteredPermissions, perm)
|
||||
}
|
||||
}
|
||||
return filteredPermissions
|
||||
}
|
||||
|
||||
func (s *SqlSchemeStore) Get(schemeId string) (*model.Scheme, error) {
|
||||
var scheme model.Scheme
|
||||
if err := s.GetReplicaX().Get(&scheme, "SELECT * from Schemes WHERE Id = ?", schemeId); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Scheme", fmt.Sprintf("schemeId=%s", schemeId))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get Scheme with schemeId=%s", schemeId)
|
||||
}
|
||||
|
||||
return &scheme, nil
|
||||
}
|
||||
|
||||
func (s *SqlSchemeStore) GetByName(schemeName string) (*model.Scheme, error) {
|
||||
var scheme model.Scheme
|
||||
|
||||
if err := s.GetReplicaX().Get(&scheme, "SELECT * from Schemes WHERE Name = ?", schemeName); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Scheme", fmt.Sprintf("schemeName=%s", schemeName))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get Scheme with schemeName=%s", schemeName)
|
||||
}
|
||||
|
||||
return &scheme, nil
|
||||
}
|
||||
|
||||
func (s *SqlSchemeStore) Delete(schemeId string) (*model.Scheme, error) {
|
||||
// Get the scheme
|
||||
scheme := model.Scheme{}
|
||||
if err := s.GetMasterX().Get(&scheme, `SELECT * from Schemes WHERE Id = ?`, schemeId); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Scheme", fmt.Sprintf("schemeId=%s", schemeId))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get Scheme with schemeId=%s", schemeId)
|
||||
}
|
||||
|
||||
// Update any teams or channels using this scheme to the default scheme.
|
||||
if scheme.Scope == model.SchemeScopeTeam {
|
||||
if _, err := s.GetMasterX().Exec(`UPDATE Teams SET SchemeId = '' WHERE SchemeId = ?`, schemeId); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update Teams with schemeId=%s", schemeId)
|
||||
}
|
||||
|
||||
s.Team().ClearCaches()
|
||||
} else if scheme.Scope == model.SchemeScopeChannel {
|
||||
if _, err := s.GetMasterX().Exec(`UPDATE Channels SET SchemeId = '' WHERE SchemeId = ?`, schemeId); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update Channels with schemeId=%s", schemeId)
|
||||
}
|
||||
}
|
||||
|
||||
// Blow away the channel caches.
|
||||
s.Channel().ClearCaches()
|
||||
|
||||
// Delete the roles belonging to the scheme.
|
||||
roleNames := []string{scheme.DefaultChannelGuestRole, scheme.DefaultChannelUserRole, scheme.DefaultChannelAdminRole}
|
||||
if scheme.Scope == model.SchemeScopeTeam {
|
||||
roleNames = append(roleNames, scheme.DefaultTeamGuestRole, scheme.DefaultTeamUserRole, scheme.DefaultTeamAdminRole)
|
||||
}
|
||||
if scheme.Scope == model.SchemeScopePlaybook {
|
||||
roleNames = append(roleNames, scheme.DefaultPlaybookAdminRole, scheme.DefaultPlaybookMemberRole)
|
||||
}
|
||||
|
||||
if scheme.Scope == model.SchemeScopeRun {
|
||||
roleNames = append(roleNames, scheme.DefaultRunAdminRole, scheme.DefaultRunMemberRole)
|
||||
}
|
||||
|
||||
time := model.GetMillis()
|
||||
|
||||
updateQuery, args, err := s.getQueryBuilder().
|
||||
Update("Roles").
|
||||
Where(sq.Eq{"Name": roleNames}).
|
||||
Set("UpdateAt", time).
|
||||
Set("DeleteAt", time).
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "status_tosql")
|
||||
}
|
||||
|
||||
if _, err = s.GetMasterX().Exec(updateQuery, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update Roles with name in (%s)", roleNames)
|
||||
}
|
||||
|
||||
// Delete the scheme itself.
|
||||
scheme.UpdateAt = time
|
||||
scheme.DeleteAt = time
|
||||
|
||||
res, err := s.GetMasterX().NamedExec(`UPDATE Schemes
|
||||
SET UpdateAt=:UpdateAt, DeleteAt=:DeleteAt, CreateAt=:CreateAt, Name=:Name, DisplayName=:DisplayName, Description=:Description, Scope=:Scope,
|
||||
DefaultTeamAdminRole=:DefaultTeamAdminRole, DefaultTeamUserRole=:DefaultTeamUserRole, DefaultTeamGuestRole=:DefaultTeamGuestRole,
|
||||
DefaultChannelAdminRole=:DefaultChannelAdminRole, DefaultChannelUserRole=:DefaultChannelUserRole, DefaultChannelGuestRole=:DefaultChannelGuestRole
|
||||
WHERE Id=:Id`, &scheme)
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update Scheme with schemeId=%s", schemeId)
|
||||
}
|
||||
|
||||
rowsChanged, err := res.RowsAffected()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get RowsAffected while updating scheme with schemeId=%s", schemeId)
|
||||
}
|
||||
if rowsChanged != 1 {
|
||||
return nil, errors.New("no record to update")
|
||||
}
|
||||
return &scheme, nil
|
||||
}
|
||||
|
||||
func (s *SqlSchemeStore) GetAllPage(scope string, offset int, limit int) ([]*model.Scheme, error) {
|
||||
schemes := []*model.Scheme{}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Schemes").
|
||||
Where(sq.Eq{"DeleteAt": 0}).
|
||||
OrderBy("CreateAt DESC").
|
||||
Limit(uint64(limit)).
|
||||
Offset(uint64(offset))
|
||||
|
||||
if scope != "" {
|
||||
query = query.Where(sq.Eq{"Scope": scope})
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "status_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&schemes, queryString, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get Schemes")
|
||||
}
|
||||
|
||||
return schemes, nil
|
||||
}
|
||||
|
||||
func (s *SqlSchemeStore) PermanentDeleteAll() error {
|
||||
if _, err := s.GetMasterX().Exec("DELETE from Schemes"); err != nil {
|
||||
return errors.Wrap(err, "failed to delete Schemes")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SqlSchemeStore) CountByScope(scope string) (int64, error) {
|
||||
var count int64
|
||||
err := s.GetReplicaX().Get(&count, `SELECT count(*) FROM Schemes WHERE Scope = ? AND DeleteAt = 0`, scope)
|
||||
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "failed to count Schemes by scope")
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s *SqlSchemeStore) CountWithoutPermission(schemeScope, permissionID string, roleScope model.RoleScope, roleType model.RoleType) (int64, error) {
|
||||
joinCol := fmt.Sprintf("Default%s%sRole", roleScope, roleType)
|
||||
query := fmt.Sprintf(`
|
||||
SELECT
|
||||
count(*)
|
||||
FROM Schemes
|
||||
JOIN Roles ON Roles.Name = Schemes.%s
|
||||
WHERE
|
||||
Schemes.DeleteAt = 0 AND
|
||||
Schemes.Scope = '%s' AND
|
||||
Roles.Permissions NOT LIKE '%%%s%%'
|
||||
`, joinCol, schemeScope, permissionID)
|
||||
|
||||
var count int64
|
||||
err := s.GetReplicaX().Get(&count, query)
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "failed to count Schemes without permission")
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/scheme_store_test.go
Обычный файл
14
server/channels/store/sqlstore/scheme_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestSchemeStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestSchemeStore)
|
||||
}
|
||||
327
server/channels/store/sqlstore/session_store.go
Обычный файл
327
server/channels/store/sqlstore/session_store.go
Обычный файл
@@ -0,0 +1,327 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
const (
|
||||
sessionsCleanupDelay = 100 * time.Millisecond
|
||||
)
|
||||
|
||||
type SqlSessionStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlSessionStore(sqlStore *SqlStore) store.SessionStore {
|
||||
return &SqlSessionStore{sqlStore}
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) Save(session *model.Session) (*model.Session, error) {
|
||||
if session.Id != "" {
|
||||
return nil, store.NewErrInvalidInput("Session", "id", session.Id)
|
||||
}
|
||||
|
||||
session.PreSave()
|
||||
|
||||
if err := session.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
jsonProps, err := json.Marshal(session.Props)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed marshalling session props")
|
||||
}
|
||||
|
||||
if me.IsBinaryParamEnabled() {
|
||||
jsonProps = AppendBinaryFlag(jsonProps)
|
||||
}
|
||||
|
||||
query, args, err := me.getQueryBuilder().
|
||||
Insert("Sessions").
|
||||
Columns("Id", "Token", "CreateAt", "ExpiresAt", "LastActivityAt", "UserId", "DeviceId", "Roles", "IsOAuth", "ExpiredNotify", "Props").
|
||||
Values(session.Id, session.Token, session.CreateAt, session.ExpiresAt, session.LastActivityAt, session.UserId, session.DeviceId, session.Roles, session.IsOAuth, session.ExpiredNotify, jsonProps).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "sessions_tosql")
|
||||
}
|
||||
if _, err = me.GetMasterX().Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to save Session with id=%s", session.Id)
|
||||
}
|
||||
|
||||
teamMembers, err := me.Team().GetTeamsForUser(context.Background(), session.UserId, "", true)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", session.UserId)
|
||||
}
|
||||
|
||||
session.TeamMembers = make([]*model.TeamMember, 0, len(teamMembers))
|
||||
for _, tm := range teamMembers {
|
||||
if tm.DeleteAt == 0 {
|
||||
session.TeamMembers = append(session.TeamMembers, tm)
|
||||
}
|
||||
}
|
||||
|
||||
return session, nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) Get(ctx context.Context, sessionIdOrToken string) (*model.Session, error) {
|
||||
sessions := []*model.Session{}
|
||||
|
||||
if err := me.DBXFromContext(ctx).Select(&sessions, "SELECT * FROM Sessions WHERE Token = ? OR Id = ? LIMIT 1", sessionIdOrToken, sessionIdOrToken); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Sessions with sessionIdOrToken=%s", sessionIdOrToken)
|
||||
}
|
||||
if len(sessions) == 0 {
|
||||
return nil, store.NewErrNotFound("Session", fmt.Sprintf("sessionIdOrToken=%s", sessionIdOrToken))
|
||||
}
|
||||
session := sessions[0]
|
||||
|
||||
tempMembers, err := me.Team().GetTeamsForUser(
|
||||
WithMaster(context.Background()),
|
||||
session.UserId, "", true)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", session.UserId)
|
||||
}
|
||||
sessions[0].TeamMembers = make([]*model.TeamMember, 0, len(tempMembers))
|
||||
for _, tm := range tempMembers {
|
||||
if tm.DeleteAt == 0 {
|
||||
sessions[0].TeamMembers = append(sessions[0].TeamMembers, tm)
|
||||
}
|
||||
}
|
||||
return session, nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) GetSessions(userId string) ([]*model.Session, error) {
|
||||
sessions := []*model.Session{}
|
||||
|
||||
if err := me.GetReplicaX().Select(&sessions, "SELECT * FROM Sessions WHERE UserId = ? ORDER BY LastActivityAt DESC", userId); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Sessions with userId=%s", userId)
|
||||
}
|
||||
|
||||
teamMembers, err := me.Team().GetTeamsForUser(context.Background(), userId, "", true)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", userId)
|
||||
}
|
||||
|
||||
for _, session := range sessions {
|
||||
session.TeamMembers = make([]*model.TeamMember, 0, len(teamMembers))
|
||||
for _, tm := range teamMembers {
|
||||
if tm.DeleteAt == 0 {
|
||||
session.TeamMembers = append(session.TeamMembers, tm)
|
||||
}
|
||||
}
|
||||
}
|
||||
return sessions, nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) GetSessionsWithActiveDeviceIds(userId string) ([]*model.Session, error) {
|
||||
query :=
|
||||
`SELECT *
|
||||
FROM
|
||||
Sessions
|
||||
WHERE
|
||||
UserId = ? AND
|
||||
ExpiresAt != 0 AND
|
||||
? <= ExpiresAt AND
|
||||
DeviceId != ''`
|
||||
|
||||
sessions := []*model.Session{}
|
||||
|
||||
if err := me.GetReplicaX().Select(&sessions, query, userId, model.GetMillis()); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find Sessions with userId=%s", userId)
|
||||
}
|
||||
return sessions, nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) GetSessionsExpired(thresholdMillis int64, mobileOnly bool, unnotifiedOnly bool) ([]*model.Session, error) {
|
||||
now := model.GetMillis()
|
||||
builder := me.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Sessions").
|
||||
Where(sq.NotEq{"ExpiresAt": 0}).
|
||||
Where(sq.Lt{"ExpiresAt": now}).
|
||||
Where(sq.Gt{"ExpiresAt": now - thresholdMillis})
|
||||
if mobileOnly {
|
||||
builder = builder.Where(sq.NotEq{"DeviceId": ""})
|
||||
}
|
||||
if unnotifiedOnly {
|
||||
builder = builder.Where(sq.NotEq{"ExpiredNotify": true})
|
||||
}
|
||||
|
||||
query, args, err := builder.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "sessions_tosql")
|
||||
}
|
||||
|
||||
sessions := []*model.Session{}
|
||||
|
||||
err = me.GetReplicaX().Select(&sessions, query, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Sessions")
|
||||
}
|
||||
return sessions, nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) UpdateExpiredNotify(sessionId string, notified bool) error {
|
||||
query, args, err := me.getQueryBuilder().
|
||||
Update("Sessions").
|
||||
Set("ExpiredNotify", notified).
|
||||
Where(sq.Eq{"Id": sessionId}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "sessions_tosql")
|
||||
}
|
||||
|
||||
_, err = me.GetMasterX().Exec(query, args...)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to update Session with id=%s", sessionId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) Remove(sessionIdOrToken string) error {
|
||||
_, err := me.GetMasterX().Exec("DELETE FROM Sessions WHERE Id = ? Or Token = ?", sessionIdOrToken, sessionIdOrToken)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Session with sessionIdOrToken=%s", sessionIdOrToken)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) RemoveAllSessions() error {
|
||||
_, err := me.GetMasterX().Exec("DELETE FROM Sessions")
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to delete all Sessions")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) PermanentDeleteSessionsByUser(userId string) error {
|
||||
_, err := me.GetMasterX().Exec("DELETE FROM Sessions WHERE UserId = ?", userId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Session with userId=%s", userId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) UpdateExpiresAt(sessionId string, time int64) error {
|
||||
_, err := me.GetMasterX().Exec("UPDATE Sessions SET ExpiresAt = ?, ExpiredNotify = false WHERE Id = ?", time, sessionId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to update Session with sessionId=%s", sessionId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (me *SqlSessionStore) GetLastSessionRowCreateAt() (int64, error) {
|
||||
query := `SELECT CREATEAT FROM Sessions ORDER BY CREATEAT DESC LIMIT 1`
|
||||
var createAt int64
|
||||
err := me.GetReplicaX().Get(&createAt, query)
|
||||
if err != nil {
|
||||
return 0, errors.Wrapf(err, "failed to get last session createat")
|
||||
}
|
||||
|
||||
return createAt, nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) UpdateLastActivityAt(sessionId string, time int64) error {
|
||||
_, err := me.GetMasterX().Exec("UPDATE Sessions SET LastActivityAt = ? WHERE Id = ?", time, sessionId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to update Session with id=%s", sessionId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) UpdateRoles(userId, roles string) (string, error) {
|
||||
if len(roles) > model.UserRolesMaxLength {
|
||||
return "", fmt.Errorf("given session roles length (%d) exceeds max storage limit (%d)", len(roles), model.UserRolesMaxLength)
|
||||
}
|
||||
|
||||
_, err := me.GetMasterX().Exec("UPDATE Sessions SET Roles = ? WHERE UserId = ?", roles, userId)
|
||||
if err != nil {
|
||||
return "", errors.Wrapf(err, "failed to update Session with userId=%s and roles=%s", userId, roles)
|
||||
}
|
||||
return userId, nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) UpdateDeviceId(id string, deviceId string, expiresAt int64) (string, error) {
|
||||
query := "UPDATE Sessions SET DeviceId = ?, ExpiresAt = ?, ExpiredNotify = false WHERE Id = ?"
|
||||
|
||||
_, err := me.GetMasterX().Exec(query, deviceId, expiresAt, id)
|
||||
if err != nil {
|
||||
return "", errors.Wrapf(err, "failed to update Session with id=%s", id)
|
||||
}
|
||||
return deviceId, nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) UpdateProps(session *model.Session) error {
|
||||
jsonProps, err := json.Marshal(session.Props)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed marshalling session props")
|
||||
}
|
||||
if me.IsBinaryParamEnabled() {
|
||||
jsonProps = AppendBinaryFlag(jsonProps)
|
||||
}
|
||||
query, args, err := me.getQueryBuilder().
|
||||
Update("Sessions").
|
||||
Set("Props", jsonProps).
|
||||
Where(sq.Eq{"Id": session.Id}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
errors.Wrap(err, "sessions_tosql")
|
||||
}
|
||||
_, err = me.GetMasterX().Exec(query, args...)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to update Session")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) AnalyticsSessionCount() (int64, error) {
|
||||
var count int64
|
||||
query :=
|
||||
`SELECT
|
||||
COUNT(*)
|
||||
FROM
|
||||
Sessions
|
||||
WHERE ExpiresAt > ?`
|
||||
if err := me.GetReplicaX().Get(&count, query, model.GetMillis()); err != nil {
|
||||
return int64(0), errors.Wrap(err, "failed to count Sessions")
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (me SqlSessionStore) Cleanup(expiryTime int64, batchSize int64) error {
|
||||
var query string
|
||||
if me.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = "DELETE FROM Sessions WHERE Id IN (SELECT Id FROM Sessions WHERE ExpiresAt != 0 AND ? > ExpiresAt LIMIT ?)"
|
||||
} else {
|
||||
query = "DELETE FROM Sessions WHERE ExpiresAt != 0 AND ? > ExpiresAt LIMIT ?"
|
||||
}
|
||||
|
||||
var rowsAffected int64 = 1
|
||||
|
||||
for rowsAffected > 0 {
|
||||
sqlResult, err := me.GetMasterX().Exec(query, expiryTime, batchSize)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "unable to delete sessions")
|
||||
}
|
||||
var rowErr error
|
||||
rowsAffected, rowErr = sqlResult.RowsAffected()
|
||||
if rowErr != nil {
|
||||
return errors.Wrap(err, "unable to delete sessions")
|
||||
}
|
||||
|
||||
time.Sleep(sessionsCleanupDelay)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/session_store_test.go
Обычный файл
14
server/channels/store/sqlstore/session_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestSessionStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestSessionStore)
|
||||
}
|
||||
823
server/channels/store/sqlstore/shared_channel_store.go
Обычный файл
823
server/channels/store/sqlstore/shared_channel_store.go
Обычный файл
@@ -0,0 +1,823 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultGetUsersForSyncLimit = 100
|
||||
)
|
||||
|
||||
type SqlSharedChannelStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlSharedChannelStore(sqlStore *SqlStore) store.SharedChannelStore {
|
||||
return &SqlSharedChannelStore{
|
||||
SqlStore: sqlStore,
|
||||
}
|
||||
}
|
||||
|
||||
// Save inserts a new shared channel record.
|
||||
func (s SqlSharedChannelStore) Save(sc *model.SharedChannel) (sh *model.SharedChannel, err error) {
|
||||
sc.PreSave()
|
||||
if err := sc.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// make sure the shared channel is associated with a real channel.
|
||||
channel, err := s.stores.channel.Get(sc.ChannelId, true)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid channel: %w", err)
|
||||
}
|
||||
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
|
||||
query, args, err := s.getQueryBuilder().Insert("SharedChannels").
|
||||
Columns("ChannelId", "TeamId", "Home", "ReadOnly", "ShareName", "ShareDisplayName", "SharePurpose", "ShareHeader", "CreatorId", "CreateAt", "UpdateAt", "RemoteId").
|
||||
Values(sc.ChannelId, sc.TeamId, sc.Home, sc.ReadOnly, sc.ShareName, sc.ShareDisplayName, sc.SharePurpose, sc.ShareHeader, sc.CreatorId, sc.CreateAt, sc.UpdateAt, sc.RemoteId).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "savesharedchannel_tosql")
|
||||
}
|
||||
if _, err := transaction.Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "save_shared_channel: ChannelId=%s", sc.ChannelId)
|
||||
}
|
||||
|
||||
// set `Shared` flag in Channels table if needed
|
||||
if channel.Shared == nil || !*channel.Shared {
|
||||
if err := s.stores.channel.SetShared(channel.Id, true); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
if err := transaction.Commit(); err != nil {
|
||||
return nil, errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
return sc, nil
|
||||
}
|
||||
|
||||
// Get fetches a shared channel by channel_id.
|
||||
func (s SqlSharedChannelStore) Get(channelId string) (*model.SharedChannel, error) {
|
||||
var sc model.SharedChannel
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("SharedChannels").
|
||||
Where(sq.Eq{"SharedChannels.ChannelId": channelId})
|
||||
|
||||
squery, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "getsharedchannel_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Get(&sc, squery, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("SharedChannel", channelId)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to find shared channel with ChannelId=%s", channelId)
|
||||
}
|
||||
return &sc, nil
|
||||
}
|
||||
|
||||
// HasChannel returns whether a given channelID is a shared channel or not.
|
||||
func (s SqlSharedChannelStore) HasChannel(channelID string) (bool, error) {
|
||||
builder := s.getQueryBuilder().
|
||||
Select("1").
|
||||
Prefix("SELECT EXISTS (").
|
||||
From("SharedChannels").
|
||||
Where(sq.Eq{"SharedChannels.ChannelId": channelID}).
|
||||
Suffix(")")
|
||||
|
||||
query, args, err := builder.ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrapf(err, "get_shared_channel_exists_tosql")
|
||||
}
|
||||
|
||||
var exists bool
|
||||
if err := s.GetReplicaX().Get(&exists, query, args...); err != nil {
|
||||
return exists, errors.Wrapf(err, "failed to get shared channel for channel_id=%s", channelID)
|
||||
}
|
||||
return exists, nil
|
||||
}
|
||||
|
||||
// GetAll fetches a paginated list of shared channels filtered by SharedChannelSearchOpts.
|
||||
func (s SqlSharedChannelStore) GetAll(offset, limit int, opts model.SharedChannelFilterOpts) ([]*model.SharedChannel, error) {
|
||||
if opts.ExcludeHome && opts.ExcludeRemote {
|
||||
return nil, errors.New("cannot exclude home and remote shared channels")
|
||||
}
|
||||
|
||||
safeConv := func(offset, limit int) (uint64, uint64, error) {
|
||||
if offset < 0 {
|
||||
return 0, 0, errors.New("offset must be positive integer")
|
||||
}
|
||||
if limit < 0 {
|
||||
return 0, 0, errors.New("limit must be positive integer")
|
||||
}
|
||||
return uint64(offset), uint64(limit), nil
|
||||
}
|
||||
|
||||
safeOffset, safeLimit, err := safeConv(offset, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query := s.getSharedChannelsQuery(opts, false)
|
||||
query = query.OrderBy("sc.ShareDisplayName, sc.ShareName").Limit(safeLimit).Offset(safeOffset)
|
||||
|
||||
squery, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to create query")
|
||||
}
|
||||
|
||||
channels := []*model.SharedChannel{}
|
||||
err = s.GetReplicaX().Select(&channels, squery, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get shared channels")
|
||||
}
|
||||
return channels, nil
|
||||
}
|
||||
|
||||
// GetAllCount returns the number of shared channels that would be fetched using SharedChannelSearchOpts.
|
||||
func (s SqlSharedChannelStore) GetAllCount(opts model.SharedChannelFilterOpts) (int64, error) {
|
||||
if opts.ExcludeHome && opts.ExcludeRemote {
|
||||
return 0, errors.New("cannot exclude home and remote shared channels")
|
||||
}
|
||||
|
||||
query := s.getSharedChannelsQuery(opts, true)
|
||||
squery, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "failed to create query")
|
||||
}
|
||||
|
||||
var count int64
|
||||
err = s.GetReplicaX().Get(&count, squery, args...)
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "failed to count channels")
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s SqlSharedChannelStore) getSharedChannelsQuery(opts model.SharedChannelFilterOpts, forCount bool) sq.SelectBuilder {
|
||||
var selectStr string
|
||||
if forCount {
|
||||
selectStr = "count(sc.ChannelId)"
|
||||
} else {
|
||||
selectStr = "sc.*"
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select(selectStr).
|
||||
From("SharedChannels AS sc")
|
||||
|
||||
if opts.MemberId != "" {
|
||||
query = query.Join("ChannelMembers AS cm ON cm.ChannelId = sc.ChannelId").
|
||||
Where(sq.Eq{"cm.UserId": opts.MemberId})
|
||||
}
|
||||
|
||||
if opts.TeamId != "" {
|
||||
query = query.Where(sq.Eq{"sc.TeamId": opts.TeamId})
|
||||
}
|
||||
|
||||
if opts.CreatorId != "" {
|
||||
query = query.Where(sq.Eq{"sc.CreatorId": opts.CreatorId})
|
||||
}
|
||||
|
||||
if opts.ExcludeHome {
|
||||
query = query.Where(sq.NotEq{"sc.Home": true})
|
||||
}
|
||||
|
||||
if opts.ExcludeRemote {
|
||||
query = query.Where(sq.Eq{"sc.Home": true})
|
||||
}
|
||||
|
||||
return query
|
||||
}
|
||||
|
||||
// Update updates the shared channel.
|
||||
func (s SqlSharedChannelStore) Update(sc *model.SharedChannel) (*model.SharedChannel, error) {
|
||||
if err := sc.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query, args, err := s.getQueryBuilder().Update("SharedChannels").Set("ChannelId", sc.ChannelId).
|
||||
Set("TeamId", sc.TeamId).
|
||||
Set("Home", sc.Home).
|
||||
Set("ReadOnly", sc.ReadOnly).
|
||||
Set("ShareName", sc.ShareName).
|
||||
Set("ShareDisplayName", sc.ShareDisplayName).
|
||||
Set("SharePurpose", sc.SharePurpose).
|
||||
Set("ShareHeader", sc.ShareHeader).
|
||||
Set("CreatorId", sc.CreatorId).
|
||||
Set("CreateAt", sc.CreateAt).
|
||||
Set("UpdateAt", sc.UpdateAt).
|
||||
Set("RemoteId", sc.RemoteId).
|
||||
Where(sq.Eq{"ChannelId": sc.ChannelId}).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "updatesharedchannel_tosql")
|
||||
}
|
||||
res, err := s.GetMasterX().Exec(query, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update shared channel with channelId=%s", sc.ChannelId)
|
||||
}
|
||||
count, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "error while getting rows_affected")
|
||||
}
|
||||
if count != 1 {
|
||||
return nil, fmt.Errorf("expected number of shared channels to be updated is 1 but was %d", count)
|
||||
}
|
||||
return sc, nil
|
||||
}
|
||||
|
||||
// Delete deletes a single shared channel plus associated SharedChannelRemotes.
|
||||
// Returns true if shared channel found and deleted, false if not found.
|
||||
func (s SqlSharedChannelStore) Delete(channelId string) (ok bool, err error) {
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "DeleteSharedChannel: begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
|
||||
squery, args, err := s.getQueryBuilder().
|
||||
Delete("SharedChannels").
|
||||
Where(sq.Eq{"SharedChannels.ChannelId": channelId}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "delete_shared_channel_tosql")
|
||||
}
|
||||
|
||||
result, err := transaction.Exec(squery, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to delete SharedChannel")
|
||||
}
|
||||
|
||||
// Also remove remotes from SharedChannelRemotes (if any).
|
||||
squery, args, err = s.getQueryBuilder().
|
||||
Delete("SharedChannelRemotes").
|
||||
Where(sq.Eq{"ChannelId": channelId}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "delete_shared_channel_remotes_tosql")
|
||||
}
|
||||
|
||||
_, err = transaction.Exec(squery, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to delete SharedChannelRemotes")
|
||||
}
|
||||
|
||||
count, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to determine rows affected")
|
||||
}
|
||||
|
||||
if count > 0 {
|
||||
// unset the channel's Shared flag
|
||||
if err = s.Channel().SetShared(channelId, false); err != nil {
|
||||
return false, errors.Wrap(err, "error unsetting channel share flag")
|
||||
}
|
||||
}
|
||||
|
||||
if err = transaction.Commit(); err != nil {
|
||||
return false, errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
// SaveRemote inserts a new shared channel remote record.
|
||||
func (s SqlSharedChannelStore) SaveRemote(remote *model.SharedChannelRemote) (*model.SharedChannelRemote, error) {
|
||||
remote.PreSave()
|
||||
if err := remote.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// make sure the shared channel remote is associated with a real channel.
|
||||
if _, err := s.stores.channel.Get(remote.ChannelId, true); err != nil {
|
||||
return nil, fmt.Errorf("invalid channel: %w", err)
|
||||
}
|
||||
|
||||
query, args, err := s.getQueryBuilder().Insert("SharedChannelRemotes").
|
||||
Columns("Id", "ChannelId", "CreatorId", "CreateAt", "UpdateAt", "IsInviteAccepted", "IsInviteConfirmed", "RemoteId", "LastPostUpdateAt", "LastPostId").
|
||||
Values(remote.Id, remote.ChannelId, remote.CreatorId, remote.CreateAt, remote.UpdateAt, remote.IsInviteAccepted, remote.IsInviteConfirmed, remote.RemoteId, remote.LastPostUpdateAt, remote.LastPostId).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "savesharedchannelremote_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "save_shared_channel_remote: channel_id=%s, id=%s", remote.ChannelId, remote.Id)
|
||||
}
|
||||
return remote, nil
|
||||
}
|
||||
|
||||
// Update updates the shared channel remote.
|
||||
func (s SqlSharedChannelStore) UpdateRemote(remote *model.SharedChannelRemote) (*model.SharedChannelRemote, error) {
|
||||
if err := remote.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query, args, err := s.getQueryBuilder().Update("SharedChannelRemotes").
|
||||
Set("CreatorId", remote.CreatorId).
|
||||
Set("CreateAt", remote.CreateAt).
|
||||
Set("UpdateAt", remote.UpdateAt).
|
||||
Set("IsInviteAccepted", remote.IsInviteAccepted).
|
||||
Set("IsInviteConfirmed", remote.IsInviteConfirmed).
|
||||
Set("RemoteId", remote.RemoteId).
|
||||
Set("LastPostUpdateAt", remote.LastPostUpdateAt).
|
||||
Set("LastPostId", remote.LastPostId).
|
||||
Where(sq.And{
|
||||
sq.Eq{"Id": remote.Id},
|
||||
sq.Eq{"ChannelId": remote.ChannelId},
|
||||
}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "updatesharedchannelremote_tosql")
|
||||
}
|
||||
|
||||
res, err := s.GetMasterX().Exec(query, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update shared channel remote with remoteId=%s", remote.Id)
|
||||
}
|
||||
|
||||
count, err := res.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "error while getting rows_affected")
|
||||
}
|
||||
if count != 1 {
|
||||
return nil, fmt.Errorf("expected number of shared channel remotes to be updated is 1 but was %d", count)
|
||||
}
|
||||
return remote, nil
|
||||
}
|
||||
|
||||
// GetRemote fetches a shared channel remote by id.
|
||||
func (s SqlSharedChannelStore) GetRemote(id string) (*model.SharedChannelRemote, error) {
|
||||
var remote model.SharedChannelRemote
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("SharedChannelRemotes").
|
||||
Where(sq.Eq{"SharedChannelRemotes.Id": id})
|
||||
|
||||
squery, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "get_shared_channel_remote_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Get(&remote, squery, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("SharedChannelRemote", id)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to find shared channel remote with id=%s", id)
|
||||
}
|
||||
return &remote, nil
|
||||
}
|
||||
|
||||
// GetRemoteByIds fetches a shared channel remote by channel id and remote cluster id.
|
||||
func (s SqlSharedChannelStore) GetRemoteByIds(channelId string, remoteId string) (*model.SharedChannelRemote, error) {
|
||||
var remote model.SharedChannelRemote
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("SharedChannelRemotes").
|
||||
Where(sq.Eq{"SharedChannelRemotes.ChannelId": channelId}).
|
||||
Where(sq.Eq{"SharedChannelRemotes.RemoteId": remoteId})
|
||||
|
||||
squery, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "get_shared_channel_remote_by_ids_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Get(&remote, squery, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("SharedChannelRemote", fmt.Sprintf("channelId=%s, remoteId=%s", channelId, remoteId))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to find shared channel remote with channelId=%s, remoteId=%s", channelId, remoteId)
|
||||
}
|
||||
return &remote, nil
|
||||
}
|
||||
|
||||
// GetRemotes fetches all shared channel remotes associated with channel_id.
|
||||
func (s SqlSharedChannelStore) GetRemotes(opts model.SharedChannelRemoteFilterOpts) ([]*model.SharedChannelRemote, error) {
|
||||
remotes := []*model.SharedChannelRemote{}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("SharedChannelRemotes")
|
||||
|
||||
if opts.ChannelId != "" {
|
||||
query = query.Where(sq.Eq{"ChannelId": opts.ChannelId})
|
||||
}
|
||||
|
||||
if opts.RemoteId != "" {
|
||||
query = query.Where(sq.Eq{"RemoteId": opts.RemoteId})
|
||||
}
|
||||
|
||||
if !opts.InclUnconfirmed {
|
||||
query = query.Where(sq.Eq{"IsInviteConfirmed": true})
|
||||
}
|
||||
|
||||
squery, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "get_shared_channel_remotes_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&remotes, squery, args...); err != nil {
|
||||
if err != sql.ErrNoRows {
|
||||
return nil, errors.Wrapf(err, "failed to get shared channel remotes for channel_id=%s; remote_id=%s",
|
||||
opts.ChannelId, opts.RemoteId)
|
||||
}
|
||||
}
|
||||
return remotes, nil
|
||||
}
|
||||
|
||||
// HasRemote returns whether a given remoteId and channelId are present in the shared channel remotes or not.
|
||||
func (s SqlSharedChannelStore) HasRemote(channelID string, remoteId string) (bool, error) {
|
||||
builder := s.getQueryBuilder().
|
||||
Select("1").
|
||||
Prefix("SELECT EXISTS (").
|
||||
From("SharedChannelRemotes").
|
||||
Where(sq.Eq{"RemoteId": remoteId}).
|
||||
Where(sq.Eq{"ChannelId": channelID}).
|
||||
Suffix(")")
|
||||
|
||||
query, args, err := builder.ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrapf(err, "get_shared_channel_hasremote_tosql")
|
||||
}
|
||||
|
||||
var hasRemote bool
|
||||
if err := s.GetReplicaX().Get(&hasRemote, query, args...); err != nil {
|
||||
return hasRemote, errors.Wrapf(err, "failed to get channel remotes for channel_id=%s", channelID)
|
||||
}
|
||||
return hasRemote, nil
|
||||
}
|
||||
|
||||
// GetRemoteForUser returns a remote cluster for the given userId only if the user belongs to at least one channel
|
||||
// shared with the remote.
|
||||
func (s SqlSharedChannelStore) GetRemoteForUser(remoteId string, userId string) (*model.RemoteCluster, error) {
|
||||
builder := s.getQueryBuilder().
|
||||
Select("rc.*").
|
||||
From("RemoteClusters AS rc").
|
||||
Join("SharedChannelRemotes AS scr ON rc.RemoteId = scr.RemoteId").
|
||||
Join("ChannelMembers AS cm ON scr.ChannelId = cm.ChannelId").
|
||||
Where(sq.Eq{"rc.RemoteId": remoteId}).
|
||||
Where(sq.Eq{"cm.UserId": userId})
|
||||
|
||||
query, args, err := builder.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "get_remote_for_user_tosql")
|
||||
}
|
||||
|
||||
var rc model.RemoteCluster
|
||||
if err := s.GetReplicaX().Get(&rc, query, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("RemoteCluster", remoteId)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get remote for user_id=%s", userId)
|
||||
}
|
||||
return &rc, nil
|
||||
}
|
||||
|
||||
// UpdateRemoteCursor updates the LastPostUpdateAt timestamp and LastPostId for the specified SharedChannelRemote.
|
||||
func (s SqlSharedChannelStore) UpdateRemoteCursor(id string, cursor model.GetPostsSinceForSyncCursor) error {
|
||||
squery, args, err := s.getQueryBuilder().
|
||||
Update("SharedChannelRemotes").
|
||||
Set("LastPostUpdateAt", cursor.LastPostUpdateAt).
|
||||
Set("LastPostId", cursor.LastPostId).
|
||||
Where(sq.Eq{"Id": id}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "update_shared_channel_remote_cursor_tosql")
|
||||
}
|
||||
|
||||
result, err := s.GetMasterX().Exec(squery, args...)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to update cursor for SharedChannelRemote")
|
||||
}
|
||||
|
||||
count, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to determine rows affected")
|
||||
}
|
||||
if count == 0 {
|
||||
return fmt.Errorf("id not found: %s", id)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteRemote deletes a single shared channel remote.
|
||||
// Returns true if remote found and deleted, false if not found.
|
||||
func (s SqlSharedChannelStore) DeleteRemote(id string) (bool, error) {
|
||||
squery, args, err := s.getQueryBuilder().
|
||||
Delete("SharedChannelRemotes").
|
||||
Where(sq.Eq{"Id": id}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "delete_shared_channel_remote_tosql")
|
||||
}
|
||||
|
||||
result, err := s.GetMasterX().Exec(squery, args...)
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to delete SharedChannelRemote")
|
||||
}
|
||||
|
||||
count, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return false, errors.Wrap(err, "failed to determine rows affected")
|
||||
}
|
||||
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
// GetRemotesStatus returns the status for each remote invited to the
|
||||
// specified shared channel.
|
||||
func (s SqlSharedChannelStore) GetRemotesStatus(channelId string) ([]*model.SharedChannelRemoteStatus, error) {
|
||||
status := []*model.SharedChannelRemoteStatus{}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("scr.ChannelId, rc.DisplayName, rc.SiteURL, rc.LastPingAt, sc.ReadOnly, scr.IsInviteAccepted").
|
||||
From("SharedChannelRemotes scr, RemoteClusters rc, SharedChannels sc").
|
||||
Where("scr.RemoteId = rc.RemoteId").
|
||||
Where("scr.ChannelId = sc.ChannelId").
|
||||
Where(sq.Eq{"scr.ChannelId": channelId})
|
||||
|
||||
squery, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "get_shared_channel_remotes_status_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&status, squery, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("SharedChannelRemoteStatus", channelId)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get shared channel remote status for channel_id=%s", channelId)
|
||||
}
|
||||
return status, nil
|
||||
}
|
||||
|
||||
// SaveUser inserts a new shared channel user record to the SharedChannelUsers table.
|
||||
func (s SqlSharedChannelStore) SaveUser(scUser *model.SharedChannelUser) (*model.SharedChannelUser, error) {
|
||||
scUser.PreSave()
|
||||
if err := scUser.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query, args, err := s.getQueryBuilder().Insert("SharedChannelUsers").
|
||||
Columns("Id", "UserId", "ChannelId", "RemoteId", "CreateAt", "LastSyncAt").
|
||||
Values(scUser.Id, scUser.UserId, scUser.ChannelId, scUser.RemoteId, scUser.CreateAt, scUser.LastSyncAt).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "savesharedchanneluser_tosql")
|
||||
}
|
||||
if _, err := s.GetMasterX().Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "save_shared_channel_user: user_id=%s, remote_id=%s", scUser.UserId, scUser.RemoteId)
|
||||
}
|
||||
return scUser, nil
|
||||
}
|
||||
|
||||
// GetSingleUser fetches a shared channel user based on userID, channelID and remoteID.
|
||||
func (s SqlSharedChannelStore) GetSingleUser(userID string, channelID string, remoteID string) (*model.SharedChannelUser, error) {
|
||||
var scu model.SharedChannelUser
|
||||
|
||||
squery, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("SharedChannelUsers").
|
||||
Where(sq.Eq{"SharedChannelUsers.UserId": userID}).
|
||||
Where(sq.Eq{"SharedChannelUsers.RemoteId": remoteID}).
|
||||
Where(sq.Eq{"SharedChannelUsers.ChannelId": channelID}).
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "getsharedchannelsingleuser_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Get(&scu, squery, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("SharedChannelUser", userID)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to find shared channel user with UserId=%s, ChannelId=%s, RemoteId=%s", userID, channelID, remoteID)
|
||||
}
|
||||
return &scu, nil
|
||||
}
|
||||
|
||||
// GetUsersForUser fetches all shared channel user records based on userID.
|
||||
func (s SqlSharedChannelStore) GetUsersForUser(userID string) ([]*model.SharedChannelUser, error) {
|
||||
squery, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("SharedChannelUsers").
|
||||
Where(sq.Eq{"SharedChannelUsers.UserId": userID}).
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "getsharedchanneluser_tosql")
|
||||
}
|
||||
|
||||
users := []*model.SharedChannelUser{}
|
||||
if err := s.GetReplicaX().Select(&users, squery, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return make([]*model.SharedChannelUser, 0), nil
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to find shared channel user with UserId=%s", userID)
|
||||
}
|
||||
return users, nil
|
||||
}
|
||||
|
||||
// GetUsersForSync fetches all shared channel users that need to be synchronized, meaning their
|
||||
// `SharedChannelUsers.LastSyncAt` is less than or equal to `User.UpdateAt`.
|
||||
func (s SqlSharedChannelStore) GetUsersForSync(filter model.GetUsersForSyncFilter) ([]*model.User, error) {
|
||||
if filter.Limit <= 0 {
|
||||
filter.Limit = DefaultGetUsersForSyncLimit
|
||||
}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("u.*").
|
||||
Distinct().
|
||||
From("Users AS u").
|
||||
Join("SharedChannelUsers AS scu ON u.Id = scu.UserId").
|
||||
OrderBy("u.Id").
|
||||
Limit(filter.Limit)
|
||||
|
||||
if filter.CheckProfileImage {
|
||||
query = query.Where("scu.LastSyncAt < u.LastPictureUpdate")
|
||||
} else {
|
||||
query = query.Where("scu.LastSyncAt < u.UpdateAt")
|
||||
}
|
||||
|
||||
if filter.ChannelID != "" {
|
||||
query = query.Where(sq.Eq{"scu.ChannelId": filter.ChannelID})
|
||||
}
|
||||
|
||||
sqlQuery, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "getsharedchannelusersforsync_tosql")
|
||||
}
|
||||
|
||||
users := []*model.User{}
|
||||
if err := s.GetReplicaX().Select(&users, sqlQuery, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return make([]*model.User, 0), nil
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to fetch shared channel users with ChannelId=%s",
|
||||
filter.ChannelID)
|
||||
}
|
||||
return users, nil
|
||||
}
|
||||
|
||||
// UpdateUserLastSyncAt updates the LastSyncAt timestamp for the specified SharedChannelUser.
|
||||
func (s SqlSharedChannelStore) UpdateUserLastSyncAt(userID string, channelID string, remoteID string) error {
|
||||
var query string
|
||||
if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = `
|
||||
UPDATE
|
||||
SharedChannelUsers AS scu
|
||||
SET
|
||||
LastSyncAt = GREATEST(Users.UpdateAt, Users.LastPictureUpdate)
|
||||
FROM
|
||||
Users
|
||||
WHERE
|
||||
Users.Id = scu.UserId AND scu.UserId = ? AND scu.ChannelId = ? AND scu.RemoteId = ?
|
||||
`
|
||||
} else if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = `
|
||||
UPDATE
|
||||
SharedChannelUsers AS scu
|
||||
INNER JOIN
|
||||
Users ON scu.UserId = Users.Id
|
||||
SET
|
||||
LastSyncAt = GREATEST(Users.UpdateAt, Users.LastPictureUpdate)
|
||||
WHERE
|
||||
scu.UserId = ? AND scu.ChannelId = ? AND scu.RemoteId = ?
|
||||
`
|
||||
} else {
|
||||
return errors.New("unsupported DB driver " + s.DriverName())
|
||||
}
|
||||
|
||||
result, err := s.GetMasterX().Exec(query, userID, channelID, remoteID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update LastSyncAt for SharedChannelUser with userId=%s, channelId=%s, remoteId=%s: %w",
|
||||
userID, channelID, remoteID, err)
|
||||
}
|
||||
|
||||
count, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to determine rows affected")
|
||||
}
|
||||
if count == 0 {
|
||||
return fmt.Errorf("SharedChannelUser not found: userId=%s, channelId=%s, remoteId=%s", userID, channelID, remoteID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SaveAttachment inserts a new shared channel file attachment record to the SharedChannelFiles table.
|
||||
func (s SqlSharedChannelStore) SaveAttachment(attachment *model.SharedChannelAttachment) (*model.SharedChannelAttachment, error) {
|
||||
attachment.PreSave()
|
||||
if err := attachment.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query, args, err := s.getQueryBuilder().Insert("SharedChannelAttachments").
|
||||
Columns("Id", "FileId", "RemoteId", "CreateAt", "LastSyncAt").
|
||||
Values(attachment.Id, attachment.FileId, attachment.RemoteId, attachment.CreateAt, attachment.LastSyncAt).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "savesahredchannelattachment_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "save_shared_channel_attachment: file_id=%s, remote_id=%s", attachment.FileId, attachment.RemoteId)
|
||||
}
|
||||
return attachment, nil
|
||||
}
|
||||
|
||||
// UpsertAttachment inserts a new shared channel file attachment record to the SharedChannelFiles table or updates its
|
||||
// LastSyncAt.
|
||||
func (s SqlSharedChannelStore) UpsertAttachment(attachment *model.SharedChannelAttachment) (string, error) {
|
||||
attachment.PreSave()
|
||||
if err := attachment.IsValid(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
query := s.getQueryBuilder().
|
||||
Insert("SharedChannelAttachments").
|
||||
Columns("Id", "FileId", "RemoteId", "CreateAt", "LastSyncAt").
|
||||
Values(attachment.Id, attachment.FileId, attachment.RemoteId, attachment.CreateAt, attachment.LastSyncAt)
|
||||
|
||||
if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE LastSyncAt = ?", attachment.LastSyncAt))
|
||||
} else if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = query.SuffixExpr(sq.Expr("ON CONFLICT (id) DO UPDATE SET LastSyncAt = ?", attachment.LastSyncAt))
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return "", errors.Wrap(err, "upsertsharedchannelattachment_tosql")
|
||||
}
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return "", errors.Wrap(err, "failed to upsert SharedChannelAttachments")
|
||||
}
|
||||
return attachment.Id, nil
|
||||
}
|
||||
|
||||
// GetAttachment fetches a shared channel file attachment record based on file_id and remoteId.
|
||||
func (s SqlSharedChannelStore) GetAttachment(fileId string, remoteId string) (*model.SharedChannelAttachment, error) {
|
||||
var attachment model.SharedChannelAttachment
|
||||
|
||||
squery, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("SharedChannelAttachments").
|
||||
Where(sq.Eq{"SharedChannelAttachments.FileId": fileId}).
|
||||
Where(sq.Eq{"SharedChannelAttachments.RemoteId": remoteId}).
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "getsharedchannelattachment_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Get(&attachment, squery, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("SharedChannelAttachment", fileId)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to find shared channel attachment with FileId=%s, RemoteId=%s", fileId, remoteId)
|
||||
}
|
||||
return &attachment, nil
|
||||
}
|
||||
|
||||
// UpdateAttachmentLastSyncAt updates the LastSyncAt timestamp for the specified SharedChannelAttachment.
|
||||
func (s SqlSharedChannelStore) UpdateAttachmentLastSyncAt(id string, syncTime int64) error {
|
||||
squery, args, err := s.getQueryBuilder().
|
||||
Update("SharedChannelAttachments").
|
||||
Set("LastSyncAt", syncTime).
|
||||
Where(sq.Eq{"Id": id}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "update_shared_channel_attachment_last_sync_at_tosql")
|
||||
}
|
||||
|
||||
result, err := s.GetMasterX().Exec(squery, args...)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to update LastSyncAt for SharedChannelAttachment")
|
||||
}
|
||||
|
||||
count, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to determine rows affected")
|
||||
}
|
||||
if count == 0 {
|
||||
return fmt.Errorf("id not found: %s", id)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/shared_channel_store_test.go
Обычный файл
14
server/channels/store/sqlstore/shared_channel_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestSharedChannelStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestSharedChannelStore)
|
||||
}
|
||||
461
server/channels/store/sqlstore/sqlx_wrapper.go
Обычный файл
461
server/channels/store/sqlstore/sqlx_wrapper.go
Обычный файл
@@ -0,0 +1,461 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode"
|
||||
|
||||
"github.com/jmoiron/sqlx"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
)
|
||||
|
||||
type StoreTestWrapper struct {
|
||||
orig *SqlStore
|
||||
}
|
||||
|
||||
func NewStoreTestWrapper(orig *SqlStore) *StoreTestWrapper {
|
||||
return &StoreTestWrapper{orig}
|
||||
}
|
||||
|
||||
func (w *StoreTestWrapper) GetMasterX() storetest.SqlXExecutor {
|
||||
return w.orig.GetMasterX()
|
||||
}
|
||||
|
||||
func (w *StoreTestWrapper) DriverName() string {
|
||||
return w.orig.DriverName()
|
||||
}
|
||||
|
||||
type Builder interface {
|
||||
ToSql() (string, []any, 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 any, query string, args ...any) error
|
||||
GetBuilder(dest any, builder Builder) error
|
||||
NamedExec(query string, arg any) (sql.Result, error)
|
||||
Exec(query string, args ...any) (sql.Result, error)
|
||||
ExecBuilder(builder Builder) (sql.Result, error)
|
||||
ExecRaw(query string, args ...any) (sql.Result, error)
|
||||
NamedQuery(query string, arg any) (*sqlx.Rows, error)
|
||||
QueryRowX(query string, args ...any) *sqlx.Row
|
||||
QueryX(query string, args ...any) (*sqlx.Rows, error)
|
||||
Select(dest any, query string, args ...any) error
|
||||
SelectBuilder(dest any, builder Builder) error
|
||||
}
|
||||
|
||||
// namedParamRegex is used to capture all named parameters and convert them
|
||||
// to lowercase. This is necessary to be able to use a single query for both
|
||||
// Postgres and MySQL.
|
||||
// This will also lowercase any constant strings containing a :, but sqlx
|
||||
// will fail the query, so it won't be checked in inadvertently.
|
||||
var namedParamRegex = regexp.MustCompile(`:\w+`)
|
||||
|
||||
type sqlxDBWrapper struct {
|
||||
*sqlx.DB
|
||||
queryTimeout time.Duration
|
||||
trace bool
|
||||
}
|
||||
|
||||
func newSqlxDBWrapper(db *sqlx.DB, timeout time.Duration, trace bool) *sqlxDBWrapper {
|
||||
return &sqlxDBWrapper{
|
||||
DB: db,
|
||||
queryTimeout: timeout,
|
||||
trace: trace,
|
||||
}
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) Stats() sql.DBStats {
|
||||
return w.DB.Stats()
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) Beginx() (*sqlxTxWrapper, error) {
|
||||
tx, err := w.DB.Beginx()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return newSqlxTxWrapper(tx, w.queryTimeout, w.trace), nil
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) BeginXWithIsolation(opts *sql.TxOptions) (*sqlxTxWrapper, error) {
|
||||
tx, err := w.DB.BeginTxx(context.Background(), opts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return newSqlxTxWrapper(tx, w.queryTimeout, w.trace), nil
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) Get(dest any, query string, args ...any) error {
|
||||
query = w.DB.Rebind(query)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.DB.GetContext(ctx, dest, query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) GetBuilder(dest any, 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 any) (sql.Result, error) {
|
||||
if w.DB.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = namedParamRegex.ReplaceAllStringFunc(query, strings.ToLower)
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), arg)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.DB.NamedExecContext(ctx, query, arg)
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) Exec(query string, args ...any) (sql.Result, error) {
|
||||
query = w.DB.Rebind(query)
|
||||
|
||||
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 ...any) (sql.Result, error) {
|
||||
query = w.DB.Rebind(query)
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.DB.ExecContext(context.Background(), 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 *sqlxDBWrapper) ExecRaw(query string, args ...any) (sql.Result, error) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.DB.ExecContext(ctx, query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) NamedQuery(query string, arg any) (*sqlx.Rows, error) {
|
||||
if w.DB.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = namedParamRegex.ReplaceAllStringFunc(query, strings.ToLower)
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), arg)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.DB.NamedQueryContext(ctx, query, arg)
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) QueryRowX(query string, args ...any) *sqlx.Row {
|
||||
query = w.DB.Rebind(query)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.DB.QueryRowxContext(ctx, query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) QueryX(query string, args ...any) (*sqlx.Rows, error) {
|
||||
query = w.DB.Rebind(query)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.DB.QueryxContext(ctx, query, args)
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) Select(dest any, query string, args ...any) error {
|
||||
return w.SelectCtx(context.Background(), dest, query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) SelectCtx(ctx context.Context, dest any, query string, args ...any) error {
|
||||
query = w.DB.Rebind(query)
|
||||
ctx, cancel := context.WithTimeout(ctx, w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.DB.SelectContext(ctx, dest, query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxDBWrapper) SelectBuilder(dest any, 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
|
||||
trace bool
|
||||
}
|
||||
|
||||
func newSqlxTxWrapper(tx *sqlx.Tx, timeout time.Duration, trace bool) *sqlxTxWrapper {
|
||||
return &sqlxTxWrapper{
|
||||
Tx: tx,
|
||||
queryTimeout: timeout,
|
||||
trace: trace,
|
||||
}
|
||||
}
|
||||
|
||||
func (w *sqlxTxWrapper) Get(dest any, query string, args ...any) error {
|
||||
query = w.Tx.Rebind(query)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.Tx.GetContext(ctx, dest, query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxTxWrapper) GetBuilder(dest any, 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 ...any) (sql.Result, error) {
|
||||
query = w.Tx.Rebind(query)
|
||||
|
||||
return w.ExecRaw(query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxTxWrapper) ExecNoTimeout(query string, args ...any) (sql.Result, error) {
|
||||
query = w.Tx.Rebind(query)
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
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 ...any) (sql.Result, error) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.Tx.ExecContext(ctx, query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxTxWrapper) NamedExec(query string, arg any) (sql.Result, error) {
|
||||
if w.Tx.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = namedParamRegex.ReplaceAllStringFunc(query, strings.ToLower)
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), arg)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.Tx.NamedExecContext(ctx, query, arg)
|
||||
}
|
||||
|
||||
func (w *sqlxTxWrapper) NamedQuery(query string, arg any) (*sqlx.Rows, error) {
|
||||
if w.Tx.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = namedParamRegex.ReplaceAllStringFunc(query, strings.ToLower)
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), arg)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
// There is no tx.NamedQueryContext support in the sqlx API. (https://github.com/jmoiron/sqlx/issues/447)
|
||||
// So we need to implement this ourselves.
|
||||
type result struct {
|
||||
rows *sqlx.Rows
|
||||
err error
|
||||
}
|
||||
|
||||
// Need to add a buffer of 1 to prevent goroutine leak.
|
||||
resChan := make(chan *result, 1)
|
||||
go func() {
|
||||
rows, err := w.Tx.NamedQuery(query, arg)
|
||||
resChan <- &result{
|
||||
rows: rows,
|
||||
err: err,
|
||||
}
|
||||
}()
|
||||
|
||||
// staticcheck fails to check that res gets re-assigned later.
|
||||
res := &result{} //nolint:staticcheck
|
||||
select {
|
||||
case res = <-resChan:
|
||||
case <-ctx.Done():
|
||||
res = &result{
|
||||
rows: nil,
|
||||
err: ctx.Err(),
|
||||
}
|
||||
}
|
||||
|
||||
return res.rows, res.err
|
||||
}
|
||||
|
||||
func (w *sqlxTxWrapper) QueryRowX(query string, args ...any) *sqlx.Row {
|
||||
query = w.Tx.Rebind(query)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.Tx.QueryRowxContext(ctx, query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxTxWrapper) QueryX(query string, args ...any) (*sqlx.Rows, error) {
|
||||
query = w.Tx.Rebind(query)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.Tx.QueryxContext(ctx, query, args)
|
||||
}
|
||||
|
||||
func (w *sqlxTxWrapper) Select(dest any, query string, args ...any) error {
|
||||
query = w.Tx.Rebind(query)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), w.queryTimeout)
|
||||
defer cancel()
|
||||
|
||||
if w.trace {
|
||||
defer func(then time.Time) {
|
||||
printArgs(query, time.Since(then), args)
|
||||
}(time.Now())
|
||||
}
|
||||
|
||||
return w.Tx.SelectContext(ctx, dest, query, args...)
|
||||
}
|
||||
|
||||
func (w *sqlxTxWrapper) SelectBuilder(dest any, 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,
|
||||
// but we ignore it for now until someone complains.
|
||||
if unicode.IsSpace(r) && r != ' ' {
|
||||
return -1
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func printArgs(query string, dur time.Duration, args ...any) {
|
||||
query = strings.Map(removeSpace, query)
|
||||
fields := make([]mlog.Field, 0, len(args)+1)
|
||||
fields = append(fields, mlog.Duration("duration", dur))
|
||||
for i, arg := range args {
|
||||
fields = append(fields, mlog.Any("arg"+strconv.Itoa(i), arg))
|
||||
}
|
||||
mlog.Debug(query, fields...)
|
||||
}
|
||||
87
server/channels/store/sqlstore/sqlx_wrapper_test.go
Обычный файл
87
server/channels/store/sqlstore/sqlx_wrapper_test.go
Обычный файл
@@ -0,0 +1,87 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
)
|
||||
|
||||
func TestSqlX(t *testing.T) {
|
||||
t.Run("NamedQuery", func(t *testing.T) {
|
||||
testDrivers := []string{
|
||||
model.DatabaseDriverPostgres,
|
||||
model.DatabaseDriverMysql,
|
||||
}
|
||||
|
||||
for _, driver := range testDrivers {
|
||||
settings, err := makeSqlSettings(driver)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
*settings.QueryTimeout = 1
|
||||
store := &SqlStore{
|
||||
rrCounter: 0,
|
||||
srCounter: 0,
|
||||
settings: settings,
|
||||
}
|
||||
|
||||
store.initConnection()
|
||||
|
||||
defer store.Close()
|
||||
|
||||
tx, err := store.GetMasterX().Beginx()
|
||||
require.NoError(t, err)
|
||||
|
||||
var query string
|
||||
if store.DriverName() == model.DatabaseDriverMysql {
|
||||
query = `SELECT SLEEP(:Timeout);`
|
||||
} else if store.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = `SELECT pg_sleep(:timeout);`
|
||||
}
|
||||
arg := struct{ Timeout int }{Timeout: 2}
|
||||
_, err = tx.NamedQuery(query, arg)
|
||||
require.Equal(t, context.DeadlineExceeded, err)
|
||||
require.NoError(t, tx.Commit())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("NamedParse", func(t *testing.T) {
|
||||
queries := []struct {
|
||||
in string
|
||||
out string
|
||||
}{
|
||||
{
|
||||
in: `SELECT pg_sleep(:Timeout)`,
|
||||
out: `SELECT pg_sleep(:timeout)`,
|
||||
},
|
||||
{
|
||||
in: `SELECT u.Username FROM Bots
|
||||
LIMIT
|
||||
:Limit
|
||||
OFFSET
|
||||
:Offset`,
|
||||
out: `SELECT u.Username FROM Bots
|
||||
LIMIT
|
||||
:limit
|
||||
OFFSET
|
||||
:offset`,
|
||||
},
|
||||
{
|
||||
in: `UPDATE OAuthAccessData SET Token =:Token`,
|
||||
out: `UPDATE OAuthAccessData SET Token =:token`,
|
||||
},
|
||||
}
|
||||
for _, q := range queries {
|
||||
out := namedParamRegex.ReplaceAllStringFunc(q.in, strings.ToLower)
|
||||
assert.Equal(t, q.out, out)
|
||||
}
|
||||
})
|
||||
}
|
||||
229
server/channels/store/sqlstore/status_store.go
Обычный файл
229
server/channels/store/sqlstore/status_store.go
Обычный файл
@@ -0,0 +1,229 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlStatusStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlStatusStore(sqlStore *SqlStore) store.StatusStore {
|
||||
return &SqlStatusStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlStatusStore) SaveOrUpdate(st *model.Status) error {
|
||||
query := s.getQueryBuilder().
|
||||
Insert("Status").
|
||||
Columns("UserId", "Status", "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 = ?",
|
||||
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 = ?",
|
||||
st.Status, st.Manual, st.LastActivityAt, st.DNDEndTime, st.PrevStatus))
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "status_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to upsert Status")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlStatusStore) Get(userId string) (*model.Status, error) {
|
||||
var status model.Status
|
||||
|
||||
if err := s.GetReplicaX().Get(&status, "SELECT * FROM Status WHERE UserId = ?", userId); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Status", fmt.Sprintf("userId=%s", userId))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get Status with userId=%s", userId)
|
||||
}
|
||||
return &status, nil
|
||||
}
|
||||
|
||||
func (s SqlStatusStore) GetByIds(userIds []string) ([]*model.Status, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("UserId, Status, Manual, LastActivityAt").
|
||||
From("Status").
|
||||
Where(sq.Eq{"UserId": userIds})
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "status_tosql")
|
||||
}
|
||||
rows, err := s.GetReplicaX().DB.Query(queryString, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Statuses")
|
||||
}
|
||||
statuses := []*model.Status{}
|
||||
defer rows.Close()
|
||||
for rows.Next() {
|
||||
var status model.Status
|
||||
if err = rows.Scan(&status.UserId, &status.Status, &status.Manual, &status.LastActivityAt); err != nil {
|
||||
return nil, errors.Wrap(err, "unable to scan from rows")
|
||||
}
|
||||
statuses = append(statuses, &status)
|
||||
}
|
||||
if err = rows.Err(); err != nil {
|
||||
return nil, errors.Wrap(err, "failed while iterating over rows")
|
||||
}
|
||||
|
||||
return statuses, nil
|
||||
}
|
||||
|
||||
// MySQL doesn't have support for RETURNING clause, so we use a transaction to get the updated rows.
|
||||
func (s SqlStatusStore) updateExpiredStatuses(t *sqlxTxWrapper) ([]*model.Status, error) {
|
||||
statuses := []*model.Status{}
|
||||
currUnixTime := time.Now().UTC().Unix()
|
||||
selectQuery, selectParams, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Status").
|
||||
Where(
|
||||
sq.And{
|
||||
sq.Eq{"Status": model.StatusDnd},
|
||||
sq.Gt{"DNDEndTime": 0},
|
||||
sq.LtOrEq{"DNDEndTime": currUnixTime},
|
||||
},
|
||||
).ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "status_tosql")
|
||||
}
|
||||
err = t.Select(&statuses, selectQuery, selectParams...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "updateExpiredStatusesT: failed to get expired dnd statuses")
|
||||
}
|
||||
updateQuery, args, err := s.getQueryBuilder().
|
||||
Update("Status").
|
||||
Where(
|
||||
sq.And{
|
||||
sq.Eq{"Status": model.StatusDnd},
|
||||
sq.Gt{"DNDEndTime": 0},
|
||||
sq.LtOrEq{"DNDEndTime": currUnixTime},
|
||||
},
|
||||
).
|
||||
Set("Status", sq.Expr("PrevStatus")).
|
||||
Set("PrevStatus", model.StatusDnd).
|
||||
Set("DNDEndTime", 0).
|
||||
Set("Manual", false).
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "status_tosql")
|
||||
}
|
||||
|
||||
if _, err := t.Exec(updateQuery, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "updateExpiredStatusesT: failed to update statuses")
|
||||
}
|
||||
|
||||
return statuses, nil
|
||||
}
|
||||
|
||||
func (s SqlStatusStore) UpdateExpiredDNDStatuses() (_ []*model.Status, err error) {
|
||||
if s.DriverName() == model.DatabaseDriverMysql {
|
||||
transaction, terr := s.GetMasterX().Beginx()
|
||||
if terr != nil {
|
||||
return nil, errors.Wrap(terr, "UpdateExpiredDNDStatuses: begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &terr)
|
||||
statuses, terr := s.updateExpiredStatuses(transaction)
|
||||
if terr != nil {
|
||||
return nil, errors.Wrap(terr, "UpdateExpiredDNDStatuses: updateExpiredDNDStatusesT")
|
||||
}
|
||||
if terr = transaction.Commit(); terr != nil {
|
||||
return nil, errors.Wrap(terr, "UpdateExpiredDNDStatuses: commit_transaction")
|
||||
}
|
||||
|
||||
for _, status := range statuses {
|
||||
status.Status = status.PrevStatus
|
||||
status.PrevStatus = model.StatusDnd
|
||||
status.DNDEndTime = 0
|
||||
status.Manual = false
|
||||
}
|
||||
|
||||
return statuses, nil
|
||||
}
|
||||
|
||||
queryString, args, err := s.getQueryBuilder().
|
||||
Update("Status").
|
||||
Where(
|
||||
sq.And{
|
||||
sq.Eq{"Status": model.StatusDnd},
|
||||
sq.Gt{"DNDEndTime": 0},
|
||||
sq.LtOrEq{"DNDEndTime": time.Now().UTC().Unix()},
|
||||
},
|
||||
).
|
||||
Set("Status", sq.Expr("PrevStatus")).
|
||||
Set("PrevStatus", model.StatusDnd).
|
||||
Set("DNDEndTime", 0).
|
||||
Set("Manual", false).
|
||||
Suffix("RETURNING *").
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "status_tosql")
|
||||
}
|
||||
|
||||
rows, err := s.GetMasterX().Query(queryString, args...)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find Statuses")
|
||||
}
|
||||
defer rows.Close()
|
||||
statuses := []*model.Status{}
|
||||
for rows.Next() {
|
||||
var status model.Status
|
||||
if err = rows.Scan(&status.UserId, &status.Status, &status.Manual, &status.LastActivityAt,
|
||||
&status.DNDEndTime, &status.PrevStatus); err != nil {
|
||||
return nil, errors.Wrap(err, "unable to scan from rows")
|
||||
}
|
||||
statuses = append(statuses, &status)
|
||||
}
|
||||
if err = rows.Err(); err != nil {
|
||||
return nil, errors.Wrap(err, "failed while iterating over rows")
|
||||
}
|
||||
|
||||
return statuses, nil
|
||||
}
|
||||
|
||||
func (s SqlStatusStore) ResetAll() error {
|
||||
if _, err := s.GetMasterX().Exec("UPDATE Status SET Status = ? WHERE Manual = false", model.StatusOffline); err != nil {
|
||||
return errors.Wrap(err, "failed to update Statuses")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlStatusStore) GetTotalActiveUsersCount() (int64, error) {
|
||||
time := model.GetMillis() - (1000 * 60 * 60 * 24)
|
||||
var count int64
|
||||
err := s.GetReplicaX().Get(&count, "SELECT COUNT(UserId) FROM Status WHERE LastActivityAt > ?", time)
|
||||
if err != nil {
|
||||
return count, errors.Wrap(err, "failed to count active users")
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s SqlStatusStore) UpdateLastActivityAt(userId string, lastActivityAt int64) error {
|
||||
if _, err := s.GetMasterX().Exec("UPDATE Status SET LastActivityAt = ? WHERE UserId = ?", lastActivityAt, userId); err != nil {
|
||||
return errors.Wrapf(err, "failed to update last activity for userId=%s", userId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/status_store_test.go
Обычный файл
14
server/channels/store/sqlstore/status_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestStatusStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestStatusStore)
|
||||
}
|
||||
1286
server/channels/store/sqlstore/store.go
Обычный файл
1286
server/channels/store/sqlstore/store.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
933
server/channels/store/sqlstore/store_test.go
Обычный файл
933
server/channels/store/sqlstore/store_test.go
Обычный файл
@@ -0,0 +1,933 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/go-sql-driver/mysql"
|
||||
"github.com/lib/pq"
|
||||
"github.com/pkg/errors"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/plugin/plugintest/mock"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/db"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/einterfaces/mocks"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/searchtest"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
type storeType struct {
|
||||
Name string
|
||||
SqlSettings *model.SqlSettings
|
||||
SqlStore *SqlStore
|
||||
Store store.Store
|
||||
}
|
||||
|
||||
var storeTypes []*storeType
|
||||
|
||||
func newStoreType(name, driver string) *storeType {
|
||||
return &storeType{
|
||||
Name: name,
|
||||
SqlSettings: storetest.MakeSqlSettings(driver, false),
|
||||
}
|
||||
}
|
||||
|
||||
func StoreTest(t *testing.T, f func(*testing.T, store.Store)) {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
tearDownStores()
|
||||
panic(err)
|
||||
}
|
||||
}()
|
||||
for _, st := range storeTypes {
|
||||
st := st
|
||||
t.Run(st.Name, func(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.SkipNow()
|
||||
}
|
||||
f(t, st.Store)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func StoreTestWithSearchTestEngine(t *testing.T, f func(*testing.T, store.Store, *searchtest.SearchTestEngine)) {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
tearDownStores()
|
||||
panic(err)
|
||||
}
|
||||
}()
|
||||
|
||||
for _, st := range storeTypes {
|
||||
st := st
|
||||
searchTestEngine := &searchtest.SearchTestEngine{
|
||||
Driver: *st.SqlSettings.DriverName,
|
||||
}
|
||||
|
||||
t.Run(st.Name, func(t *testing.T) { f(t, st.Store, searchTestEngine) })
|
||||
}
|
||||
}
|
||||
|
||||
func StoreTestWithSqlStore(t *testing.T, f func(*testing.T, store.Store, storetest.SqlStore)) {
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
tearDownStores()
|
||||
panic(err)
|
||||
}
|
||||
}()
|
||||
for _, st := range storeTypes {
|
||||
st := st
|
||||
t.Run(st.Name, func(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.SkipNow()
|
||||
}
|
||||
f(t, st.Store, &StoreTestWrapper{st.SqlStore})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func initStores() {
|
||||
if testing.Short() {
|
||||
return
|
||||
}
|
||||
// In CI, we already run the entire test suite for both mysql and postgres in parallel.
|
||||
// So we just run the tests for the current database set.
|
||||
if os.Getenv("IS_CI") == "true" {
|
||||
switch os.Getenv("MM_SQLSETTINGS_DRIVERNAME") {
|
||||
case "mysql":
|
||||
storeTypes = append(storeTypes, newStoreType("MySQL", model.DatabaseDriverMysql))
|
||||
case "postgres":
|
||||
storeTypes = append(storeTypes, newStoreType("PostgreSQL", model.DatabaseDriverPostgres))
|
||||
}
|
||||
} else {
|
||||
storeTypes = append(storeTypes,
|
||||
newStoreType("MySQL", model.DatabaseDriverMysql),
|
||||
newStoreType("PostgreSQL", model.DatabaseDriverPostgres),
|
||||
)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
if err := recover(); err != nil {
|
||||
tearDownStores()
|
||||
panic(err)
|
||||
}
|
||||
}()
|
||||
var wg sync.WaitGroup
|
||||
for _, st := range storeTypes {
|
||||
st := st
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
st.SqlStore = New(*st.SqlSettings, nil)
|
||||
st.Store = st.SqlStore
|
||||
st.Store.DropAllTables()
|
||||
st.Store.MarkSystemRanUnitTests()
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
var tearDownStoresOnce sync.Once
|
||||
|
||||
func tearDownStores() {
|
||||
if testing.Short() {
|
||||
return
|
||||
}
|
||||
tearDownStoresOnce.Do(func() {
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(len(storeTypes))
|
||||
for _, st := range storeTypes {
|
||||
st := st
|
||||
go func() {
|
||||
if st.Store != nil {
|
||||
st.Store.Close()
|
||||
}
|
||||
if st.SqlSettings != nil {
|
||||
storetest.CleanupSqlSettings(st.SqlSettings)
|
||||
}
|
||||
wg.Done()
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
})
|
||||
}
|
||||
|
||||
// This test was used to consistently reproduce the race
|
||||
// before the fix in MM-28397.
|
||||
// Keeping it here to help avoiding future regressions.
|
||||
func TestStoreLicenseRace(t *testing.T) {
|
||||
settings, err := makeSqlSettings(model.DatabaseDriverPostgres)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
|
||||
store := New(*settings, nil)
|
||||
defer func() {
|
||||
store.Close()
|
||||
storetest.CleanupSqlSettings(settings)
|
||||
}()
|
||||
|
||||
wg := sync.WaitGroup{}
|
||||
wg.Add(3)
|
||||
|
||||
go func() {
|
||||
store.UpdateLicense(&model.License{})
|
||||
wg.Done()
|
||||
}()
|
||||
|
||||
go func() {
|
||||
store.GetReplicaX()
|
||||
wg.Done()
|
||||
}()
|
||||
|
||||
go func() {
|
||||
store.GetSearchReplicaX()
|
||||
wg.Done()
|
||||
}()
|
||||
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func TestGetReplica(t *testing.T) {
|
||||
t.Parallel()
|
||||
testCases := []struct {
|
||||
Description string
|
||||
DataSourceReplicaNum int
|
||||
DataSourceSearchReplicaNum int
|
||||
}{
|
||||
{
|
||||
"no replicas",
|
||||
0,
|
||||
0,
|
||||
},
|
||||
{
|
||||
"one source replica",
|
||||
1,
|
||||
0,
|
||||
},
|
||||
{
|
||||
"multiple source replicas",
|
||||
3,
|
||||
0,
|
||||
},
|
||||
{
|
||||
"one source search replica",
|
||||
0,
|
||||
1,
|
||||
},
|
||||
{
|
||||
"multiple source search replicas",
|
||||
0,
|
||||
3,
|
||||
},
|
||||
{
|
||||
"one source replica, one source search replica",
|
||||
1,
|
||||
1,
|
||||
},
|
||||
{
|
||||
"one source replica, multiple source search replicas",
|
||||
1,
|
||||
3,
|
||||
},
|
||||
{
|
||||
"multiple source replica, one source search replica",
|
||||
3,
|
||||
1,
|
||||
},
|
||||
{
|
||||
"multiple source replica, multiple source search replicas",
|
||||
3,
|
||||
3,
|
||||
},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
testCase := testCase
|
||||
t.Run(testCase.Description+" with license", func(t *testing.T) {
|
||||
settings, err := makeSqlSettings(model.DatabaseDriverPostgres)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
|
||||
dataSourceReplicas := []string{}
|
||||
dataSourceSearchReplicas := []string{}
|
||||
for i := 0; i < testCase.DataSourceReplicaNum; i++ {
|
||||
dataSourceReplicas = append(dataSourceReplicas, *settings.DataSource)
|
||||
}
|
||||
for i := 0; i < testCase.DataSourceSearchReplicaNum; i++ {
|
||||
dataSourceSearchReplicas = append(dataSourceSearchReplicas, *settings.DataSource)
|
||||
}
|
||||
|
||||
settings.DataSourceReplicas = dataSourceReplicas
|
||||
settings.DataSourceSearchReplicas = dataSourceSearchReplicas
|
||||
store := New(*settings, nil)
|
||||
defer func() {
|
||||
store.Close()
|
||||
storetest.CleanupSqlSettings(settings)
|
||||
}()
|
||||
|
||||
store.UpdateLicense(&model.License{})
|
||||
|
||||
replicas := make(map[*sqlxDBWrapper]bool)
|
||||
for i := 0; i < 5; i++ {
|
||||
replicas[store.GetReplicaX()] = true
|
||||
}
|
||||
|
||||
searchReplicas := make(map[*sqlxDBWrapper]bool)
|
||||
for i := 0; i < 5; i++ {
|
||||
searchReplicas[store.GetSearchReplicaX()] = true
|
||||
}
|
||||
|
||||
if testCase.DataSourceReplicaNum > 0 {
|
||||
// If replicas were defined, ensure none are the master.
|
||||
assert.Len(t, replicas, testCase.DataSourceReplicaNum)
|
||||
|
||||
for replica := range replicas {
|
||||
assert.NotSame(t, store.GetMasterX(), replica)
|
||||
}
|
||||
|
||||
} else if assert.Len(t, replicas, 1) {
|
||||
// Otherwise ensure the replicas contains only the master.
|
||||
for replica := range replicas {
|
||||
assert.Same(t, store.GetMasterX(), replica)
|
||||
}
|
||||
}
|
||||
|
||||
if testCase.DataSourceSearchReplicaNum > 0 {
|
||||
// If search replicas were defined, ensure none are the master nor the replicas.
|
||||
assert.Len(t, searchReplicas, testCase.DataSourceSearchReplicaNum)
|
||||
|
||||
for searchReplica := range searchReplicas {
|
||||
assert.NotSame(t, store.GetMasterX(), searchReplica)
|
||||
for replica := range replicas {
|
||||
assert.NotSame(t, searchReplica, replica)
|
||||
}
|
||||
}
|
||||
} else if testCase.DataSourceReplicaNum > 0 {
|
||||
assert.Equal(t, len(replicas), len(searchReplicas))
|
||||
for k := range replicas {
|
||||
assert.True(t, searchReplicas[k])
|
||||
}
|
||||
} else if testCase.DataSourceReplicaNum == 0 && assert.Len(t, searchReplicas, 1) {
|
||||
// Otherwise ensure the search replicas contains the master.
|
||||
for searchReplica := range searchReplicas {
|
||||
assert.Same(t, store.GetMasterX(), searchReplica)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
t.Run(testCase.Description+" without license", func(t *testing.T) {
|
||||
settings, err := makeSqlSettings(model.DatabaseDriverPostgres)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
|
||||
dataSourceReplicas := []string{}
|
||||
dataSourceSearchReplicas := []string{}
|
||||
for i := 0; i < testCase.DataSourceReplicaNum; i++ {
|
||||
dataSourceReplicas = append(dataSourceReplicas, *settings.DataSource)
|
||||
}
|
||||
for i := 0; i < testCase.DataSourceSearchReplicaNum; i++ {
|
||||
dataSourceSearchReplicas = append(dataSourceSearchReplicas, *settings.DataSource)
|
||||
}
|
||||
|
||||
settings.DataSourceReplicas = dataSourceReplicas
|
||||
settings.DataSourceSearchReplicas = dataSourceSearchReplicas
|
||||
store := New(*settings, nil)
|
||||
defer func() {
|
||||
store.Close()
|
||||
storetest.CleanupSqlSettings(settings)
|
||||
}()
|
||||
|
||||
replicas := make(map[*sqlxDBWrapper]bool)
|
||||
for i := 0; i < 5; i++ {
|
||||
replicas[store.GetReplicaX()] = true
|
||||
}
|
||||
|
||||
searchReplicas := make(map[*sqlxDBWrapper]bool)
|
||||
for i := 0; i < 5; i++ {
|
||||
searchReplicas[store.GetSearchReplicaX()] = true
|
||||
}
|
||||
|
||||
if testCase.DataSourceReplicaNum > 0 {
|
||||
// If replicas were defined, ensure none are the master.
|
||||
assert.Len(t, replicas, 1)
|
||||
|
||||
for replica := range replicas {
|
||||
assert.Same(t, store.GetMasterX(), replica)
|
||||
}
|
||||
|
||||
} else if assert.Len(t, replicas, 1) {
|
||||
// Otherwise ensure the replicas contains only the master.
|
||||
for replica := range replicas {
|
||||
assert.Same(t, store.GetMasterX(), replica)
|
||||
}
|
||||
}
|
||||
|
||||
if testCase.DataSourceSearchReplicaNum > 0 {
|
||||
// If search replicas were defined, ensure none are the master nor the replicas.
|
||||
assert.Len(t, searchReplicas, 1)
|
||||
|
||||
for searchReplica := range searchReplicas {
|
||||
assert.Same(t, store.GetMasterX(), searchReplica)
|
||||
}
|
||||
|
||||
} else if testCase.DataSourceReplicaNum > 0 {
|
||||
assert.Equal(t, len(replicas), len(searchReplicas))
|
||||
for k := range replicas {
|
||||
assert.True(t, searchReplicas[k])
|
||||
}
|
||||
} else if assert.Len(t, searchReplicas, 1) {
|
||||
// Otherwise ensure the search replicas contains the master.
|
||||
for searchReplica := range searchReplicas {
|
||||
assert.Same(t, store.GetMasterX(), searchReplica)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDbVersion(t *testing.T) {
|
||||
testDrivers := []string{
|
||||
model.DatabaseDriverPostgres,
|
||||
model.DatabaseDriverMysql,
|
||||
}
|
||||
|
||||
for _, d := range testDrivers {
|
||||
driver := d
|
||||
t.Run("Should return db version for "+driver, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
settings, err := makeSqlSettings(driver)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
|
||||
store := New(*settings, nil)
|
||||
|
||||
version, err := store.GetDbVersion(false)
|
||||
require.NoError(t, err)
|
||||
require.Regexp(t, regexp.MustCompile(`\d+\.\d+(\.\d+)?`), version)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureMinimumDBVersion(t *testing.T) {
|
||||
tests := []struct {
|
||||
driver string
|
||||
ver string
|
||||
ok bool
|
||||
err string
|
||||
}{
|
||||
{
|
||||
driver: model.DatabaseDriverPostgres,
|
||||
ver: "100001",
|
||||
ok: true,
|
||||
err: "",
|
||||
},
|
||||
{
|
||||
driver: model.DatabaseDriverPostgres,
|
||||
ver: "90603",
|
||||
ok: false,
|
||||
err: "minimum Postgres version requirements not met",
|
||||
},
|
||||
{
|
||||
driver: model.DatabaseDriverPostgres,
|
||||
ver: "12.34.1",
|
||||
ok: false,
|
||||
err: "cannot parse DB version",
|
||||
},
|
||||
{
|
||||
driver: model.DatabaseDriverMysql,
|
||||
ver: "10.4.5-MariaDB",
|
||||
ok: true,
|
||||
err: "",
|
||||
},
|
||||
{
|
||||
driver: model.DatabaseDriverMysql,
|
||||
ver: "5.6.99-test",
|
||||
ok: false,
|
||||
err: "minimum MySQL version requirements not met",
|
||||
},
|
||||
{
|
||||
driver: model.DatabaseDriverMysql,
|
||||
ver: "34-55.12",
|
||||
ok: false,
|
||||
err: "cannot parse MySQL DB version",
|
||||
},
|
||||
{
|
||||
driver: model.DatabaseDriverMysql,
|
||||
ver: "8.0.0-log",
|
||||
ok: true,
|
||||
err: "",
|
||||
},
|
||||
}
|
||||
|
||||
pg := model.DatabaseDriverPostgres
|
||||
pgSettings := &model.SqlSettings{
|
||||
DriverName: &pg,
|
||||
}
|
||||
my := model.DatabaseDriverMysql
|
||||
mySettings := &model.SqlSettings{
|
||||
DriverName: &my,
|
||||
}
|
||||
for _, tc := range tests {
|
||||
store := &SqlStore{}
|
||||
switch tc.driver {
|
||||
case pg:
|
||||
store.settings = pgSettings
|
||||
case my:
|
||||
store.settings = mySettings
|
||||
}
|
||||
ok, err := store.ensureMinimumDBVersion(tc.ver)
|
||||
assert.Equal(t, tc.ok, ok)
|
||||
if tc.err != "" {
|
||||
assert.Contains(t, err.Error(), tc.err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBinaryParamEnabled(t *testing.T) {
|
||||
tests := []struct {
|
||||
store SqlStore
|
||||
expected bool
|
||||
}{
|
||||
{
|
||||
store: SqlStore{
|
||||
settings: &model.SqlSettings{
|
||||
DriverName: model.NewString(model.DatabaseDriverPostgres),
|
||||
DataSource: model.NewString("postgres://mmuser:mostest@localhost/loadtest?sslmode=disable\u0026binary_parameters=yes"),
|
||||
},
|
||||
},
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
store: SqlStore{
|
||||
settings: &model.SqlSettings{
|
||||
DriverName: model.NewString(model.DatabaseDriverMysql),
|
||||
DataSource: model.NewString("postgres://mmuser:mostest@localhost/loadtest?sslmode=disable\u0026binary_parameters=yes"),
|
||||
},
|
||||
},
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
store: SqlStore{
|
||||
settings: &model.SqlSettings{
|
||||
DriverName: model.NewString(model.DatabaseDriverPostgres),
|
||||
DataSource: model.NewString("postgres://mmuser:mostest@localhost/loadtest?sslmode=disable&binary_parameters=yes"),
|
||||
},
|
||||
},
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
store: SqlStore{
|
||||
settings: &model.SqlSettings{
|
||||
DriverName: model.NewString(model.DatabaseDriverPostgres),
|
||||
DataSource: model.NewString("postgres://mmuser:mostest@localhost/loadtest?sslmode=disable"),
|
||||
},
|
||||
},
|
||||
expected: false,
|
||||
},
|
||||
}
|
||||
|
||||
for i := range tests {
|
||||
ok, err := tests[i].store.computeBinaryParam()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tests[i].expected, ok)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
func TestUpAndDownMigrations(t *testing.T) {
|
||||
testDrivers := []string{
|
||||
model.DatabaseDriverPostgres,
|
||||
model.DatabaseDriverMysql,
|
||||
}
|
||||
|
||||
for _, driver := range testDrivers {
|
||||
t.Run("Should be reversible for "+driver, func(t *testing.T) {
|
||||
settings, err := makeSqlSettings(driver)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
store := New(*settings, nil)
|
||||
defer store.Close()
|
||||
|
||||
err = store.migrate(migrationsDirectionDown)
|
||||
assert.NoError(t, err, "downing migrations should not error")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetAllConns(t *testing.T) {
|
||||
t.Parallel()
|
||||
testCases := []struct {
|
||||
Description string
|
||||
DataSourceReplicaNum int
|
||||
DataSourceSearchReplicaNum int
|
||||
ExpectedNumConnections int
|
||||
}{
|
||||
{
|
||||
"no replicas",
|
||||
0,
|
||||
0,
|
||||
1,
|
||||
},
|
||||
{
|
||||
"one source replica",
|
||||
1,
|
||||
0,
|
||||
2,
|
||||
},
|
||||
{
|
||||
"multiple source replicas",
|
||||
3,
|
||||
0,
|
||||
4,
|
||||
},
|
||||
{
|
||||
"one source search replica",
|
||||
0,
|
||||
1,
|
||||
1,
|
||||
},
|
||||
{
|
||||
"multiple source search replicas",
|
||||
0,
|
||||
3,
|
||||
1,
|
||||
},
|
||||
{
|
||||
"one source replica, one source search replica",
|
||||
1,
|
||||
1,
|
||||
2,
|
||||
},
|
||||
{
|
||||
"one source replica, multiple source search replicas",
|
||||
1,
|
||||
3,
|
||||
2,
|
||||
},
|
||||
{
|
||||
"multiple source replica, one source search replica",
|
||||
3,
|
||||
1,
|
||||
4,
|
||||
},
|
||||
{
|
||||
"multiple source replica, multiple source search replicas",
|
||||
3,
|
||||
3,
|
||||
4,
|
||||
},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
testCase := testCase
|
||||
t.Run(testCase.Description, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
settings, err := makeSqlSettings(model.DatabaseDriverPostgres)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
dataSourceReplicas := []string{}
|
||||
dataSourceSearchReplicas := []string{}
|
||||
for i := 0; i < testCase.DataSourceReplicaNum; i++ {
|
||||
dataSourceReplicas = append(dataSourceReplicas, *settings.DataSource)
|
||||
}
|
||||
for i := 0; i < testCase.DataSourceSearchReplicaNum; i++ {
|
||||
dataSourceSearchReplicas = append(dataSourceSearchReplicas, *settings.DataSource)
|
||||
}
|
||||
|
||||
settings.DataSourceReplicas = dataSourceReplicas
|
||||
settings.DataSourceSearchReplicas = dataSourceSearchReplicas
|
||||
store := New(*settings, nil)
|
||||
defer func() {
|
||||
store.Close()
|
||||
storetest.CleanupSqlSettings(settings)
|
||||
}()
|
||||
|
||||
assert.Len(t, store.GetAllConns(), testCase.ExpectedNumConnections)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsDuplicate(t *testing.T) {
|
||||
testErrors := map[error]bool{
|
||||
&pq.Error{Code: "42P06"}: false,
|
||||
&pq.Error{Code: PGDupTableErrorCode}: true,
|
||||
&mysql.MySQLError{Number: uint16(1000)}: false,
|
||||
&mysql.MySQLError{Number: MySQLDupTableErrorCode}: true,
|
||||
errors.New("Random error"): false,
|
||||
}
|
||||
|
||||
for e, b := range testErrors {
|
||||
err := e
|
||||
expected := b
|
||||
t.Run(fmt.Sprintf("Should return %t for %s", expected, err.Error()), func(t *testing.T) {
|
||||
t.Parallel()
|
||||
assert.Equal(t, expected, IsDuplicate(err))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestVersionString(t *testing.T) {
|
||||
versions := []struct {
|
||||
input int
|
||||
driver string
|
||||
output string
|
||||
}{
|
||||
{
|
||||
input: 100000,
|
||||
driver: model.DatabaseDriverPostgres,
|
||||
output: "10.0",
|
||||
},
|
||||
{
|
||||
input: 90603,
|
||||
driver: model.DatabaseDriverPostgres,
|
||||
output: "9.603",
|
||||
},
|
||||
{
|
||||
input: 120005,
|
||||
driver: model.DatabaseDriverPostgres,
|
||||
output: "12.5",
|
||||
},
|
||||
{
|
||||
input: 5708,
|
||||
driver: model.DatabaseDriverMysql,
|
||||
output: "5.7.8",
|
||||
},
|
||||
{
|
||||
input: 8000,
|
||||
driver: model.DatabaseDriverMysql,
|
||||
output: "8.0.0",
|
||||
},
|
||||
}
|
||||
|
||||
for _, v := range versions {
|
||||
out := versionString(v.input, v.driver)
|
||||
assert.Equal(t, v.output, out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplicaLagQuery(t *testing.T) {
|
||||
testDrivers := []string{
|
||||
model.DatabaseDriverPostgres,
|
||||
model.DatabaseDriverMysql,
|
||||
}
|
||||
|
||||
for _, driver := range testDrivers {
|
||||
t.Run(driver, func(t *testing.T) {
|
||||
settings, err := makeSqlSettings(driver)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
var query string
|
||||
var tableName string
|
||||
// Just any random query which returns a row in (string, int) format.
|
||||
switch driver {
|
||||
case model.DatabaseDriverPostgres:
|
||||
query = `SELECT relname, count(relname) FROM pg_class WHERE relname='posts' GROUP BY relname`
|
||||
tableName = "posts"
|
||||
case model.DatabaseDriverMysql:
|
||||
query = `SELECT table_name, count(table_name) FROM information_schema.tables WHERE table_name='Posts' and table_schema=Database() GROUP BY table_name`
|
||||
tableName = "Posts"
|
||||
}
|
||||
|
||||
settings.ReplicaLagSettings = []*model.ReplicaLagSettings{{
|
||||
DataSource: model.NewString(*settings.DataSource),
|
||||
QueryAbsoluteLag: model.NewString(query),
|
||||
QueryTimeLag: model.NewString(query),
|
||||
}}
|
||||
|
||||
mockMetrics := &mocks.MetricsInterface{}
|
||||
mockMetrics.On("SetReplicaLagAbsolute", tableName, float64(1))
|
||||
mockMetrics.On("SetReplicaLagTime", tableName, float64(1))
|
||||
mockMetrics.On("RegisterDBCollector", mock.AnythingOfType("*sql.DB"), "master")
|
||||
|
||||
store := &SqlStore{
|
||||
rrCounter: 0,
|
||||
srCounter: 0,
|
||||
settings: settings,
|
||||
metrics: mockMetrics,
|
||||
}
|
||||
|
||||
store.initConnection()
|
||||
store.stores.post = newSqlPostStore(store, mockMetrics)
|
||||
err = store.migrate(migrationsDirectionUp)
|
||||
require.NoError(t, err)
|
||||
|
||||
defer store.Close()
|
||||
|
||||
err = store.ReplicaLagAbs()
|
||||
require.NoError(t, err)
|
||||
err = store.ReplicaLagTime()
|
||||
require.NoError(t, err)
|
||||
mockMetrics.AssertExpectations(t)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
var errDriverMismatch = errors.New("database drivers mismatch")
|
||||
var errDriverUnsupported = errors.New("database driver not supported")
|
||||
|
||||
func makeSqlSettings(driver string) (*model.SqlSettings, error) {
|
||||
// When running under CI, only one database engine container is launched
|
||||
// so here we must error out if the requested driver doesn't match.
|
||||
if os.Getenv("IS_CI") == "true" {
|
||||
envDriver := os.Getenv("MM_SQLSETTINGS_DRIVERNAME")
|
||||
if envDriver != "" && envDriver != driver {
|
||||
return nil, errDriverMismatch
|
||||
}
|
||||
}
|
||||
|
||||
switch driver {
|
||||
case model.DatabaseDriverPostgres:
|
||||
return storetest.MakeSqlSettings(driver, false), nil
|
||||
case model.DatabaseDriverMysql:
|
||||
return storetest.MakeSqlSettings(driver, false), nil
|
||||
}
|
||||
|
||||
return nil, errDriverUnsupported
|
||||
}
|
||||
|
||||
func TestExecNoTimeout(t *testing.T) {
|
||||
StoreTest(t, func(t *testing.T, ss store.Store) {
|
||||
sqlStore := ss.(*SqlStore)
|
||||
var query string
|
||||
timeout := sqlStore.masterX.queryTimeout
|
||||
sqlStore.masterX.queryTimeout = 1
|
||||
defer func() {
|
||||
sqlStore.masterX.queryTimeout = timeout
|
||||
}()
|
||||
if sqlStore.DriverName() == model.DatabaseDriverMysql {
|
||||
query = `SELECT SLEEP(2);`
|
||||
} else if sqlStore.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = `SELECT pg_sleep(2);`
|
||||
}
|
||||
_, err := sqlStore.GetMasterX().ExecNoTimeout(query)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func TestMySQLReadTimeout(t *testing.T) {
|
||||
settings, err := makeSqlSettings(model.DatabaseDriverMysql)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
dataSource := *settings.DataSource
|
||||
config, err := mysql.ParseDSN(dataSource)
|
||||
require.NoError(t, err)
|
||||
|
||||
config.ReadTimeout = 1 * time.Second
|
||||
dataSource = config.FormatDSN()
|
||||
settings.DataSource = &dataSource
|
||||
|
||||
store := &SqlStore{
|
||||
settings: settings,
|
||||
}
|
||||
store.initConnection()
|
||||
defer store.Close()
|
||||
|
||||
_, err = store.GetMasterX().ExecNoTimeout(`SELECT SLEEP(3)`)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestGetDBSchemaVersion(t *testing.T) {
|
||||
testDrivers := []string{
|
||||
model.DatabaseDriverPostgres,
|
||||
model.DatabaseDriverMysql,
|
||||
}
|
||||
|
||||
assets := db.Assets()
|
||||
|
||||
for _, d := range testDrivers {
|
||||
driver := d
|
||||
t.Run("Should return latest version number of applied migrations for "+driver, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
settings, err := makeSqlSettings(driver)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
store := New(*settings, nil)
|
||||
|
||||
assetsList, err := assets.ReadDir(filepath.Join("migrations", driver))
|
||||
require.NoError(t, err)
|
||||
|
||||
var assetNamesForDriver []string
|
||||
for _, entry := range assetsList {
|
||||
assetNamesForDriver = append(assetNamesForDriver, entry.Name())
|
||||
}
|
||||
sort.Strings(assetNamesForDriver)
|
||||
|
||||
require.NotEmpty(t, assetNamesForDriver)
|
||||
lastMigration := assetNamesForDriver[len(assetNamesForDriver)-1]
|
||||
expectedVersion := strings.Split(lastMigration, "_")[0]
|
||||
|
||||
version, err := store.GetDBSchemaVersion()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, expectedVersion, fmt.Sprintf("%06d", version))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetAppliedMigrations(t *testing.T) {
|
||||
testDrivers := []string{
|
||||
model.DatabaseDriverPostgres,
|
||||
model.DatabaseDriverMysql,
|
||||
}
|
||||
|
||||
assets := db.Assets()
|
||||
|
||||
for _, d := range testDrivers {
|
||||
driver := d
|
||||
t.Run("Should return db applied migrations for "+driver, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
settings, err := makeSqlSettings(driver)
|
||||
if err != nil {
|
||||
t.Skip(err)
|
||||
}
|
||||
store := New(*settings, nil)
|
||||
|
||||
assetsList, err := assets.ReadDir(filepath.Join("migrations", driver))
|
||||
require.NoError(t, err)
|
||||
|
||||
var migrationsFromFiles []model.AppliedMigration
|
||||
for _, entry := range assetsList {
|
||||
if strings.HasSuffix(entry.Name(), ".up.sql") {
|
||||
versionString := strings.Split(entry.Name(), "_")[0]
|
||||
version, vErr := strconv.Atoi(versionString)
|
||||
require.NoError(t, vErr)
|
||||
|
||||
name := strings.TrimSuffix(strings.TrimLeft(entry.Name(), versionString+"_"), ".up.sql")
|
||||
|
||||
migrationsFromFiles = append(migrationsFromFiles, model.AppliedMigration{
|
||||
Version: version,
|
||||
Name: name,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
require.NotEmpty(t, migrationsFromFiles)
|
||||
|
||||
migrations, err := store.GetAppliedMigrations()
|
||||
require.NoError(t, err)
|
||||
require.ElementsMatch(t, migrationsFromFiles, migrations)
|
||||
})
|
||||
}
|
||||
}
|
||||
156
server/channels/store/sqlstore/system_store.go
Обычный файл
156
server/channels/store/sqlstore/system_store.go
Обычный файл
@@ -0,0 +1,156 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/utils"
|
||||
)
|
||||
|
||||
type SqlSystemStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlSystemStore(sqlStore *SqlStore) store.SystemStore {
|
||||
return &SqlSystemStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlSystemStore) Save(system *model.System) error {
|
||||
query := "INSERT INTO Systems (Name, Value) VALUES (:Name, :Value)"
|
||||
if _, err := s.GetMasterX().NamedExec(query, system); err != nil {
|
||||
return errors.Wrapf(err, "failed to save system property with name=%s", system.Name)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlSystemStore) SaveOrUpdate(system *model.System) error {
|
||||
query := s.getQueryBuilder().
|
||||
Insert("Systems").
|
||||
Columns("Name", "Value").
|
||||
Values(system.Name, system.Value)
|
||||
|
||||
if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = query.SuffixExpr(sq.Expr("ON DUPLICATE KEY UPDATE Value = ?", system.Value))
|
||||
} else {
|
||||
query = query.SuffixExpr(sq.Expr("ON CONFLICT (name) DO UPDATE SET Value = ?", system.Value))
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "system_tosql")
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().Exec(queryString, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to upsert system property")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlSystemStore) SaveOrUpdateWithWarnMetricHandling(system *model.System) error {
|
||||
if err := s.SaveOrUpdate(system); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if strings.HasPrefix(system.Name, model.WarnMetricStatusStorePrefix) &&
|
||||
(system.Value == model.WarnMetricStatusRunonce || system.Value == model.WarnMetricStatusLimitReached) {
|
||||
if err := s.SaveOrUpdate(&model.System{
|
||||
Name: model.SystemWarnMetricLastRunTimestampKey,
|
||||
Value: strconv.FormatInt(utils.MillisFromTime(time.Now()), 10),
|
||||
}); err != nil {
|
||||
return errors.Wrapf(err, "failed to save system property with name=%s", model.SystemWarnMetricLastRunTimestampKey)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlSystemStore) Update(system *model.System) error {
|
||||
query := "UPDATE Systems SET Value=:Value WHERE Name=:Name"
|
||||
if _, err := s.GetMasterX().NamedExec(query, system); err != nil {
|
||||
return errors.Wrapf(err, "failed to update system property with name=%s", system.Name)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlSystemStore) Get() (model.StringMap, error) {
|
||||
systems := []model.System{}
|
||||
props := make(model.StringMap)
|
||||
|
||||
if err := s.GetReplicaX().Select(&systems, "SELECT * FROM Systems"); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get System list")
|
||||
}
|
||||
|
||||
for _, prop := range systems {
|
||||
props[prop.Name] = prop.Value
|
||||
}
|
||||
|
||||
return props, nil
|
||||
}
|
||||
|
||||
func (s SqlSystemStore) GetByName(name string) (*model.System, error) {
|
||||
var system model.System
|
||||
if err := s.GetMasterX().Get(&system, "SELECT * FROM Systems WHERE Name = ?", name); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("System", fmt.Sprintf("name=%s", system.Name))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get system property with name=%s", system.Name)
|
||||
}
|
||||
|
||||
return &system, nil
|
||||
}
|
||||
|
||||
func (s SqlSystemStore) PermanentDeleteByName(name string) (*model.System, error) {
|
||||
var system model.System
|
||||
if _, err := s.GetMasterX().Exec("DELETE FROM Systems WHERE Name = ?", name); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to permanent delete system property with name=%s", system.Name)
|
||||
}
|
||||
|
||||
return &system, nil
|
||||
}
|
||||
|
||||
// InsertIfExists inserts a given system value if it does not already exist. If a value
|
||||
// already exists, it returns the old one, else returns the new one.
|
||||
func (s SqlSystemStore) InsertIfExists(system *model.System) (_ *model.System, err error) {
|
||||
tx, err := s.GetMasterX().BeginXWithIsolation(&sql.TxOptions{
|
||||
Isolation: sql.LevelSerializable,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(tx, &err)
|
||||
|
||||
var origSystem model.System
|
||||
if err := tx.Get(&origSystem, `SELECT * FROM Systems
|
||||
WHERE Name = ?`, system.Name); err != nil && err != sql.ErrNoRows {
|
||||
return nil, errors.Wrapf(err, "failed to get system property with name=%s", system.Name)
|
||||
}
|
||||
|
||||
if origSystem.Value != "" {
|
||||
// Already a value exists, return that.
|
||||
return &origSystem, nil
|
||||
}
|
||||
|
||||
// Key does not exist, need to insert.
|
||||
if _, err := tx.NamedExec("INSERT INTO Systems (Name, Value) VALUES (:Name, :Value)", system); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to save system property with name=%s", system.Name)
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
return system, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/system_store_test.go
Обычный файл
14
server/channels/store/sqlstore/system_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestSystemStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestSystemStore)
|
||||
}
|
||||
1692
server/channels/store/sqlstore/team_store.go
Обычный файл
1692
server/channels/store/sqlstore/team_store.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
521
server/channels/store/sqlstore/team_store_test.go
Обычный файл
521
server/channels/store/sqlstore/team_store_test.go
Обычный файл
@@ -0,0 +1,521 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestTeamStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestTeamStore)
|
||||
}
|
||||
|
||||
func TestTeamStoreInternalDataTypes(t *testing.T) {
|
||||
t.Run("NewTeamMemberFromModel", func(t *testing.T) { testNewTeamMemberFromModel(t) })
|
||||
t.Run("TeamMemberWithSchemeRolesToModel", func(t *testing.T) { testTeamMemberWithSchemeRolesToModel(t) })
|
||||
}
|
||||
|
||||
func testNewTeamMemberFromModel(t *testing.T) {
|
||||
m := model.TeamMember{
|
||||
TeamId: model.NewId(),
|
||||
UserId: model.NewId(),
|
||||
Roles: "team_user team_admin custom_role",
|
||||
DeleteAt: 12345,
|
||||
SchemeGuest: false,
|
||||
SchemeUser: true,
|
||||
SchemeAdmin: true,
|
||||
ExplicitRoles: "custom_role",
|
||||
}
|
||||
|
||||
db := NewTeamMemberFromModel(&m)
|
||||
|
||||
assert.Equal(t, m.TeamId, db.TeamId)
|
||||
assert.Equal(t, m.UserId, db.UserId)
|
||||
assert.Equal(t, m.DeleteAt, db.DeleteAt)
|
||||
assert.Equal(t, true, db.SchemeGuest.Valid)
|
||||
assert.Equal(t, true, db.SchemeUser.Valid)
|
||||
assert.Equal(t, true, db.SchemeAdmin.Valid)
|
||||
assert.Equal(t, m.SchemeGuest, db.SchemeGuest.Bool)
|
||||
assert.Equal(t, m.SchemeUser, db.SchemeUser.Bool)
|
||||
assert.Equal(t, m.SchemeAdmin, db.SchemeAdmin.Bool)
|
||||
assert.Equal(t, m.ExplicitRoles, db.Roles)
|
||||
}
|
||||
|
||||
func testTeamMemberWithSchemeRolesToModel(t *testing.T) {
|
||||
// Test all the non-role-related properties here.
|
||||
t.Run("BasicProperties", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
TeamId: model.NewId(),
|
||||
UserId: model.NewId(),
|
||||
Roles: "custom_role",
|
||||
DeleteAt: 12345,
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: true},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, db.TeamId, m.TeamId)
|
||||
assert.Equal(t, db.UserId, m.UserId)
|
||||
assert.Equal(t, "custom_role team_user team_admin", m.Roles)
|
||||
assert.Equal(t, db.DeleteAt, m.DeleteAt)
|
||||
assert.Equal(t, db.SchemeGuest.Bool, m.SchemeGuest)
|
||||
assert.Equal(t, db.SchemeUser.Bool, m.SchemeUser)
|
||||
assert.Equal(t, db.SchemeAdmin.Bool, m.SchemeAdmin)
|
||||
assert.Equal(t, db.Roles, m.ExplicitRoles)
|
||||
})
|
||||
|
||||
// Example data *before* the Phase 2 migration has taken place.
|
||||
t.Run("Unmigrated_NoScheme_User", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "team_user",
|
||||
SchemeGuest: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: false, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "team_user", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Unmigrated_NoScheme_Admin", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "team_user team_admin",
|
||||
SchemeGuest: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: false, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "team_user team_admin", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, true, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Unmigrated_NoScheme_CustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: false, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Unmigrated_NoScheme_UserAndCustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "team_user custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: false, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role team_user", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Unmigrated_NoScheme_AdminAndCustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "team_user team_admin custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: false, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role team_user team_admin", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, true, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Unmigrated_NoScheme_NoRoles", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "",
|
||||
SchemeGuest: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: false, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: false, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
// Example data *after* the Phase 2 migration has taken place.
|
||||
t.Run("Migrated_NoScheme_User", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "team_user", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_NoScheme_Admin", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: true},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "team_user team_admin", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, true, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_NoScheme_Guest", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "team_guest", m.Roles)
|
||||
assert.Equal(t, true, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_NoScheme_CustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_NoScheme_UserAndCustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role team_user", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_NoScheme_AdminAndCustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: true},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role team_user team_admin", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, true, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_NoScheme_GuestAndCustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role team_guest", m.Roles)
|
||||
assert.Equal(t, true, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_NoScheme_NoRoles", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: false},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: false},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
// Example data with a team scheme.
|
||||
t.Run("Migrated_TeamScheme_User", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: true, String: "tscheme_guest"},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: true, String: "tscheme_user"},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: true, String: "tscheme_admin"},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "tscheme_user", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_TeamScheme_Admin", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: true},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: true, String: "tscheme_guest"},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: true, String: "tscheme_user"},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: true, String: "tscheme_admin"},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "tscheme_user tscheme_admin", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, true, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_TeamScheme_Guest", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: true, String: "tscheme_guest"},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: true, String: "tscheme_user"},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: true, String: "tscheme_admin"},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "tscheme_guest", m.Roles)
|
||||
assert.Equal(t, true, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_TeamScheme_CustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: true, String: "tscheme_guest"},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: true, String: "tscheme_user"},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: true, String: "tscheme_admin"},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_TeamScheme_UserAndCustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: true, String: "tscheme_guest"},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: true, String: "tscheme_user"},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: true, String: "tscheme_admin"},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role tscheme_user", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_TeamScheme_AdminAndCustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: true},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: true, String: "tscheme_guest"},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: true, String: "tscheme_user"},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: true, String: "tscheme_admin"},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role tscheme_user tscheme_admin", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, true, m.SchemeUser)
|
||||
assert.Equal(t, true, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_TeamScheme_GuestAndCustomRole", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "custom_role",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: true},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: true, String: "tscheme_guest"},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: true, String: "tscheme_user"},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: true, String: "tscheme_admin"},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "custom_role tscheme_guest", m.Roles)
|
||||
assert.Equal(t, true, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "custom_role", m.ExplicitRoles)
|
||||
})
|
||||
|
||||
t.Run("Migrated_TeamScheme_NoRoles", func(t *testing.T) {
|
||||
db := teamMemberWithSchemeRoles{
|
||||
Roles: "",
|
||||
SchemeGuest: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeUser: sql.NullBool{Valid: true, Bool: false},
|
||||
SchemeAdmin: sql.NullBool{Valid: true, Bool: false},
|
||||
TeamSchemeDefaultGuestRole: sql.NullString{Valid: true, String: "tscheme_guest"},
|
||||
TeamSchemeDefaultUserRole: sql.NullString{Valid: true, String: "tscheme_user"},
|
||||
TeamSchemeDefaultAdminRole: sql.NullString{Valid: true, String: "tscheme_admin"},
|
||||
}
|
||||
|
||||
m := db.ToModel()
|
||||
|
||||
assert.Equal(t, "", m.Roles)
|
||||
assert.Equal(t, false, m.SchemeGuest)
|
||||
assert.Equal(t, false, m.SchemeUser)
|
||||
assert.Equal(t, false, m.SchemeAdmin)
|
||||
assert.Equal(t, "", m.ExplicitRoles)
|
||||
})
|
||||
}
|
||||
92
server/channels/store/sqlstore/terms_of_service_store.go
Обычный файл
92
server/channels/store/sqlstore/terms_of_service_store.go
Обычный файл
@@ -0,0 +1,92 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/einterfaces"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlTermsOfServiceStore struct {
|
||||
*SqlStore
|
||||
metrics einterfaces.MetricsInterface
|
||||
}
|
||||
|
||||
func newSqlTermsOfServiceStore(sqlStore *SqlStore, metrics einterfaces.MetricsInterface) store.TermsOfServiceStore {
|
||||
return SqlTermsOfServiceStore{sqlStore, metrics}
|
||||
}
|
||||
|
||||
func (s SqlTermsOfServiceStore) Save(termsOfService *model.TermsOfService) (*model.TermsOfService, error) {
|
||||
if termsOfService.Id != "" {
|
||||
return nil, store.NewErrInvalidInput("TermsOfService", "Id", termsOfService.Id)
|
||||
}
|
||||
|
||||
termsOfService.PreSave()
|
||||
|
||||
if err := termsOfService.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
query := `INSERT INTO TermsOfService
|
||||
(Id, CreateAt, UserId, Text)
|
||||
VALUES
|
||||
(:Id, :CreateAt, :UserId, :Text)
|
||||
`
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(query, termsOfService); err != nil {
|
||||
return nil, errors.Wrapf(err, "could not save a new TermsOfService")
|
||||
}
|
||||
|
||||
return termsOfService, nil
|
||||
}
|
||||
|
||||
func (s SqlTermsOfServiceStore) GetLatest(allowFromCache bool) (*model.TermsOfService, error) {
|
||||
var termsOfService model.TermsOfService
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("TermsOfService").
|
||||
OrderBy("CreateAt DESC").
|
||||
Limit(uint64(1))
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not build sql query to get latest TOS")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Get(&termsOfService, queryString, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("TermsOfService", "CreateAt=latest")
|
||||
}
|
||||
return nil, errors.Wrap(err, "could not find latest TermsOfService")
|
||||
}
|
||||
|
||||
return &termsOfService, nil
|
||||
}
|
||||
|
||||
func (s SqlTermsOfServiceStore) Get(id string, allowFromCache bool) (*model.TermsOfService, error) {
|
||||
var termsOfService model.TermsOfService
|
||||
queryString, _, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("TermsOfService").
|
||||
Where("id = ?").
|
||||
ToSql()
|
||||
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "terms_of_service_to_sql")
|
||||
}
|
||||
|
||||
err = s.GetReplicaX().Get(&termsOfService, queryString, id)
|
||||
if err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("TermsOfService", "id")
|
||||
}
|
||||
return nil, errors.Wrapf(err, "could not find TermsOfService with id=%s", id)
|
||||
}
|
||||
return &termsOfService, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/terms_of_service_store_test.go
Обычный файл
14
server/channels/store/sqlstore/terms_of_service_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestTermsOfServiceStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestTermsOfServiceStore)
|
||||
}
|
||||
1193
server/channels/store/sqlstore/thread_store.go
Обычный файл
1193
server/channels/store/sqlstore/thread_store.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
14
server/channels/store/sqlstore/thread_store_test.go
Обычный файл
14
server/channels/store/sqlstore/thread_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestThreadStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestThreadStore)
|
||||
}
|
||||
93
server/channels/store/sqlstore/tokens_store.go
Обычный файл
93
server/channels/store/sqlstore/tokens_store.go
Обычный файл
@@ -0,0 +1,93 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
)
|
||||
|
||||
type SqlTokenStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlTokenStore(sqlStore *SqlStore) store.TokenStore {
|
||||
return &SqlTokenStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlTokenStore) Save(token *model.Token) error {
|
||||
if err := token.IsValid(); err != nil {
|
||||
return err
|
||||
}
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Insert("Tokens").
|
||||
Columns("Token", "CreateAt", "Type", "Extra").
|
||||
Values(token.Token, token.CreateAt, token.Type, token.Extra).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "token_tosql")
|
||||
}
|
||||
if _, err := s.GetMasterX().Exec(query, args...); err != nil {
|
||||
return errors.Wrap(err, "failed to save Token")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlTokenStore) Delete(token string) error {
|
||||
if _, err := s.GetMasterX().Exec("DELETE FROM Tokens WHERE Token = ?", token); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Token with value %s", token)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlTokenStore) GetByToken(tokenString string) (*model.Token, error) {
|
||||
var token model.Token
|
||||
|
||||
if err := s.GetReplicaX().Get(&token, "SELECT * FROM Tokens WHERE Token = ?", tokenString); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("Token", fmt.Sprintf("Token=%s", tokenString))
|
||||
}
|
||||
|
||||
return nil, errors.Wrapf(err, "failed to get Token with value %s", tokenString)
|
||||
}
|
||||
|
||||
return &token, nil
|
||||
}
|
||||
|
||||
func (s SqlTokenStore) Cleanup(expiryTime int64) {
|
||||
if _, err := s.GetMasterX().Exec("DELETE FROM Tokens WHERE CreateAt < ?", expiryTime); err != nil {
|
||||
mlog.Error("Unable to cleanup token store.")
|
||||
}
|
||||
}
|
||||
|
||||
func (s SqlTokenStore) GetAllTokensByType(tokenType string) ([]*model.Token, error) {
|
||||
tokens := []*model.Token{}
|
||||
query, args, err := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("Tokens").
|
||||
Where(sq.Eq{"Type": tokenType}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "could not build sql query to get all tokens by type")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&tokens, query, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to get all tokens of Type=%s", tokenType)
|
||||
}
|
||||
return tokens, nil
|
||||
}
|
||||
|
||||
func (s SqlTokenStore) RemoveAllTokensByType(tokenType string) error {
|
||||
if _, err := s.GetMasterX().Exec("DELETE FROM Tokens WHERE Type = ?", tokenType); err != nil {
|
||||
return errors.Wrapf(err, "failed to remove all Tokens with Type=%s", tokenType)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/tokens_store_test.go
Обычный файл
14
server/channels/store/sqlstore/tokens_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestTokensStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestTokensStore)
|
||||
}
|
||||
81
server/channels/store/sqlstore/true_up_review_store.go
Обычный файл
81
server/channels/store/sqlstore/true_up_review_store.go
Обычный файл
@@ -0,0 +1,81 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"strconv"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
// SqlLicenseStore encapsulates the database writes and reads for
|
||||
// model.LicenseRecord objects.
|
||||
type SqlTrueUpReviewStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlTrueUpReviewStore(sqlStore *SqlStore) store.TrueUpReviewStore {
|
||||
return &SqlTrueUpReviewStore{sqlStore}
|
||||
}
|
||||
|
||||
func trueUpReviewStatusColumns() []string {
|
||||
return []string{
|
||||
"DueDate",
|
||||
"Completed",
|
||||
}
|
||||
}
|
||||
|
||||
func (s *SqlTrueUpReviewStore) GetTrueUpReviewStatus(dueDate int64) (*model.TrueUpReviewStatus, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("TrueUpReviewHistory").
|
||||
Where(sq.Eq{"DueDate": dueDate})
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "get_trueUpReviewStatusRecord_tosql")
|
||||
}
|
||||
var trueUpReviewStatus model.TrueUpReviewStatus
|
||||
if err := s.GetReplicaX().Get(&trueUpReviewStatus, queryString, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("TrueUpReviewStatus", strconv.FormatInt(dueDate, 10))
|
||||
}
|
||||
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &trueUpReviewStatus, nil
|
||||
}
|
||||
|
||||
func (s *SqlTrueUpReviewStore) CreateTrueUpReviewStatusRecord(reviewStatus *model.TrueUpReviewStatus) (*model.TrueUpReviewStatus, error) {
|
||||
builder := s.getQueryBuilder().Insert("TrueUpReviewHistory").Columns(trueUpReviewStatusColumns()...).Values(reviewStatus.ToSlice()...)
|
||||
query, args, err := builder.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "create_trueUpReviewStatusRecord_tosql")
|
||||
}
|
||||
|
||||
if _, err = s.GetMasterX().Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "fail to create true up review status record")
|
||||
}
|
||||
|
||||
return reviewStatus, nil
|
||||
}
|
||||
|
||||
func (s *SqlTrueUpReviewStore) Update(reviewStatus *model.TrueUpReviewStatus) (*model.TrueUpReviewStatus, error) {
|
||||
query := s.getQueryBuilder().
|
||||
Update("TrueUpReviewHistory").
|
||||
Set("Completed", reviewStatus.Completed).
|
||||
Where(sq.Eq{"DueDate": reviewStatus.DueDate})
|
||||
|
||||
if _, err := s.GetMasterX().ExecBuilder(query); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update true up review status with DueDate=%d", reviewStatus.DueDate)
|
||||
}
|
||||
|
||||
return reviewStatus, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/true_up_review_store_test.go
Обычный файл
14
server/channels/store/sqlstore/true_up_review_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestTrueUpReviewStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestTrueUpReviewStatusStore)
|
||||
}
|
||||
139
server/channels/store/sqlstore/upload_session_store.go
Обычный файл
139
server/channels/store/sqlstore/upload_session_store.go
Обычный файл
@@ -0,0 +1,139 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlUploadSessionStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlUploadSessionStore(sqlStore *SqlStore) store.UploadSessionStore {
|
||||
return &SqlUploadSessionStore{
|
||||
SqlStore: sqlStore,
|
||||
}
|
||||
}
|
||||
|
||||
func (us SqlUploadSessionStore) Save(session *model.UploadSession) (*model.UploadSession, error) {
|
||||
if session == nil {
|
||||
return nil, errors.New("SqlUploadSessionStore.Save: session should not be nil")
|
||||
}
|
||||
session.PreSave()
|
||||
if err := session.IsValid(); err != nil {
|
||||
return nil, errors.Wrap(err, "SqlUploadSessionStore.Save: validation failed")
|
||||
}
|
||||
query, args, err := us.getQueryBuilder().
|
||||
Insert("UploadSessions").
|
||||
Columns("Id", "Type", "CreateAt", "UserId", "ChannelId", "Filename", "Path", "FileSize", "FileOffset", "RemoteId", "ReqFileId").
|
||||
Values(session.Id, session.Type, session.CreateAt, session.UserId, session.ChannelId, session.Filename, session.Path, session.FileSize, session.FileOffset, session.RemoteId, session.ReqFileId).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "SqlUploadSessionStore.Save: failed to build query")
|
||||
}
|
||||
if _, err := us.GetMasterX().Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "SqlUploadSessionStore.Save: failed to insert")
|
||||
}
|
||||
return session, nil
|
||||
}
|
||||
|
||||
func (us SqlUploadSessionStore) Update(session *model.UploadSession) error {
|
||||
if session == nil {
|
||||
return errors.New("SqlUploadSessionStore.Update: session should not be nil")
|
||||
}
|
||||
if err := session.IsValid(); err != nil {
|
||||
return errors.Wrap(err, "SqlUploadSessionStore.Update: validation failed")
|
||||
}
|
||||
query, args, err := us.getQueryBuilder().
|
||||
Update("UploadSessions").
|
||||
Set("Type", session.Type).
|
||||
Set("CreateAt", session.CreateAt).
|
||||
Set("UserId", session.UserId).
|
||||
Set("ChannelId", session.ChannelId).
|
||||
Set("Filename", session.Filename).
|
||||
Set("Path", session.Path).
|
||||
Set("FileSize", session.FileSize).
|
||||
Set("FileOffset", session.FileOffset).
|
||||
Set("RemoteId", session.RemoteId).
|
||||
Set("ReqFileId", session.ReqFileId).
|
||||
Where(sq.Eq{"Id": session.Id}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "SqlUploadSessionStore.Update: failed to build query")
|
||||
}
|
||||
if _, err := us.GetMasterX().Exec(query, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return store.NewErrNotFound("UploadSession", session.Id)
|
||||
}
|
||||
return errors.Wrapf(err, "SqlUploadSessionStore.Update: failed to update session with id=%s", session.Id)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (us SqlUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) {
|
||||
if !model.IsValidId(id) {
|
||||
return nil, errors.New("SqlUploadSessionStore.Get: id is not valid")
|
||||
}
|
||||
query, args, err := us.getQueryBuilder().
|
||||
Select("*").
|
||||
From("UploadSessions").
|
||||
Where(sq.Eq{"Id": id}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "SqlUploadSessionStore.Get: failed to build query")
|
||||
}
|
||||
var session model.UploadSession
|
||||
if err := us.DBXFromContext(ctx).Get(&session, query, args...); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("UploadSession", id)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "SqlUploadSessionStore.Get: failed to select session with id=%s", id)
|
||||
}
|
||||
return &session, nil
|
||||
}
|
||||
|
||||
func (us SqlUploadSessionStore) GetForUser(userId string) ([]*model.UploadSession, error) {
|
||||
query, args, err := us.getQueryBuilder().
|
||||
Select("*").
|
||||
From("UploadSessions").
|
||||
Where(sq.Eq{"UserId": userId}).
|
||||
OrderBy("CreateAt ASC").
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "SqlUploadSessionStore.GetForUser: failed to build query")
|
||||
}
|
||||
sessions := []*model.UploadSession{}
|
||||
if err := us.GetReplicaX().Select(&sessions, query, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "SqlUploadSessionStore.GetForUser: failed to select")
|
||||
}
|
||||
return sessions, nil
|
||||
}
|
||||
|
||||
func (us SqlUploadSessionStore) Delete(id string) error {
|
||||
if !model.IsValidId(id) {
|
||||
return errors.New("SqlUploadSessionStore.Delete: id is not valid")
|
||||
}
|
||||
|
||||
query, args, err := us.getQueryBuilder().
|
||||
Delete("UploadSessions").
|
||||
Where(sq.Eq{"Id": id}).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "SqlUploadSessionStore.Delete: failed to build query")
|
||||
}
|
||||
|
||||
if _, err := us.GetMasterX().Exec(query, args...); err != nil {
|
||||
return errors.Wrap(err, "SqlUploadSessionStore.Delete: failed to delete")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/upload_session_store_test.go
Обычный файл
14
server/channels/store/sqlstore/upload_session_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestUploadSessionStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestUploadSessionStore)
|
||||
}
|
||||
238
server/channels/store/sqlstore/user_access_token_store.go
Обычный файл
238
server/channels/store/sqlstore/user_access_token_store.go
Обычный файл
@@ -0,0 +1,238 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlUserAccessTokenStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlUserAccessTokenStore(sqlStore *SqlStore) store.UserAccessTokenStore {
|
||||
return &SqlUserAccessTokenStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) Save(token *model.UserAccessToken) (*model.UserAccessToken, error) {
|
||||
token.PreSave()
|
||||
|
||||
if err := token.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query, args, err := s.getQueryBuilder().Insert("UserAccessTokens").
|
||||
Columns("Id", "Token", "UserId", "Description", "IsActive").
|
||||
Values(token.Id, token.Token, token.UserId, token.Description, token.IsActive).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "UserAccessToken_tosql")
|
||||
}
|
||||
if _, err := s.GetMasterX().Exec(query, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to save UserAccessToken")
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) Delete(tokenId string) (err error) {
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
|
||||
if err := s.deleteSessionsAndTokensById(transaction, tokenId); err == nil {
|
||||
if err := transaction.Commit(); err != nil {
|
||||
// don't need to rollback here since the transaction is already closed
|
||||
return errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) deleteSessionsAndTokensById(transaction *sqlxTxWrapper, tokenId string) error {
|
||||
|
||||
query := ""
|
||||
if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = "DELETE FROM Sessions s USING UserAccessTokens o WHERE o.Token = s.Token AND o.Id = ?"
|
||||
} else if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = "DELETE s.* FROM Sessions s INNER JOIN UserAccessTokens o ON o.Token = s.Token WHERE o.Id = ?"
|
||||
}
|
||||
|
||||
if _, err := transaction.Exec(query, tokenId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Sessions with UserAccessToken id=%s", tokenId)
|
||||
}
|
||||
|
||||
return s.deleteTokensById(transaction, tokenId)
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) deleteTokensById(transaction *sqlxTxWrapper, tokenId string) error {
|
||||
|
||||
if _, err := transaction.Exec("DELETE FROM UserAccessTokens WHERE Id = ?", tokenId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete UserAccessToken id=%s", tokenId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) DeleteAllForUser(userId string) (err error) {
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
if err := s.deleteSessionsandTokensByUser(transaction, userId); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := transaction.Commit(); err != nil {
|
||||
// don't need to rollback here since the transaction is already closed
|
||||
return errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) deleteSessionsandTokensByUser(transaction *sqlxTxWrapper, userId string) error {
|
||||
query := ""
|
||||
if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = "DELETE FROM Sessions s USING UserAccessTokens o WHERE o.Token = s.Token AND o.UserId = ?"
|
||||
} else if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = "DELETE s.* FROM Sessions s INNER JOIN UserAccessTokens o ON o.Token = s.Token WHERE o.UserId = ?"
|
||||
}
|
||||
|
||||
if _, err := transaction.Exec(query, userId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Sessions with UserAccessToken userId=%s", userId)
|
||||
}
|
||||
|
||||
return s.deleteTokensByUser(transaction, userId)
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) deleteTokensByUser(transaction *sqlxTxWrapper, userId string) error {
|
||||
if _, err := transaction.Exec("DELETE FROM UserAccessTokens WHERE UserId = ?", userId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete UserAccessToken userId=%s", userId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) Get(tokenId string) (*model.UserAccessToken, error) {
|
||||
var token model.UserAccessToken
|
||||
|
||||
if err := s.GetReplicaX().Get(&token, "SELECT * FROM UserAccessTokens WHERE Id = ?", tokenId); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("UserAccessToken", tokenId)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get UserAccessToken with id=%s", tokenId)
|
||||
}
|
||||
|
||||
return &token, nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) GetAll(offset, limit int) ([]*model.UserAccessToken, error) {
|
||||
tokens := []*model.UserAccessToken{}
|
||||
|
||||
if err := s.GetReplicaX().Select(&tokens, "SELECT * FROM UserAccessTokens LIMIT ? OFFSET ?", limit, offset); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find UserAccessTokens")
|
||||
}
|
||||
|
||||
return tokens, nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) GetByToken(tokenString string) (*model.UserAccessToken, error) {
|
||||
var token model.UserAccessToken
|
||||
|
||||
if err := s.GetReplicaX().Get(&token, "SELECT * FROM UserAccessTokens WHERE Token = ?", tokenString); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("UserAccessToken", fmt.Sprintf("token=%s", tokenString))
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get UserAccessToken with token=%s", tokenString)
|
||||
}
|
||||
|
||||
return &token, nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) GetByUser(userId string, offset, limit int) ([]*model.UserAccessToken, error) {
|
||||
tokens := []*model.UserAccessToken{}
|
||||
|
||||
if err := s.GetReplicaX().Select(&tokens, "SELECT * FROM UserAccessTokens WHERE UserId = ? LIMIT ? OFFSET ?", userId, limit, offset); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find UserAccessTokens with userId=%s", userId)
|
||||
}
|
||||
|
||||
return tokens, nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) Search(term string) ([]*model.UserAccessToken, error) {
|
||||
term = sanitizeSearchTerm(term, "\\")
|
||||
tokens := []*model.UserAccessToken{}
|
||||
params := []any{term, term, term}
|
||||
query := `
|
||||
SELECT
|
||||
uat.*
|
||||
FROM UserAccessTokens uat
|
||||
INNER JOIN Users u
|
||||
ON uat.UserId = u.Id
|
||||
WHERE uat.Id LIKE ? OR uat.UserId LIKE ? OR u.Username LIKE ?`
|
||||
|
||||
if err := s.GetReplicaX().Select(&tokens, query, params...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find UserAccessTokens by term with value '%s'", term)
|
||||
}
|
||||
|
||||
return tokens, nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) UpdateTokenEnable(tokenId string) error {
|
||||
if _, err := s.GetMasterX().Exec("UPDATE UserAccessTokens SET IsActive = TRUE WHERE Id = ?", tokenId); err != nil {
|
||||
return errors.Wrapf(err, "failed to update UserAccessTokens with id=%s", tokenId)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) UpdateTokenDisable(tokenId string) (err error) {
|
||||
transaction, err := s.GetMasterX().Beginx()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "begin_transaction")
|
||||
}
|
||||
defer finalizeTransactionX(transaction, &err)
|
||||
|
||||
if err := s.deleteSessionsAndDisableToken(transaction, tokenId); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := transaction.Commit(); err != nil {
|
||||
// don't need to rollback here since the transaction is already closed
|
||||
return errors.Wrap(err, "commit_transaction")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) deleteSessionsAndDisableToken(transaction *sqlxTxWrapper, tokenId string) error {
|
||||
query := ""
|
||||
if s.DriverName() == model.DatabaseDriverPostgres {
|
||||
query = "DELETE FROM Sessions s USING UserAccessTokens o WHERE o.Token = s.Token AND o.Id = ?"
|
||||
} else if s.DriverName() == model.DatabaseDriverMysql {
|
||||
query = "DELETE s.* FROM Sessions s INNER JOIN UserAccessTokens o ON o.Token = s.Token WHERE o.Id = ?"
|
||||
}
|
||||
|
||||
if _, err := transaction.Exec(query, tokenId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete Sessions with UserAccessToken id=%s", tokenId)
|
||||
}
|
||||
|
||||
return s.updateTokenDisable(transaction, tokenId)
|
||||
}
|
||||
|
||||
func (s SqlUserAccessTokenStore) updateTokenDisable(transaction *sqlxTxWrapper, tokenId string) error {
|
||||
if _, err := transaction.Exec("UPDATE UserAccessTokens SET IsActive = FALSE WHERE Id = ?", tokenId); err != nil {
|
||||
return errors.Wrapf(err, "failed to update UserAccessToken with id=%s", tokenId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/user_access_token_store_test.go
Обычный файл
14
server/channels/store/sqlstore/user_access_token_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestUserAccessTokenStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestUserAccessTokenStore)
|
||||
}
|
||||
2202
server/channels/store/sqlstore/user_store.go
Обычный файл
2202
server/channels/store/sqlstore/user_store.go
Обычный файл
Разница между файлами не показана из-за своего большого размера
Загрузить разницу
19
server/channels/store/sqlstore/user_store_test.go
Обычный файл
19
server/channels/store/sqlstore/user_store_test.go
Обычный файл
@@ -0,0 +1,19 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/searchtest"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestUserStore(t *testing.T) {
|
||||
StoreTestWithSqlStore(t, storetest.TestUserStore)
|
||||
}
|
||||
|
||||
func TestSearchUserStore(t *testing.T) {
|
||||
StoreTestWithSearchTestEngine(t, searchtest.TestSearchUserStore)
|
||||
}
|
||||
87
server/channels/store/sqlstore/user_terms_of_service.go
Обычный файл
87
server/channels/store/sqlstore/user_terms_of_service.go
Обычный файл
@@ -0,0 +1,87 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlUserTermsOfServiceStore struct {
|
||||
*SqlStore
|
||||
}
|
||||
|
||||
func newSqlUserTermsOfServiceStore(sqlStore *SqlStore) store.UserTermsOfServiceStore {
|
||||
return SqlUserTermsOfServiceStore{sqlStore}
|
||||
}
|
||||
|
||||
func (s SqlUserTermsOfServiceStore) GetByUser(userId string) (*model.UserTermsOfService, error) {
|
||||
var userTermsOfService model.UserTermsOfService
|
||||
query := `
|
||||
SELECT *
|
||||
FROM UserTermsOfService
|
||||
WHERE UserId = ?
|
||||
`
|
||||
if err := s.GetReplicaX().Get(&userTermsOfService, query, userId); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("UserTermsOfService", "userId="+userId)
|
||||
}
|
||||
|
||||
return nil, errors.Wrapf(err, "failed to get UserTermsOfService with userId=%s", userId)
|
||||
}
|
||||
|
||||
return &userTermsOfService, nil
|
||||
}
|
||||
|
||||
func (s SqlUserTermsOfServiceStore) Save(userTermsOfService *model.UserTermsOfService) (*model.UserTermsOfService, error) {
|
||||
userTermsOfService.PreSave()
|
||||
if err := userTermsOfService.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
query := `
|
||||
UPDATE UserTermsOfService
|
||||
SET UserId = :UserId, TermsOfServiceId = :TermsOfServiceId, CreateAt = :CreateAt
|
||||
WHERE UserId = :UserId
|
||||
`
|
||||
result, err := s.GetMasterX().NamedExec(query, userTermsOfService)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update UserTermsOfService with userId=%s and termsOfServiceId=%s", userTermsOfService.UserId, userTermsOfService.TermsOfServiceId)
|
||||
}
|
||||
|
||||
updatedRows, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to retrieve the number of affected rows for the update of UserTermsOfService")
|
||||
}
|
||||
if updatedRows == 0 {
|
||||
query := `
|
||||
INSERT INTO UserTermsOfService
|
||||
(UserId, TermsOfServiceId, CreateAt)
|
||||
VALUES
|
||||
(:UserId, :TermsOfServiceId, :CreateAt)
|
||||
`
|
||||
if _, err := s.GetMasterX().NamedExec(query, userTermsOfService); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to save UserTermsOfService with userId=%s and termsOfServiceId=%s", userTermsOfService.UserId, userTermsOfService.TermsOfServiceId)
|
||||
}
|
||||
}
|
||||
|
||||
return userTermsOfService, nil
|
||||
}
|
||||
|
||||
func (s SqlUserTermsOfServiceStore) Delete(userId, termsOfServiceId string) error {
|
||||
query := `
|
||||
DELETE
|
||||
FROM UserTermsOfService
|
||||
WHERE UserId = ? AND TermsOfServiceId = ?
|
||||
`
|
||||
if _, err := s.GetMasterX().Exec(query, userId, termsOfServiceId); err != nil {
|
||||
return errors.Wrapf(err, "failed to delete UserTermsOfService with userId=%s and termsOfServiceId=%s", userId, termsOfServiceId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestUserTermsOfServiceStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestUserTermsOfServiceStore)
|
||||
}
|
||||
208
server/channels/store/sqlstore/utils.go
Обычный файл
208
server/channels/store/sqlstore/utils.go
Обычный файл
@@ -0,0 +1,208 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"io"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"unicode"
|
||||
|
||||
"github.com/wiggin77/merror"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/platform/shared/mlog"
|
||||
|
||||
"github.com/go-sql-driver/mysql"
|
||||
)
|
||||
|
||||
var escapeLikeSearchChar = []string{
|
||||
"%",
|
||||
"_",
|
||||
}
|
||||
|
||||
func sanitizeSearchTerm(term string, escapeChar string) string {
|
||||
term = strings.Replace(term, escapeChar, "", -1)
|
||||
|
||||
for _, c := range escapeLikeSearchChar {
|
||||
term = strings.Replace(term, c, escapeChar+c, -1)
|
||||
}
|
||||
|
||||
return term
|
||||
}
|
||||
|
||||
// Converts a list of strings into a list of query parameters and a named parameter map that can
|
||||
// be used as part of a SQL query.
|
||||
func MapStringsToQueryParams(list []string, paramPrefix string) (string, map[string]any) {
|
||||
var keys strings.Builder
|
||||
params := make(map[string]any, len(list))
|
||||
for i, entry := range list {
|
||||
if keys.Len() > 0 {
|
||||
keys.WriteString(",")
|
||||
}
|
||||
|
||||
key := paramPrefix + strconv.Itoa(i)
|
||||
keys.WriteString(":" + key)
|
||||
params[key] = entry
|
||||
}
|
||||
|
||||
return "(" + keys.String() + ")", params
|
||||
}
|
||||
|
||||
// finalizeTransactionX ensures a transaction is closed after use, rolling back if not already committed.
|
||||
func finalizeTransactionX(transaction *sqlxTxWrapper, perr *error) {
|
||||
// Rollback returns sql.ErrTxDone if the transaction was already closed.
|
||||
if err := transaction.Rollback(); err != nil && err != sql.ErrTxDone {
|
||||
*perr = merror.Append(*perr, err)
|
||||
}
|
||||
}
|
||||
|
||||
func deferClose(c io.Closer, perr *error) {
|
||||
err := c.Close()
|
||||
*perr = merror.Append(*perr, err)
|
||||
}
|
||||
|
||||
// removeNonAlphaNumericUnquotedTerms removes all unquoted words that only contain
|
||||
// non-alphanumeric chars from given line
|
||||
func removeNonAlphaNumericUnquotedTerms(line, separator string) string {
|
||||
words := strings.Split(line, separator)
|
||||
filteredResult := make([]string, 0, len(words))
|
||||
|
||||
for _, w := range words {
|
||||
if isQuotedWord(w) || containsAlphaNumericChar(w) {
|
||||
filteredResult = append(filteredResult, strings.TrimSpace(w))
|
||||
}
|
||||
}
|
||||
return strings.Join(filteredResult, separator)
|
||||
}
|
||||
|
||||
// containsAlphaNumericChar returns true in case any letter or digit is present, false otherwise
|
||||
func containsAlphaNumericChar(s string) bool {
|
||||
for _, r := range s {
|
||||
if unicode.IsLetter(r) || unicode.IsDigit(r) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// isQuotedWord return true if the input string is quoted, false otherwise. Ex :-
|
||||
//
|
||||
// "quoted string" - will return true
|
||||
// unquoted string - will return false
|
||||
func isQuotedWord(s string) bool {
|
||||
if len(s) < 2 {
|
||||
return false
|
||||
}
|
||||
|
||||
return s[0] == '"' && s[len(s)-1] == '"'
|
||||
}
|
||||
|
||||
// constructMySQLJSONArgs returns the arg list to pass to a query along with
|
||||
// the string of placeholders which is needed to be to the JSON_SET function.
|
||||
// Use this function in this way:
|
||||
// UPDATE Table
|
||||
// SET Col = JSON_SET(Col, `+argString+`)
|
||||
// WHERE Id=?`, args...)
|
||||
// after appending the Id param to the args slice.
|
||||
func constructMySQLJSONArgs(props map[string]string) ([]any, string) {
|
||||
if len(props) == 0 {
|
||||
return nil, ""
|
||||
}
|
||||
|
||||
// Unpack the keys and values to pass to MySQL.
|
||||
args := make([]any, 0, len(props))
|
||||
for k, v := range props {
|
||||
args = append(args, "$."+k, v)
|
||||
}
|
||||
|
||||
// We calculate the number of ? to set in the query string.
|
||||
argString := strings.Repeat("?, ", len(props)*2)
|
||||
// Strip off the trailing comma.
|
||||
argString = strings.TrimSuffix(argString, ", ")
|
||||
|
||||
return args, argString
|
||||
}
|
||||
|
||||
func makeStringArgs(params []string) []any {
|
||||
args := make([]any, len(params))
|
||||
for i, name := range params {
|
||||
args[i] = name
|
||||
}
|
||||
return args
|
||||
}
|
||||
|
||||
func constructArrayArgs(ids []string) (string, []any) {
|
||||
var placeholder strings.Builder
|
||||
values := make([]any, 0, len(ids))
|
||||
for _, entry := range ids {
|
||||
if placeholder.Len() > 0 {
|
||||
placeholder.WriteString(",")
|
||||
}
|
||||
|
||||
placeholder.WriteString("?")
|
||||
values = append(values, entry)
|
||||
}
|
||||
|
||||
return "(" + placeholder.String() + ")", values
|
||||
}
|
||||
|
||||
func wrapBinaryParamStringMap(ok bool, props model.StringMap) model.StringMap {
|
||||
if props == nil {
|
||||
props = make(model.StringMap)
|
||||
}
|
||||
props[model.BinaryParamKey] = strconv.FormatBool(ok)
|
||||
return props
|
||||
}
|
||||
|
||||
// morphWriter is a target to pass to the logger instance of morph.
|
||||
// For now, everything is just logged at a debug level. If we need to log
|
||||
// errors/warnings from the library also, that needs to be seen later.
|
||||
type morphWriter struct {
|
||||
}
|
||||
|
||||
func (l *morphWriter) Write(in []byte) (int, error) {
|
||||
mlog.Debug(string(in))
|
||||
return len(in), nil
|
||||
}
|
||||
|
||||
func DSNHasBinaryParam(dsn string) (bool, error) {
|
||||
url, err := url.Parse(dsn)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return url.Query().Get("binary_parameters") == "yes", nil
|
||||
}
|
||||
|
||||
// AppendBinaryFlag updates the byte slice to work using binary_parameters=yes.
|
||||
func AppendBinaryFlag(buf []byte) []byte {
|
||||
return append([]byte{0x01}, buf...)
|
||||
}
|
||||
|
||||
// AppendMultipleStatementsFlag attached dsn parameters to MySQL dsn in order to make migrations work.
|
||||
func AppendMultipleStatementsFlag(dataSource string) (string, error) {
|
||||
config, err := mysql.ParseDSN(dataSource)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
if config.Params == nil {
|
||||
config.Params = map[string]string{}
|
||||
}
|
||||
|
||||
config.Params["multiStatements"] = "true"
|
||||
return config.FormatDSN(), nil
|
||||
}
|
||||
|
||||
// ResetReadTimeout removes the timeout constraint from the MySQL dsn.
|
||||
func ResetReadTimeout(dataSource string) (string, error) {
|
||||
config, err := mysql.ParseDSN(dataSource)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
config.ReadTimeout = 0
|
||||
return config.FormatDSN(), nil
|
||||
}
|
||||
162
server/channels/store/sqlstore/utils_test.go
Обычный файл
162
server/channels/store/sqlstore/utils_test.go
Обычный файл
@@ -0,0 +1,162 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestMapStringsToQueryParams(t *testing.T) {
|
||||
t.Run("one item", func(t *testing.T) {
|
||||
input := []string{"apple"}
|
||||
|
||||
keys, params := MapStringsToQueryParams(input, "Fruit")
|
||||
|
||||
require.Len(t, params, 1, "returned incorrect params", params)
|
||||
require.Equal(t, "apple", params["Fruit0"], "returned incorrect params", params)
|
||||
require.Equal(t, "(:Fruit0)", keys, "returned incorrect query", keys)
|
||||
})
|
||||
|
||||
t.Run("multiple items", func(t *testing.T) {
|
||||
input := []string{"carrot", "tomato", "potato"}
|
||||
|
||||
keys, params := MapStringsToQueryParams(input, "Vegetable")
|
||||
|
||||
require.Len(t, params, 3, "returned incorrect params", params)
|
||||
require.Equal(t, "carrot", params["Vegetable0"], "returned incorrect params", params)
|
||||
require.Equal(t, "tomato", params["Vegetable1"], "returned incorrect params", params)
|
||||
require.Equal(t, "potato", params["Vegetable2"], "returned incorrect params", params)
|
||||
require.Equal(t, "(:Vegetable0,:Vegetable1,:Vegetable2)", keys, "returned incorrect query", keys)
|
||||
})
|
||||
}
|
||||
|
||||
var keys string
|
||||
var params map[string]any
|
||||
|
||||
func BenchmarkMapStringsToQueryParams(b *testing.B) {
|
||||
b.Run("one item", func(b *testing.B) {
|
||||
input := []string{"apple"}
|
||||
for i := 0; i < b.N; i++ {
|
||||
keys, params = MapStringsToQueryParams(input, "Fruit")
|
||||
}
|
||||
})
|
||||
b.Run("multiple items", func(b *testing.B) {
|
||||
input := []string{"carrot", "tomato", "potato"}
|
||||
for i := 0; i < b.N; i++ {
|
||||
keys, params = MapStringsToQueryParams(input, "Vegetable")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSanitizeSearchTerm(t *testing.T) {
|
||||
term := "test"
|
||||
result := sanitizeSearchTerm(term, "\\")
|
||||
require.Equal(t, result, term)
|
||||
|
||||
term = "%%%"
|
||||
expected := "\\%\\%\\%"
|
||||
result = sanitizeSearchTerm(term, "\\")
|
||||
require.Equal(t, result, expected)
|
||||
|
||||
term = "%\\%\\%"
|
||||
expected = "\\%\\%\\%"
|
||||
result = sanitizeSearchTerm(term, "\\")
|
||||
require.Equal(t, result, expected)
|
||||
|
||||
term = "%_test_%"
|
||||
expected = "\\%\\_test\\_\\%"
|
||||
result = sanitizeSearchTerm(term, "\\")
|
||||
require.Equal(t, result, expected)
|
||||
|
||||
term = "**test_%"
|
||||
expected = "test*_*%"
|
||||
result = sanitizeSearchTerm(term, "*")
|
||||
require.Equal(t, result, expected)
|
||||
}
|
||||
|
||||
func TestRemoveNonAlphaNumericUnquotedTerms(t *testing.T) {
|
||||
const (
|
||||
sep = " "
|
||||
chineseHello = "你好"
|
||||
japaneseHello = "こんにちは"
|
||||
)
|
||||
tests := []struct {
|
||||
term string
|
||||
want string
|
||||
name string
|
||||
}{
|
||||
{term: "", want: "", name: "empty"},
|
||||
{term: "h", want: "h", name: "singleChar"},
|
||||
{term: "hello", want: "hello", name: "multiChar"},
|
||||
{term: `hel*lo "**" **& hello`, want: `hel*lo "**" hello`, name: "quoted_unquoted_english"},
|
||||
{term: japaneseHello + chineseHello, want: japaneseHello + chineseHello, name: "japanese_chinese"},
|
||||
{term: japaneseHello + ` "*" ` + chineseHello, want: japaneseHello + ` "*" ` + chineseHello, name: `quoted_japanese_and_chinese`},
|
||||
{term: japaneseHello + ` "*" &&* ` + chineseHello, want: japaneseHello + ` "*" ` + chineseHello, name: "quoted_unquoted_japanese_and_chinese"},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
got := removeNonAlphaNumericUnquotedTerms(test.term, sep)
|
||||
require.Equal(t, test.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMySQLJSONArgs(t *testing.T) {
|
||||
tests := []struct {
|
||||
props map[string]string
|
||||
args []any
|
||||
argString string
|
||||
}{
|
||||
{
|
||||
props: map[string]string{
|
||||
"desktop": "linux",
|
||||
"mobile": "android",
|
||||
"notify": "always",
|
||||
},
|
||||
args: []any{"$.desktop", "linux", "$.mobile", "android", "$.notify", "always"},
|
||||
argString: "?, ?, ?, ?, ?, ?",
|
||||
},
|
||||
{
|
||||
props: map[string]string{},
|
||||
args: nil,
|
||||
argString: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
args, argString := constructMySQLJSONArgs(test.props)
|
||||
assert.ElementsMatch(t, test.args, args)
|
||||
assert.Equal(t, test.argString, argString)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAppendMultipleStatementsFlag(t *testing.T) {
|
||||
testCases := []struct {
|
||||
Scenario string
|
||||
DSN string
|
||||
ExpectedDSN string
|
||||
}{
|
||||
{
|
||||
"Should append multiStatements param to the DSN path with existing params",
|
||||
"user:rand?&ompasswith@character@unix(/var/run/mysqld/mysqld.sock)/mattermost?writeTimeout=30s",
|
||||
"user:rand?&ompasswith@character@unix(/var/run/mysqld/mysqld.sock)/mattermost?writeTimeout=30s&multiStatements=true",
|
||||
},
|
||||
{
|
||||
"Should append multiStatements param to the DSN path with no existing params",
|
||||
"user:rand?&ompasswith@character@unix(/var/run/mysqld/mysqld.sock)/mattermost",
|
||||
"user:rand?&ompasswith@character@unix(/var/run/mysqld/mysqld.sock)/mattermost?multiStatements=true",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.Scenario, func(t *testing.T) {
|
||||
res, err := AppendMultipleStatementsFlag(tc.DSN)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tc.ExpectedDSN, res)
|
||||
})
|
||||
}
|
||||
}
|
||||
402
server/channels/store/sqlstore/webhook_store.go
Обычный файл
402
server/channels/store/sqlstore/webhook_store.go
Обычный файл
@@ -0,0 +1,402 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
|
||||
sq "github.com/mattermost/squirrel"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/model"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/einterfaces"
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store"
|
||||
)
|
||||
|
||||
type SqlWebhookStore struct {
|
||||
*SqlStore
|
||||
metrics einterfaces.MetricsInterface
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) ClearCaches() {
|
||||
}
|
||||
|
||||
func newSqlWebhookStore(sqlStore *SqlStore, metrics einterfaces.MetricsInterface) store.WebhookStore {
|
||||
return &SqlWebhookStore{
|
||||
SqlStore: sqlStore,
|
||||
metrics: metrics,
|
||||
}
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) InvalidateWebhookCache(webhookId string) {
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) SaveIncoming(webhook *model.IncomingWebhook) (*model.IncomingWebhook, error) {
|
||||
if webhook.Id != "" {
|
||||
return nil, store.NewErrInvalidInput("IncomingWebhook", "id", webhook.Id)
|
||||
}
|
||||
|
||||
webhook.PreSave()
|
||||
if err := webhook.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(`INSERT INTO IncomingWebhooks
|
||||
(Id, CreateAt, UpdateAt, DeleteAt, UserId, ChannelId, TeamId, DisplayName, Description, Username, IconURL, ChannelLocked)
|
||||
VALUES
|
||||
(:Id, :CreateAt, :UpdateAt, :DeleteAt, :UserId, :ChannelId, :TeamId, :DisplayName, :Description, :Username, :IconURL, :ChannelLocked)`, webhook); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to save IncomingWebhook with id=%s", webhook.Id)
|
||||
}
|
||||
|
||||
return webhook, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) UpdateIncoming(hook *model.IncomingWebhook) (*model.IncomingWebhook, error) {
|
||||
hook.UpdateAt = model.GetMillis()
|
||||
|
||||
_, err := s.GetMasterX().NamedExec(`UPDATE IncomingWebhooks SET
|
||||
CreateAt=:CreateAt, UpdateAt=:UpdateAt, DeleteAt=:DeleteAt, ChannelId=:ChannelId, TeamId=:TeamId, DisplayName=:DisplayName,
|
||||
Description=:Description, Username=:Username, IconURL=:IconURL, ChannelLocked=:ChannelLocked
|
||||
WHERE Id=:Id`, hook)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update IncomingWebhook with id=%s", hook.Id)
|
||||
}
|
||||
|
||||
return hook, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetIncoming(id string, allowFromCache bool) (*model.IncomingWebhook, error) {
|
||||
var webhook model.IncomingWebhook
|
||||
if err := s.GetReplicaX().Get(&webhook, "SELECT * FROM IncomingWebhooks WHERE Id = ? AND DeleteAt = 0", id); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("IncomingWebhook", id)
|
||||
}
|
||||
return nil, errors.Wrapf(err, "failed to get IncomingWebhook with id=%s", id)
|
||||
}
|
||||
|
||||
return &webhook, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) DeleteIncoming(webhookId string, time int64) error {
|
||||
_, err := s.GetMasterX().Exec("UPDATE IncomingWebhooks SET DeleteAt = ?, UpdateAt = ? WHERE Id = ?", time, time, webhookId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to update IncomingWebhook with id=%s", webhookId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) PermanentDeleteIncomingByUser(userId string) error {
|
||||
_, err := s.GetMasterX().Exec("DELETE FROM IncomingWebhooks WHERE UserId = ?", userId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete IncomingWebhook with userId=%s", userId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) PermanentDeleteIncomingByChannel(channelId string) error {
|
||||
_, err := s.GetMasterX().Exec("DELETE FROM IncomingWebhooks WHERE ChannelId = ?", channelId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete IncomingWebhook with channelId=%s", channelId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetIncomingList(offset, limit int) ([]*model.IncomingWebhook, error) {
|
||||
return s.GetIncomingListByUser("", offset, limit)
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetIncomingListByUser(userId string, offset, limit int) ([]*model.IncomingWebhook, error) {
|
||||
webhooks := []*model.IncomingWebhook{}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("IncomingWebhooks").
|
||||
Where(sq.Eq{"DeleteAt": int(0)}).Limit(uint64(limit)).Offset(uint64(offset))
|
||||
|
||||
if userId != "" {
|
||||
query = query.Where(sq.Eq{"UserId": userId})
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "incoming_webhook_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&webhooks, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find IncomingWebhooks")
|
||||
}
|
||||
|
||||
return webhooks, nil
|
||||
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetIncomingByTeamByUser(teamId string, userId string, offset, limit int) ([]*model.IncomingWebhook, error) {
|
||||
webhooks := []*model.IncomingWebhook{}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("IncomingWebhooks").
|
||||
Where(sq.And{
|
||||
sq.Eq{"TeamId": teamId},
|
||||
sq.Eq{"DeleteAt": int(0)},
|
||||
}).Limit(uint64(limit)).Offset(uint64(offset))
|
||||
|
||||
if userId != "" {
|
||||
query = query.Where(sq.Eq{"UserId": userId})
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "incoming_webhook_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&webhooks, queryString, args...); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find IncomingWebhook with teamId=%s", teamId)
|
||||
}
|
||||
|
||||
return webhooks, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetIncomingByTeam(teamId string, offset, limit int) ([]*model.IncomingWebhook, error) {
|
||||
return s.GetIncomingByTeamByUser(teamId, "", offset, limit)
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetIncomingByChannel(channelId string) ([]*model.IncomingWebhook, error) {
|
||||
webhooks := []*model.IncomingWebhook{}
|
||||
|
||||
if err := s.GetReplicaX().Select(&webhooks, "SELECT * FROM IncomingWebhooks WHERE ChannelId = ? AND DeleteAt = 0", channelId); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find IncomingWebhooks with channelId=%s", channelId)
|
||||
}
|
||||
|
||||
return webhooks, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) SaveOutgoing(webhook *model.OutgoingWebhook) (*model.OutgoingWebhook, error) {
|
||||
if webhook.Id != "" {
|
||||
return nil, store.NewErrInvalidInput("OutgoingWebhook", "id", webhook.Id)
|
||||
}
|
||||
|
||||
webhook.PreSave()
|
||||
if err := webhook.IsValid(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, err := s.GetMasterX().NamedExec(`INSERT INTO OutgoingWebhooks
|
||||
(Id, Token, CreateAt, UpdateAt, DeleteAt, CreatorId, ChannelId, TeamId, TriggerWords, TriggerWhen,
|
||||
CallbackURLs, DisplayName, Description, ContentType, Username, IconURL)
|
||||
VALUES
|
||||
(:Id, :Token, :CreateAt, :UpdateAt, :DeleteAt, :CreatorId, :ChannelId, :TeamId, :TriggerWords, :TriggerWhen,
|
||||
:CallbackURLs, :DisplayName, :Description, :ContentType, :Username, :IconURL)`, webhook); err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to save OutgoingWebhook with id=%s", webhook.Id)
|
||||
}
|
||||
|
||||
return webhook, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetOutgoing(id string) (*model.OutgoingWebhook, error) {
|
||||
|
||||
var webhook model.OutgoingWebhook
|
||||
|
||||
if err := s.GetReplicaX().Get(&webhook, "SELECT * FROM OutgoingWebhooks WHERE Id = ? AND DeleteAt = 0", id); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, store.NewErrNotFound("OutgoingWebhook", id)
|
||||
}
|
||||
|
||||
return nil, errors.Wrapf(err, "failed to get OutgoingWebhook with id=%s", id)
|
||||
}
|
||||
|
||||
return &webhook, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetOutgoingListByUser(userId string, offset, limit int) ([]*model.OutgoingWebhook, error) {
|
||||
webhooks := []*model.OutgoingWebhook{}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("OutgoingWebhooks").
|
||||
Where(sq.And{
|
||||
sq.Eq{"DeleteAt": int(0)},
|
||||
}).Limit(uint64(limit)).Offset(uint64(offset))
|
||||
|
||||
if userId != "" {
|
||||
query = query.Where(sq.Eq{"CreatorId": userId})
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "outgoing_webhook_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&webhooks, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find OutgoingWebhooks")
|
||||
}
|
||||
|
||||
return webhooks, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetOutgoingList(offset, limit int) ([]*model.OutgoingWebhook, error) {
|
||||
return s.GetOutgoingListByUser("", offset, limit)
|
||||
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetOutgoingByChannelByUser(channelId string, userId string, offset, limit int) ([]*model.OutgoingWebhook, error) {
|
||||
webhooks := []*model.OutgoingWebhook{}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("OutgoingWebhooks").
|
||||
Where(sq.And{
|
||||
sq.Eq{"ChannelId": channelId},
|
||||
sq.Eq{"DeleteAt": int(0)},
|
||||
})
|
||||
|
||||
if userId != "" {
|
||||
query = query.Where(sq.Eq{"CreatorId": userId})
|
||||
}
|
||||
if limit >= 0 && offset >= 0 {
|
||||
query = query.Limit(uint64(limit)).Offset(uint64(offset))
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "outgoing_webhook_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&webhooks, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find OutgoingWebhooks")
|
||||
}
|
||||
|
||||
return webhooks, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetOutgoingByChannel(channelId string, offset, limit int) ([]*model.OutgoingWebhook, error) {
|
||||
return s.GetOutgoingByChannelByUser(channelId, "", offset, limit)
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetOutgoingByTeamByUser(teamId string, userId string, offset, limit int) ([]*model.OutgoingWebhook, error) {
|
||||
webhooks := []*model.OutgoingWebhook{}
|
||||
|
||||
query := s.getQueryBuilder().
|
||||
Select("*").
|
||||
From("OutgoingWebhooks").
|
||||
Where(sq.And{
|
||||
sq.Eq{"TeamId": teamId},
|
||||
sq.Eq{"DeleteAt": int(0)},
|
||||
})
|
||||
|
||||
if userId != "" {
|
||||
query = query.Where(sq.Eq{"CreatorId": userId})
|
||||
}
|
||||
if limit >= 0 && offset >= 0 {
|
||||
query = query.Limit(uint64(limit)).Offset(uint64(offset))
|
||||
}
|
||||
|
||||
queryString, args, err := query.ToSql()
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "outgoing_webhook_tosql")
|
||||
}
|
||||
|
||||
if err := s.GetReplicaX().Select(&webhooks, queryString, args...); err != nil {
|
||||
return nil, errors.Wrap(err, "failed to find OutgoingWebhooks")
|
||||
}
|
||||
|
||||
return webhooks, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) GetOutgoingByTeam(teamId string, offset, limit int) ([]*model.OutgoingWebhook, error) {
|
||||
return s.GetOutgoingByTeamByUser(teamId, "", offset, limit)
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) DeleteOutgoing(webhookId string, time int64) error {
|
||||
_, err := s.GetMasterX().Exec("Update OutgoingWebhooks SET DeleteAt = ?, UpdateAt = ? WHERE Id = ?", time, time, webhookId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to update OutgoingWebhook with id=%s", webhookId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) PermanentDeleteOutgoingByUser(userId string) error {
|
||||
_, err := s.GetMasterX().Exec("DELETE FROM OutgoingWebhooks WHERE CreatorId = ?", userId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete OutgoingWebhook with creatorId=%s", userId)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) PermanentDeleteOutgoingByChannel(channelId string) error {
|
||||
_, err := s.GetMasterX().Exec("DELETE FROM OutgoingWebhooks WHERE ChannelId = ?", channelId)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to delete OutgoingWebhook with channelId=%s", channelId)
|
||||
}
|
||||
|
||||
s.ClearCaches()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) UpdateOutgoing(hook *model.OutgoingWebhook) (*model.OutgoingWebhook, error) {
|
||||
hook.UpdateAt = model.GetMillis()
|
||||
|
||||
_, err := s.GetMasterX().NamedExec(`UPDATE OutgoingWebhooks SET
|
||||
CreateAt = :CreateAt, UpdateAt = :UpdateAt, DeleteAt = :DeleteAt, Token = :Token, CreatorId = :CreatorId,
|
||||
ChannelId = :ChannelId, TeamId = :TeamId, TriggerWords = :TriggerWords, TriggerWhen = :TriggerWhen,
|
||||
CallbackURLs = :CallbackURLs, DisplayName = :DisplayName, Description = :Description,
|
||||
ContentType = :ContentType, Username = :Username, IconURL = :IconURL WHERE Id = :Id`, hook)
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to update OutgoingWebhook with id=%s", hook.Id)
|
||||
}
|
||||
|
||||
return hook, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) AnalyticsIncomingCount(teamId string) (int64, error) {
|
||||
queryBuilder :=
|
||||
s.getQueryBuilder().
|
||||
Select("COUNT(*)").
|
||||
From("IncomingWebhooks").
|
||||
Where("DeleteAt = 0")
|
||||
|
||||
if teamId != "" {
|
||||
queryBuilder = queryBuilder.Where("TeamId", teamId)
|
||||
}
|
||||
|
||||
queryString, args, err := queryBuilder.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "incoming_webhook_tosql")
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := s.GetReplicaX().Get(&count, queryString, args...); err != nil {
|
||||
return 0, errors.Wrap(err, "failed to count IncomingWebhooks")
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (s SqlWebhookStore) AnalyticsOutgoingCount(teamId string) (int64, error) {
|
||||
queryBuilder :=
|
||||
s.getQueryBuilder().
|
||||
Select("COUNT(*)").
|
||||
From("OutgoingWebhooks").
|
||||
Where("DeleteAt = 0")
|
||||
|
||||
if teamId != "" {
|
||||
queryBuilder = queryBuilder.Where("TeamId", teamId)
|
||||
}
|
||||
|
||||
queryString, args, err := queryBuilder.ToSql()
|
||||
if err != nil {
|
||||
return 0, errors.Wrap(err, "outgoing_webhook_tosql")
|
||||
}
|
||||
|
||||
var count int64
|
||||
if err := s.GetReplicaX().Get(&count, queryString, args...); err != nil {
|
||||
return 0, errors.Wrap(err, "failed to count OutgoingWebhooks")
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
14
server/channels/store/sqlstore/webhook_store_test.go
Обычный файл
14
server/channels/store/sqlstore/webhook_store_test.go
Обычный файл
@@ -0,0 +1,14 @@
|
||||
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
|
||||
// See LICENSE.txt for license information.
|
||||
|
||||
package sqlstore
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/mattermost/mattermost-server/v6/server/channels/store/storetest"
|
||||
)
|
||||
|
||||
func TestWebhookStore(t *testing.T) {
|
||||
StoreTest(t, storetest.TestWebhookStore)
|
||||
}
|
||||
Ссылка в новой задаче
Block a user