Doug Lauder
2023-03-22 17:22:27 -04:00
коммит произвёл GitHub
родитель b61c096497
Коммит c943ed6859
13276 изменённых файлов: 1695615 добавлений и 223189 удалений

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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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 Обычный файл
Просмотреть файл

@@ -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
}

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

@@ -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)
}

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

@@ -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)
}

Разница между файлами не показана из-за своего большого размера Загрузить разницу

Разница между файлами не показана из-за своего большого размера Загрузить разницу

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

@@ -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)
}

Разница между файлами не показана из-за своего большого размера Загрузить разницу

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

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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.")
}
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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 Обычный файл
Просмотреть файл

@@ -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()
}

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

@@ -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 Обычный файл
Просмотреть файл

@@ -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
}

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

@@ -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 Обычный файл
Просмотреть файл

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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 Обычный файл

Разница между файлами не показана из-за своего большого размера Загрузить разницу

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

@@ -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 Обычный файл
Просмотреть файл

@@ -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 Обычный файл
Просмотреть файл

@@ -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)
}

Разница между файлами не показана из-за своего большого размера Загрузить разницу

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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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 Обычный файл
Просмотреть файл

@@ -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()
}

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

@@ -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
}

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

@@ -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 Обычный файл
Просмотреть файл

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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 Обычный файл

Разница между файлами не показана из-за своего большого размера Загрузить разницу

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

@@ -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)
}
}
})
}
}

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

@@ -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
}

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

@@ -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)
}
})
}

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

@@ -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(&noticeStates, 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`, &noticeStates[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(&noticeStates, sql, args...); err != nil {
return nil, errors.Wrapf(err, "failed to get ProductNoticeViewState with userId=%s", userId)
}
return noticeStates, 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 TestProductNoticesStore(t *testing.T) {
StoreTest(t, storetest.TestProductNoticesStore)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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 Обычный файл
Просмотреть файл

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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...)
}

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

@@ -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)
}
})
}

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

@@ -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
}

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

@@ -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 Обычный файл

Разница между файлами не показана из-за своего большого размера Загрузить разницу

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)
})
}
}

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

@@ -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
}

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

@@ -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 Обычный файл

Разница между файлами не показана из-за своего большого размера Загрузить разницу

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

@@ -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)
})
}

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

@@ -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
}

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

@@ -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 Обычный файл

Разница между файлами не показана из-за своего большого размера Загрузить разницу

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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)
}

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

@@ -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
}

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

@@ -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 Обычный файл

Разница между файлами не показана из-за своего большого размера Загрузить разницу

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

@@ -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)
}

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

@@ -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 Обычный файл
Просмотреть файл

@@ -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 Обычный файл
Просмотреть файл

@@ -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)
})
}
}

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

@@ -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
}

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

@@ -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)
}