Migrate store methods to use request.Context instead of context.Context (#24836)

Этот коммит содержится в:
Ben Schumacher
2023-10-11 13:08:55 +02:00
коммит произвёл GitHub
родитель 0d5a8b8841
Коммит 13c05a571f
127 изменённых файлов: 1030 добавлений и 921 удалений

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

@@ -811,7 +811,7 @@ func (s *OpenTracingLayerChannelStore) CreateDirectChannel(userID *model.User, o
return result, err
}
func (s *OpenTracingLayerChannelStore) CreateInitialSidebarCategories(userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) {
func (s *OpenTracingLayerChannelStore) CreateInitialSidebarCategories(c request.CTX, userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ChannelStore.CreateInitialSidebarCategories")
s.Root.Store.SetContext(newCtx)
@@ -820,7 +820,7 @@ func (s *OpenTracingLayerChannelStore) CreateInitialSidebarCategories(userID str
}()
defer span.Finish()
result, err := s.ChannelStore.CreateInitialSidebarCategories(userID, opts)
result, err := s.ChannelStore.CreateInitialSidebarCategories(c, userID, opts)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -3209,7 +3209,7 @@ func (s *OpenTracingLayerComplianceStore) GetAll(offset int, limit int) (model.C
return result, err
}
func (s *OpenTracingLayerComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
func (s *OpenTracingLayerComplianceStore) MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "ComplianceStore.MessageExport")
s.Root.Store.SetContext(newCtx)
@@ -3218,7 +3218,7 @@ func (s *OpenTracingLayerComplianceStore) MessageExport(ctx context.Context, cur
}()
defer span.Finish()
result, resultVar1, err := s.ComplianceStore.MessageExport(ctx, cursor, limit)
result, resultVar1, err := s.ComplianceStore.MessageExport(c, cursor, limit)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -3443,7 +3443,7 @@ func (s *OpenTracingLayerEmojiStore) Delete(emoji *model.Emoji, timestamp int64)
return err
}
func (s *OpenTracingLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error) {
func (s *OpenTracingLayerEmojiStore) Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "EmojiStore.Get")
s.Root.Store.SetContext(newCtx)
@@ -3452,7 +3452,7 @@ func (s *OpenTracingLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCa
}()
defer span.Finish()
result, err := s.EmojiStore.Get(ctx, id, allowFromCache)
result, err := s.EmojiStore.Get(c, id, allowFromCache)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -3461,7 +3461,7 @@ func (s *OpenTracingLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCa
return result, err
}
func (s *OpenTracingLayerEmojiStore) GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error) {
func (s *OpenTracingLayerEmojiStore) GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "EmojiStore.GetByName")
s.Root.Store.SetContext(newCtx)
@@ -3470,7 +3470,7 @@ func (s *OpenTracingLayerEmojiStore) GetByName(ctx request.CTX, name string, all
}()
defer span.Finish()
result, err := s.EmojiStore.GetByName(ctx, name, allowFromCache)
result, err := s.EmojiStore.GetByName(c, name, allowFromCache)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -3497,7 +3497,7 @@ func (s *OpenTracingLayerEmojiStore) GetList(offset int, limit int, sort string)
return result, err
}
func (s *OpenTracingLayerEmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error) {
func (s *OpenTracingLayerEmojiStore) GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "EmojiStore.GetMultipleByName")
s.Root.Store.SetContext(newCtx)
@@ -3506,7 +3506,7 @@ func (s *OpenTracingLayerEmojiStore) GetMultipleByName(ctx request.CTX, names []
}()
defer span.Finish()
result, err := s.EmojiStore.GetMultipleByName(ctx, names)
result, err := s.EmojiStore.GetMultipleByName(c, names)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -5197,7 +5197,7 @@ func (s *OpenTracingLayerJobStore) UpdateStatusOptimistically(id string, current
return result, err
}
func (s *OpenTracingLayerLicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) {
func (s *OpenTracingLayerLicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "LicenseStore.Get")
s.Root.Store.SetContext(newCtx)
@@ -5206,7 +5206,7 @@ func (s *OpenTracingLayerLicenseStore) Get(ctx context.Context, id string) (*mod
}()
defer span.Finish()
result, err := s.LicenseStore.Get(ctx, id)
result, err := s.LicenseStore.Get(c, id)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -8281,7 +8281,7 @@ func (s *OpenTracingLayerSessionStore) Cleanup(expiryTime int64, batchSize int64
return err
}
func (s *OpenTracingLayerSessionStore) Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error) {
func (s *OpenTracingLayerSessionStore) Get(c request.CTX, sessionIDOrToken string) (*model.Session, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "SessionStore.Get")
s.Root.Store.SetContext(newCtx)
@@ -8290,7 +8290,7 @@ func (s *OpenTracingLayerSessionStore) Get(ctx context.Context, sessionIDOrToken
}()
defer span.Finish()
result, err := s.SessionStore.Get(ctx, sessionIDOrToken)
result, err := s.SessionStore.Get(c, sessionIDOrToken)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -8299,7 +8299,7 @@ func (s *OpenTracingLayerSessionStore) Get(ctx context.Context, sessionIDOrToken
return result, err
}
func (s *OpenTracingLayerSessionStore) GetSessions(userID string) ([]*model.Session, error) {
func (s *OpenTracingLayerSessionStore) GetSessions(c *request.Context, userID string) ([]*model.Session, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "SessionStore.GetSessions")
s.Root.Store.SetContext(newCtx)
@@ -8308,7 +8308,7 @@ func (s *OpenTracingLayerSessionStore) GetSessions(userID string) ([]*model.Sess
}()
defer span.Finish()
result, err := s.SessionStore.GetSessions(userID)
result, err := s.SessionStore.GetSessions(c, userID)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -8407,7 +8407,7 @@ func (s *OpenTracingLayerSessionStore) RemoveAllSessions() error {
return err
}
func (s *OpenTracingLayerSessionStore) Save(session *model.Session) (*model.Session, error) {
func (s *OpenTracingLayerSessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "SessionStore.Save")
s.Root.Store.SetContext(newCtx)
@@ -8416,7 +8416,7 @@ func (s *OpenTracingLayerSessionStore) Save(session *model.Session) (*model.Sess
}()
defer span.Finish()
result, err := s.SessionStore.Save(session)
result, err := s.SessionStore.Save(c, session)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -9626,7 +9626,7 @@ func (s *OpenTracingLayerTeamStore) GetMany(ids []string) ([]*model.Team, error)
return result, err
}
func (s *OpenTracingLayerTeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) {
func (s *OpenTracingLayerTeamStore) GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "TeamStore.GetMember")
s.Root.Store.SetContext(newCtx)
@@ -9635,7 +9635,7 @@ func (s *OpenTracingLayerTeamStore) GetMember(ctx context.Context, teamID string
}()
defer span.Finish()
result, err := s.TeamStore.GetMember(ctx, teamID, userID)
result, err := s.TeamStore.GetMember(c, teamID, userID)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -9734,7 +9734,7 @@ func (s *OpenTracingLayerTeamStore) GetTeamsByUserId(userID string) ([]*model.Te
return result, err
}
func (s *OpenTracingLayerTeamStore) GetTeamsForUser(ctx context.Context, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
func (s *OpenTracingLayerTeamStore) GetTeamsForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "TeamStore.GetTeamsForUser")
s.Root.Store.SetContext(newCtx)
@@ -9743,7 +9743,7 @@ func (s *OpenTracingLayerTeamStore) GetTeamsForUser(ctx context.Context, userID
}()
defer span.Finish()
result, err := s.TeamStore.GetTeamsForUser(ctx, userID, excludeTeamID, includeDeleted)
result, err := s.TeamStore.GetTeamsForUser(c, userID, excludeTeamID, includeDeleted)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)
@@ -10840,7 +10840,7 @@ func (s *OpenTracingLayerUploadSessionStore) Delete(id string) error {
return err
}
func (s *OpenTracingLayerUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) {
func (s *OpenTracingLayerUploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) {
origCtx := s.Root.Store.Context()
span, newCtx := tracing.StartSpanWithParentByContext(s.Root.Store.Context(), "UploadSessionStore.Get")
s.Root.Store.SetContext(newCtx)
@@ -10849,7 +10849,7 @@ func (s *OpenTracingLayerUploadSessionStore) Get(ctx context.Context, id string)
}()
defer span.Finish()
result, err := s.UploadSessionStore.Get(ctx, id)
result, err := s.UploadSessionStore.Get(c, id)
if err != nil {
span.LogFields(spanlog.Error(err))
ext.Error.Set(span, true)

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

@@ -871,11 +871,11 @@ func (s *RetryLayerChannelStore) CreateDirectChannel(userID *model.User, otherUs
}
func (s *RetryLayerChannelStore) CreateInitialSidebarCategories(userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) {
func (s *RetryLayerChannelStore) CreateInitialSidebarCategories(c request.CTX, userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) {
tries := 0
for {
result, err := s.ChannelStore.CreateInitialSidebarCategories(userID, opts)
result, err := s.ChannelStore.CreateInitialSidebarCategories(c, userID, opts)
if err == nil {
return result, nil
}
@@ -3577,11 +3577,11 @@ func (s *RetryLayerComplianceStore) GetAll(offset int, limit int) (model.Complia
}
func (s *RetryLayerComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
func (s *RetryLayerComplianceStore) MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
tries := 0
for {
result, resultVar1, err := s.ComplianceStore.MessageExport(ctx, cursor, limit)
result, resultVar1, err := s.ComplianceStore.MessageExport(c, cursor, limit)
if err == nil {
return result, resultVar1, nil
}
@@ -3850,11 +3850,11 @@ func (s *RetryLayerEmojiStore) Delete(emoji *model.Emoji, timestamp int64) error
}
func (s *RetryLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error) {
func (s *RetryLayerEmojiStore) Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error) {
tries := 0
for {
result, err := s.EmojiStore.Get(ctx, id, allowFromCache)
result, err := s.EmojiStore.Get(c, id, allowFromCache)
if err == nil {
return result, nil
}
@@ -3871,11 +3871,11 @@ func (s *RetryLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bo
}
func (s *RetryLayerEmojiStore) GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error) {
func (s *RetryLayerEmojiStore) GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error) {
tries := 0
for {
result, err := s.EmojiStore.GetByName(ctx, name, allowFromCache)
result, err := s.EmojiStore.GetByName(c, name, allowFromCache)
if err == nil {
return result, nil
}
@@ -3913,11 +3913,11 @@ func (s *RetryLayerEmojiStore) GetList(offset int, limit int, sort string) ([]*m
}
func (s *RetryLayerEmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error) {
func (s *RetryLayerEmojiStore) GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error) {
tries := 0
for {
result, err := s.EmojiStore.GetMultipleByName(ctx, names)
result, err := s.EmojiStore.GetMultipleByName(c, names)
if err == nil {
return result, nil
}
@@ -5878,11 +5878,11 @@ func (s *RetryLayerJobStore) UpdateStatusOptimistically(id string, currentStatus
}
func (s *RetryLayerLicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) {
func (s *RetryLayerLicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) {
tries := 0
for {
result, err := s.LicenseStore.Get(ctx, id)
result, err := s.LicenseStore.Get(c, id)
if err == nil {
return result, nil
}
@@ -9430,11 +9430,11 @@ func (s *RetryLayerSessionStore) Cleanup(expiryTime int64, batchSize int64) erro
}
func (s *RetryLayerSessionStore) Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error) {
func (s *RetryLayerSessionStore) Get(c request.CTX, sessionIDOrToken string) (*model.Session, error) {
tries := 0
for {
result, err := s.SessionStore.Get(ctx, sessionIDOrToken)
result, err := s.SessionStore.Get(c, sessionIDOrToken)
if err == nil {
return result, nil
}
@@ -9451,11 +9451,11 @@ func (s *RetryLayerSessionStore) Get(ctx context.Context, sessionIDOrToken strin
}
func (s *RetryLayerSessionStore) GetSessions(userID string) ([]*model.Session, error) {
func (s *RetryLayerSessionStore) GetSessions(c *request.Context, userID string) ([]*model.Session, error) {
tries := 0
for {
result, err := s.SessionStore.GetSessions(userID)
result, err := s.SessionStore.GetSessions(c, userID)
if err == nil {
return result, nil
}
@@ -9577,11 +9577,11 @@ func (s *RetryLayerSessionStore) RemoveAllSessions() error {
}
func (s *RetryLayerSessionStore) Save(session *model.Session) (*model.Session, error) {
func (s *RetryLayerSessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) {
tries := 0
for {
result, err := s.SessionStore.Save(session)
result, err := s.SessionStore.Save(c, session)
if err == nil {
return result, nil
}
@@ -10990,11 +10990,11 @@ func (s *RetryLayerTeamStore) GetMany(ids []string) ([]*model.Team, error) {
}
func (s *RetryLayerTeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) {
func (s *RetryLayerTeamStore) GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) {
tries := 0
for {
result, err := s.TeamStore.GetMember(ctx, teamID, userID)
result, err := s.TeamStore.GetMember(c, teamID, userID)
if err == nil {
return result, nil
}
@@ -11116,11 +11116,11 @@ func (s *RetryLayerTeamStore) GetTeamsByUserId(userID string) ([]*model.Team, er
}
func (s *RetryLayerTeamStore) GetTeamsForUser(ctx context.Context, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
func (s *RetryLayerTeamStore) GetTeamsForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
tries := 0
for {
result, err := s.TeamStore.GetTeamsForUser(ctx, userID, excludeTeamID, includeDeleted)
result, err := s.TeamStore.GetTeamsForUser(c, userID, excludeTeamID, includeDeleted)
if err == nil {
return result, nil
}
@@ -12388,11 +12388,11 @@ func (s *RetryLayerUploadSessionStore) Delete(id string) error {
}
func (s *RetryLayerUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) {
func (s *RetryLayerUploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) {
tries := 0
for {
result, err := s.UploadSessionStore.Get(ctx, id)
result, err := s.UploadSessionStore.Get(c, id)
if err == nil {
return result, nil
}

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

@@ -4,13 +4,13 @@
package sqlstore
import (
"context"
"fmt"
sq "github.com/mattermost/squirrel"
"github.com/pkg/errors"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -20,14 +20,14 @@ type dbSelecter interface {
Select(i any, query string, args ...any) error
}
func (s SqlChannelStore) CreateInitialSidebarCategories(userId string, opts *store.SidebarCategorySearchOpts) (_ *model.OrderedSidebarCategories, err error) {
func (s SqlChannelStore) CreateInitialSidebarCategories(c request.CTX, userId string, opts *store.SidebarCategorySearchOpts) (_ *model.OrderedSidebarCategories, err error) {
transaction, err := s.GetMasterX().Beginx()
if err != nil {
return nil, errors.Wrap(err, "CreateInitialSidebarCategories: begin_transaction")
}
defer finalizeTransactionX(transaction, &err)
teamsWithExclude, err := s.SqlStore.stores.team.GetTeamsForUser(context.Background(), userId, opts.TeamID, false)
teamsWithExclude, err := s.SqlStore.stores.team.GetTeamsForUser(c, userId, opts.TeamID, false)
if err != nil {
return nil, errors.Wrap(err, "CreateInitialSidebarCategories: GetTeamsForUser")
}

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

@@ -4,7 +4,6 @@
package sqlstore
import (
"context"
"database/sql"
"fmt"
"strings"
@@ -13,6 +12,7 @@ import (
"github.com/pkg/errors"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -271,7 +271,7 @@ func (s SqlComplianceStore) ComplianceExport(job *model.Compliance, cursor model
return append(channelPosts, directMessagePosts...), cursor, nil
}
func (s SqlComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
func (s SqlComplianceStore) MessageExport(c request.CTX, 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 :=
@@ -318,7 +318,7 @@ func (s SqlComplianceStore) MessageExport(ctx context.Context, cursor model.Mess
LIMIT ?`
cposts := []*model.MessageExport{}
if err := s.GetReplicaX().SelectCtx(ctx, &cposts, query, args...); err != nil {
if err := s.GetReplicaX().SelectCtx(c.Context(), &cposts, query, args...); err != nil {
return nil, cursor, errors.Wrap(err, "unable to export messages")
}
if len(cposts) > 0 {

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

@@ -35,8 +35,8 @@ func RequestContextWithMaster(c request.CTX) request.CTX {
return c
}
// hasMaster is a helper function to check whether master DB should be selected or not.
func hasMaster(ctx context.Context) bool {
// 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
@@ -47,7 +47,7 @@ func hasMaster(ctx context.Context) bool {
// 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) {
if HasMaster(ctx) {
return ss.GetMasterX()
}
return ss.GetReplicaX()

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

@@ -14,5 +14,5 @@ func TestContextMaster(t *testing.T) {
ctx := context.Background()
m := WithMaster(ctx)
assert.True(t, hasMaster(m))
assert.True(t, HasMaster(m))
}

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

@@ -9,6 +9,7 @@ import (
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -316,10 +317,10 @@ func createScheme(ss store.Store) *model.Scheme {
return s
}
func createSession(ss store.Store, userId string) *model.Session {
func createSession(c *request.Context, ss store.Store, userId string) *model.Session {
m := model.Session{}
m.UserId = userId
s, _ := ss.Session().Save(&m)
s, _ := ss.Session().Save(c, &m)
return s
}
@@ -762,6 +763,7 @@ func TestCheckSchemesTeamsIntegrity(t *testing.T) {
func TestCheckSessionsAuditsIntegrity(t *testing.T) {
StoreTest(t, func(t *testing.T, ss store.Store) {
c := request.TestContext(t)
store := ss.(*SqlStore)
dbmap := store.GetMasterX()
@@ -774,7 +776,7 @@ func TestCheckSessionsAuditsIntegrity(t *testing.T) {
t.Run("should generate a report with one record", func(t *testing.T) {
userId := model.NewId()
session := createSession(ss, model.NewId())
session := createSession(c, ss, model.NewId())
sessionId := session.Id
audit := createAudit(ss, userId, sessionId)
dbmap.Exec(`DELETE FROM Sessions WHERE Id=?`, session.Id)
@@ -1492,6 +1494,7 @@ func TestCheckUsersReactionsIntegrity(t *testing.T) {
func TestCheckUsersSessionsIntegrity(t *testing.T) {
StoreTest(t, func(t *testing.T, ss store.Store) {
c := request.TestContext(t)
store := ss.(*SqlStore)
dbmap := store.GetMasterX()
@@ -1504,7 +1507,7 @@ func TestCheckUsersSessionsIntegrity(t *testing.T) {
t.Run("should generate a report with one record", func(t *testing.T) {
userId := model.NewId()
session := createSession(ss, userId)
session := createSession(c, ss, userId)
result := checkUsersSessionsIntegrity(store)
require.NoError(t, result.Err)
data := result.Data.(model.RelationalIntegrityCheckData)

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

@@ -4,12 +4,11 @@
package sqlstore
import (
"context"
sq "github.com/mattermost/squirrel"
"github.com/pkg/errors"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -58,7 +57,7 @@ func (ls SqlLicenseStore) Save(license *model.LicenseRecord) error {
// 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(ctx context.Context, id string) (*model.LicenseRecord, error) {
func (ls SqlLicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) {
query := ls.getQueryBuilder().
Select("Id, CreateAt, Bytes").
From("Licenses").
@@ -70,7 +69,7 @@ func (ls SqlLicenseStore) Get(ctx context.Context, id string) (*model.LicenseRec
}
license := &model.LicenseRecord{}
if err := ls.DBXFromContext(ctx).Get(license, queryString, args...); err != nil {
if err := ls.DBXFromContext(c.Context()).Get(license, queryString, args...); err != nil {
return nil, store.NewErrNotFound("License", id)
}
return license, nil

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

@@ -4,7 +4,6 @@
package sqlstore
import (
"context"
"encoding/json"
"fmt"
"time"
@@ -13,6 +12,7 @@ import (
"github.com/pkg/errors"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -28,7 +28,7 @@ func newSqlSessionStore(sqlStore *SqlStore) store.SessionStore {
return &SqlSessionStore{sqlStore}
}
func (me SqlSessionStore) Save(session *model.Session) (*model.Session, error) {
func (me SqlSessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) {
if session.Id != "" {
return nil, store.NewErrInvalidInput("Session", "id", session.Id)
}
@@ -59,7 +59,7 @@ func (me SqlSessionStore) Save(session *model.Session) (*model.Session, error) {
return nil, errors.Wrapf(err, "failed to save Session with id=%s", session.Id)
}
teamMembers, err := me.Team().GetTeamsForUser(context.Background(), session.UserId, "", true)
teamMembers, err := me.Team().GetTeamsForUser(c, session.UserId, "", true)
if err != nil {
return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", session.UserId)
}
@@ -74,10 +74,10 @@ func (me SqlSessionStore) Save(session *model.Session) (*model.Session, error) {
return session, nil
}
func (me SqlSessionStore) Get(ctx context.Context, sessionIdOrToken string) (*model.Session, error) {
func (me SqlSessionStore) Get(c request.CTX, 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 {
if err := me.DBXFromContext(c.Context()).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 {
@@ -86,7 +86,7 @@ func (me SqlSessionStore) Get(ctx context.Context, sessionIdOrToken string) (*mo
session := sessions[0]
tempMembers, err := me.Team().GetTeamsForUser(
WithMaster(context.Background()),
RequestContextWithMaster(c),
session.UserId, "", true)
if err != nil {
return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", session.UserId)
@@ -100,14 +100,14 @@ func (me SqlSessionStore) Get(ctx context.Context, sessionIdOrToken string) (*mo
return session, nil
}
func (me SqlSessionStore) GetSessions(userId string) ([]*model.Session, error) {
func (me SqlSessionStore) GetSessions(c *request.Context, 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)
teamMembers, err := me.Team().GetTeamsForUser(c, userId, "", true)
if err != nil {
return nil, errors.Wrapf(err, "failed to find TeamMembers for Session with userId=%s", userId)
}

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

@@ -4,7 +4,6 @@
package sqlstore
import (
"context"
"database/sql"
"fmt"
"strings"
@@ -13,6 +12,7 @@ import (
"github.com/pkg/errors"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
"github.com/mattermost/mattermost/server/v8/channels/utils"
)
@@ -945,7 +945,7 @@ func (s SqlTeamStore) UpdateMember(member *model.TeamMember) (*model.TeamMember,
}
// GetMember returns a single member of the team that matches the teamId and userId provided as parameters.
func (s SqlTeamStore) GetMember(ctx context.Context, teamId string, userId string) (*model.TeamMember, error) {
func (s SqlTeamStore) GetMember(ctx request.CTX, teamId string, userId string) (*model.TeamMember, error) {
query := s.getTeamMembersWithSchemeSelectQuery().
Where(sq.Eq{"TeamMembers.TeamId": teamId}).
Where(sq.Eq{"TeamMembers.UserId": userId})
@@ -956,7 +956,7 @@ func (s SqlTeamStore) GetMember(ctx context.Context, teamId string, userId strin
}
var dbMember teamMemberWithSchemeRoles
err = s.DBXFromContext(ctx).Get(&dbMember, queryString, args...)
err = s.DBXFromContext(ctx.Context()).Get(&dbMember, queryString, args...)
if err != nil {
if err == sql.ErrNoRows {
return nil, store.NewErrNotFound("TeamMember", fmt.Sprintf("teamId=%s, userId=%s", teamId, userId))
@@ -1092,7 +1092,7 @@ func (s SqlTeamStore) GetMembersByIds(teamId string, userIds []string, restricti
}
// GetTeamsForUser returns a list of teams that the user is a member of. Expects userId to be passed as a parameter. It can also negative the teamID passed.
func (s SqlTeamStore) GetTeamsForUser(ctx context.Context, userId, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
func (s SqlTeamStore) GetTeamsForUser(ctx request.CTX, userId, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
query := s.getTeamMembersWithSchemeSelectQuery().
Where(sq.Eq{"TeamMembers.UserId": userId})
@@ -1110,7 +1110,7 @@ func (s SqlTeamStore) GetTeamsForUser(ctx context.Context, userId, excludeTeamID
}
dbMembers := teamMemberWithSchemeRolesList{}
err = s.SqlStore.DBXFromContext(ctx).Select(&dbMembers, queryString, args...)
err = s.SqlStore.DBXFromContext(ctx.Context()).Select(&dbMembers, queryString, args...)
if err != nil {
return nil, errors.Wrapf(err, "failed to find TeamMembers with userId=%s", userId)
}

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

@@ -4,13 +4,13 @@
package sqlstore
import (
"context"
"database/sql"
sq "github.com/mattermost/squirrel"
"github.com/pkg/errors"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -79,7 +79,7 @@ func (us SqlUploadSessionStore) Update(session *model.UploadSession) error {
return nil
}
func (us SqlUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) {
func (us SqlUploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) {
if !model.IsValidId(id) {
return nil, errors.New("SqlUploadSessionStore.Get: id is not valid")
}
@@ -92,7 +92,7 @@ func (us SqlUploadSessionStore) Get(ctx context.Context, id string) (*model.Uplo
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 := us.DBXFromContext(c.Context()).Get(&session, query, args...); err != nil {
if err == sql.ErrNoRows {
return nil, store.NewErrNotFound("UploadSession", id)
}

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

@@ -138,12 +138,12 @@ type TeamStore interface {
SaveMember(member *model.TeamMember, maxUsersPerTeam int) (*model.TeamMember, error)
UpdateMember(member *model.TeamMember) (*model.TeamMember, error)
UpdateMultipleMembers(members []*model.TeamMember) ([]*model.TeamMember, error)
GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error)
GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error)
GetMembers(teamID string, offset int, limit int, teamMembersGetOptions *model.TeamMembersGetOptions) ([]*model.TeamMember, error)
GetMembersByIds(teamID string, userIds []string, restrictions *model.ViewUsersRestrictions) ([]*model.TeamMember, error)
GetTotalMemberCount(teamID string, restrictions *model.ViewUsersRestrictions) (int64, error)
GetActiveMemberCount(teamID string, restrictions *model.ViewUsersRestrictions) (int64, error)
GetTeamsForUser(ctx context.Context, userID, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error)
GetTeamsForUser(c request.CTX, userID, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error)
GetTeamsForUserWithPagination(userID string, page, perPage int) ([]*model.TeamMember, error)
GetChannelUnreadsForAllTeams(excludeTeamID, userID string) ([]*model.ChannelUnread, error)
GetChannelUnreadsForTeam(teamID, userID string) ([]*model.ChannelUnread, error)
@@ -278,7 +278,7 @@ type ChannelStore interface {
MigrateChannelMembers(fromChannelID string, fromUserID string) (map[string]string, error)
ResetAllChannelSchemes() error
ClearAllCustomRoleAssignments() error
CreateInitialSidebarCategories(userID string, opts *SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error)
CreateInitialSidebarCategories(c request.CTX, userID string, opts *SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error)
GetSidebarCategoriesForTeamForUser(userID, teamID string) (*model.OrderedSidebarCategories, error)
GetSidebarCategories(userID string, opts *SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error)
GetSidebarCategory(categoryID string) (*model.SidebarCategoryWithChannels, error)
@@ -488,9 +488,9 @@ type BotStore interface {
}
type SessionStore interface {
Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error)
Save(session *model.Session) (*model.Session, error)
GetSessions(userID string) ([]*model.Session, error)
Get(c request.CTX, sessionIDOrToken string) (*model.Session, error)
Save(c request.CTX, session *model.Session) (*model.Session, error)
GetSessions(c *request.Context, userID string) ([]*model.Session, error)
GetSessionsWithActiveDeviceIds(userID string) ([]*model.Session, error)
GetSessionsExpired(thresholdMillis int64, mobileOnly bool, unnotifiedOnly bool) ([]*model.Session, error)
UpdateExpiredNotify(sessionid string, notified bool) error
@@ -537,7 +537,7 @@ type ComplianceStore interface {
Get(id string) (*model.Compliance, error)
GetAll(offset, limit int) (model.Compliances, error)
ComplianceExport(compliance *model.Compliance, cursor model.ComplianceExportCursor, limit int) ([]*model.CompliancePost, model.ComplianceExportCursor, error)
MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error)
MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error)
}
type OAuthStore interface {
@@ -642,7 +642,7 @@ type PreferenceStore interface {
type LicenseStore interface {
Save(license *model.LicenseRecord) error
Get(ctx context.Context, id string) (*model.LicenseRecord, error)
Get(c request.CTX, id string) (*model.LicenseRecord, error)
GetAll() ([]*model.LicenseRecord, error)
}
@@ -665,9 +665,9 @@ type DesktopTokensStore interface {
type EmojiStore interface {
Save(emoji *model.Emoji) (*model.Emoji, error)
Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error)
GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error)
GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error)
Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error)
GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error)
GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error)
GetList(offset, limit int, sort string) ([]*model.Emoji, error)
Delete(emoji *model.Emoji, timestamp int64) error
Search(name string, prefixOnly bool, limit int) ([]*model.Emoji, error)
@@ -712,7 +712,7 @@ type FileInfoStore interface {
type UploadSessionStore interface {
Save(session *model.UploadSession) (*model.UploadSession, error)
Update(session *model.UploadSession) error
Get(ctx context.Context, id string) (*model.UploadSession, error)
Get(c request.CTX, id string) (*model.UploadSession, error)
GetForUser(userID string) ([]*model.UploadSession, error)
Delete(id string) error
}

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

@@ -13,6 +13,7 @@ import (
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -53,6 +54,8 @@ func setupTeam(t *testing.T, ss store.Store, userIds ...string) *model.Team {
}
func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
c := request.TestContext(t)
t.Run("should create initial favorites/channels/DMs categories", func(t *testing.T) {
userId := model.NewId()
@@ -63,7 +66,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
assert.NoError(t, nErr)
require.Len(t, res.Categories, 3)
assert.Equal(t, model.SidebarCategoryFavorites, res.Categories[0].Type)
@@ -85,11 +88,11 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
res, nErr = ss.Channel().CreateInitialSidebarCategories(userId2, opts)
res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId2, opts)
assert.NoError(t, nErr)
assert.Len(t, res.Categories, 3)
assert.Equal(t, model.SidebarCategoryFavorites, res.Categories[0].Type)
@@ -111,7 +114,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -119,7 +122,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team2.Id,
ExcludeTeam: false,
}
res, nErr = ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
assert.NoError(t, nErr)
assert.Len(t, res.Categories, 3)
assert.Equal(t, model.SidebarCategoryFavorites, res.Categories[0].Type)
@@ -140,7 +143,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -149,7 +152,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
require.Equal(t, res, initialCategories)
// Calling CreateInitialSidebarCategories a second time shouldn't create any new categories
res, nErr = ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
assert.NoError(t, nErr)
assert.NotEmpty(t, res)
@@ -175,7 +178,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
_, _ = ss.Channel().CreateInitialSidebarCategories(userId, opts)
_, _ = ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
}()
}
@@ -233,7 +236,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
categories, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
categories, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.Len(t, categories.Categories, 3)
assert.Equal(t, model.SidebarCategoryFavorites, categories.Categories[0].Type)
@@ -302,7 +305,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
categories, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
categories, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.Len(t, categories.Categories, 3)
assert.Equal(t, model.SidebarCategoryFavorites, categories.Categories[0].Type)
@@ -370,7 +373,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
categories, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
categories, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.Len(t, categories.Categories, 3)
assert.Equal(t, model.SidebarCategoryFavorites, categories.Categories[0].Type)
@@ -419,7 +422,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
categories, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
categories, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.Len(t, categories.Categories, 3)
assert.Equal(t, model.SidebarCategoryFavorites, categories.Categories[0].Type)
@@ -468,7 +471,7 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
TeamID: t1.Id,
ExcludeTeam: true,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -479,6 +482,8 @@ func testCreateInitialSidebarCategories(t *testing.T, ss store.Store) {
}
func testCreateSidebarCategory(t *testing.T, ss store.Store) {
c := request.TestContext(t)
t.Run("Creating category without initial categories should fail", func(t *testing.T) {
userId := model.NewId()
teamId := model.NewId()
@@ -504,7 +509,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -534,7 +539,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -575,7 +580,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -617,7 +622,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -682,7 +687,7 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
// Create the category
@@ -707,6 +712,8 @@ func testCreateSidebarCategory(t *testing.T, ss store.Store) {
}
func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
c := request.TestContext(t)
t.Run("should return a custom category with its Channels field set", func(t *testing.T) {
userId := model.NewId()
team := setupTeam(t, ss, userId)
@@ -719,7 +726,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -753,7 +760,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -821,7 +828,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -864,7 +871,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
ExcludeTeam: false,
}
// Create the initial categories and find the channels category
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -931,7 +938,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -976,7 +983,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1018,7 +1025,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1052,7 +1059,7 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
TeamID: otherTeam.Id,
ExcludeTeam: false,
}
res, nErr = ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1075,6 +1082,8 @@ func testGetSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
}
func testGetSidebarCategories(t *testing.T, ss store.Store) {
c := request.TestContext(t)
t.Run("should return channels in the same order between different ways of getting categories", func(t *testing.T) {
userId := model.NewId()
team := setupTeam(t, ss, userId)
@@ -1083,7 +1092,7 @@ func testGetSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1140,7 +1149,7 @@ func testGetSidebarCategories(t *testing.T, ss store.Store) {
}
for _, id := range teamIds {
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, &store.SidebarCategorySearchOpts{TeamID: id})
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, &store.SidebarCategorySearchOpts{TeamID: id})
require.NoError(t, nErr)
require.NotEmpty(t, res)
}
@@ -1176,6 +1185,8 @@ func testGetSidebarCategories(t *testing.T, ss store.Store) {
}
func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
c := request.TestContext(t)
t.Run("ensure the query to update SidebarCategories hasn't been polluted by UpdateSidebarCategoryOrder", func(t *testing.T) {
userId := model.NewId()
team := setupTeam(t, ss, userId)
@@ -1185,7 +1196,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, err := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, err := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, err)
require.NotEmpty(t, res)
@@ -1223,7 +1234,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, err := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, err := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, err)
require.NotEmpty(t, res)
@@ -1254,7 +1265,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1325,7 +1336,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1390,7 +1401,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1461,7 +1472,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1475,7 +1486,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team2.Id,
ExcludeTeam: false,
}
res, nErr = ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1570,7 +1581,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1583,7 +1594,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
require.Equal(t, model.SidebarCategoryChannels, channelsCategory.Type)
// Create the other users' categories
res, nErr = ss.Channel().CreateInitialSidebarCategories(userId2, opts)
res, nErr = ss.Channel().CreateInitialSidebarCategories(c, userId2, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1743,7 +1754,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1802,7 +1813,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1894,7 +1905,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -1962,7 +1973,7 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -2017,6 +2028,8 @@ func testUpdateSidebarCategories(t *testing.T, ss store.Store) {
}
func setupInitialSidebarCategories(t *testing.T, ss store.Store) (string, string) {
c := request.TestContext(t)
userId := model.NewId()
team := setupTeam(t, ss, userId)
@@ -2024,7 +2037,7 @@ func setupInitialSidebarCategories(t *testing.T, ss store.Store) (string, string
TeamID: team.Id,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -2036,6 +2049,8 @@ func setupInitialSidebarCategories(t *testing.T, ss store.Store) (string, string
}
func testClearSidebarOnTeamLeave(t *testing.T, ss store.Store, s SqlStore) {
c := request.TestContext(t)
t.Run("should delete all sidebar categories and channels on the team", func(t *testing.T) {
userId, teamId := setupInitialSidebarCategories(t, ss)
@@ -2151,7 +2166,7 @@ func testClearSidebarOnTeamLeave(t *testing.T, ss store.Store, s SqlStore) {
TeamID: team2.Id,
ExcludeTeam: false,
}
res, err := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, err := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, err)
require.NotEmpty(t, res)
@@ -2334,6 +2349,8 @@ func testDeleteSidebarCategory(t *testing.T, ss store.Store, s SqlStore) {
}
func testUpdateSidebarChannelsByPreferences(t *testing.T, ss store.Store) {
c := request.TestContext(t)
t.Run("Should be able to update sidebar channels", func(t *testing.T) {
userId := model.NewId()
teamId := model.NewId()
@@ -2342,7 +2359,7 @@ func testUpdateSidebarChannelsByPreferences(t *testing.T, ss store.Store) {
TeamID: teamId,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
require.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -2371,7 +2388,7 @@ func testUpdateSidebarChannelsByPreferences(t *testing.T, ss store.Store) {
TeamID: teamId,
ExcludeTeam: false,
}
res, nErr := ss.Channel().CreateInitialSidebarCategories(userId, opts)
res, nErr := ss.Channel().CreateInitialSidebarCategories(c, userId, opts)
assert.NoError(t, nErr)
require.NotEmpty(t, res)
@@ -2391,6 +2408,8 @@ func testUpdateSidebarChannelsByPreferences(t *testing.T, ss store.Store) {
// in the hope of triggering a deadlock. This is a best-effort test case, and is not guaranteed
// to catch a bug.
func testSidebarCategoryDeadlock(t *testing.T, ss store.Store) {
c := request.TestContext(t)
userID := model.NewId()
team := setupTeam(t, ss, userID)
@@ -2413,7 +2432,7 @@ func testSidebarCategoryDeadlock(t *testing.T, ss store.Store) {
TeamID: team.Id,
ExcludeTeam: false,
}
res, err := ss.Channel().CreateInitialSidebarCategories(userID, opts)
res, err := ss.Channel().CreateInitialSidebarCategories(c, userID, opts)
require.NoError(t, err)
require.NotEmpty(t, res)

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

@@ -4,7 +4,6 @@
package storetest
import (
"context"
"encoding/json"
"testing"
"time"
@@ -13,6 +12,7 @@ import (
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -396,11 +396,13 @@ func testComplianceExportDirectMessages(t *testing.T, ss store.Store) {
}
func testMessageExportPublicChannel(t *testing.T, ss store.Store) {
c := request.TestContext(t)
defer cleanupStoreState(t, ss)
// get the starting number of message export entries
startTime := model.GetMillis()
messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
require.NoError(t, err)
assert.Equal(t, 0, len(messages))
@@ -470,7 +472,7 @@ func testMessageExportPublicChannel(t *testing.T, ss store.Store) {
// fetch the message exports for both posts that user1 sent
messageExportMap := map[string]model.MessageExport{}
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
require.NoError(t, err)
assert.Equal(t, 2, len(messages))
@@ -500,11 +502,13 @@ func testMessageExportPublicChannel(t *testing.T, ss store.Store) {
}
func testMessageExportPrivateChannel(t *testing.T, ss store.Store) {
c := request.TestContext(t)
defer cleanupStoreState(t, ss)
// get the starting number of message export entries
startTime := model.GetMillis()
messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
require.NoError(t, err)
assert.Equal(t, 0, len(messages))
@@ -574,7 +578,7 @@ func testMessageExportPrivateChannel(t *testing.T, ss store.Store) {
// fetch the message exports for both posts that user1 sent
messageExportMap := map[string]model.MessageExport{}
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
require.NoError(t, err)
assert.Equal(t, 2, len(messages))
@@ -606,11 +610,13 @@ func testMessageExportPrivateChannel(t *testing.T, ss store.Store) {
}
func testMessageExportDirectMessageChannel(t *testing.T, ss store.Store) {
c := request.TestContext(t)
defer cleanupStoreState(t, ss)
// get the starting number of message export entries
startTime := model.GetMillis()
messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
require.NoError(t, err)
assert.Equal(t, 0, len(messages))
@@ -665,7 +671,7 @@ func testMessageExportDirectMessageChannel(t *testing.T, ss store.Store) {
// fetch the message export for the post that user1 sent
messageExportMap := map[string]model.MessageExport{}
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
require.NoError(t, err)
assert.Equal(t, 1, len(messages))
@@ -687,11 +693,13 @@ func testMessageExportDirectMessageChannel(t *testing.T, ss store.Store) {
}
func testMessageExportGroupMessageChannel(t *testing.T, ss store.Store) {
c := request.TestContext(t)
defer cleanupStoreState(t, ss)
// get the starting number of message export entries
startTime := model.GetMillis()
messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
require.NoError(t, err)
assert.Equal(t, 0, len(messages))
@@ -763,7 +771,7 @@ func testMessageExportGroupMessageChannel(t *testing.T, ss store.Store) {
// fetch the message export for the post that user1 sent
messageExportMap := map[string]model.MessageExport{}
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 10}, 10)
require.NoError(t, err)
assert.Equal(t, 1, len(messages))
@@ -785,10 +793,13 @@ func testMessageExportGroupMessageChannel(t *testing.T, ss store.Store) {
// post,edit,export
func testEditExportMessage(t *testing.T, ss store.Store) {
c := request.TestContext(t)
defer cleanupStoreState(t, ss)
// get the starting number of message export entries
startTime := model.GetMillis()
messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 0, len(messages))
@@ -843,7 +854,7 @@ func testEditExportMessage(t *testing.T, ss store.Store) {
require.NoError(t, err)
// fetch the message exports from the start
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 2, len(messages))
@@ -877,10 +888,12 @@ func testEditExportMessage(t *testing.T, ss store.Store) {
// post, export, edit, export
func testEditAfterExportMessage(t *testing.T, ss store.Store) {
c := request.TestContext(t)
defer cleanupStoreState(t, ss)
// get the starting number of message export entries
startTime := model.GetMillis()
messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 0, len(messages))
@@ -928,7 +941,7 @@ func testEditAfterExportMessage(t *testing.T, ss store.Store) {
require.NoError(t, err)
// fetch the message exports from the start
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 1, len(messages))
@@ -954,7 +967,7 @@ func testEditAfterExportMessage(t *testing.T, ss store.Store) {
require.NoError(t, err)
// fetch the message exports after edit
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: postEditTime - 1}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: postEditTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 2, len(messages))
@@ -988,10 +1001,12 @@ func testEditAfterExportMessage(t *testing.T, ss store.Store) {
// post, delete, export
func testDeleteExportMessage(t *testing.T, ss store.Store) {
c := request.TestContext(t)
defer cleanupStoreState(t, ss)
// get the starting number of message export entries
startTime := model.GetMillis()
messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 0, len(messages))
@@ -1044,7 +1059,7 @@ func testDeleteExportMessage(t *testing.T, ss store.Store) {
require.NoError(t, err)
// fetch the message exports from the start
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 1, len(messages))
@@ -1073,10 +1088,12 @@ func testDeleteExportMessage(t *testing.T, ss store.Store) {
// post,export,delete,export
func testDeleteAfterExportMessage(t *testing.T, ss store.Store) {
c := request.TestContext(t)
defer cleanupStoreState(t, ss)
// get the starting number of message export entries
startTime := model.GetMillis()
messages, _, err := ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
messages, _, err := ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 0, len(messages))
@@ -1124,7 +1141,7 @@ func testDeleteAfterExportMessage(t *testing.T, ss store.Store) {
require.NoError(t, err)
// fetch the message exports from the start
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: startTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 1, len(messages))
@@ -1147,7 +1164,7 @@ func testDeleteAfterExportMessage(t *testing.T, ss store.Store) {
require.NoError(t, err)
// fetch the message exports after delete
messages, _, err = ss.Compliance().MessageExport(context.Background(), model.MessageExportCursor{LastPostUpdateAt: postDeleteTime - 1}, 10)
messages, _, err = ss.Compliance().MessageExport(c, model.MessageExportCursor{LastPostUpdateAt: postDeleteTime - 1}, 10)
require.NoError(t, err)
assert.Equal(t, 1, len(messages))

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

@@ -4,12 +4,12 @@
package storetest
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -36,6 +36,8 @@ func testLicenseStoreSave(t *testing.T, ss store.Store) {
}
func testLicenseStoreGet(t *testing.T, ss store.Store) {
c := request.TestContext(t)
l1 := model.LicenseRecord{}
l1.Id = model.NewId()
l1.Bytes = "junk"
@@ -43,11 +45,11 @@ func testLicenseStoreGet(t *testing.T, ss store.Store) {
err := ss.License().Save(&l1)
require.NoError(t, err)
record, err := ss.License().Get(context.Background(), l1.Id)
record, err := ss.License().Get(c, l1.Id)
require.NoError(t, err, "couldn't get license")
require.Equal(t, record.Bytes, l1.Bytes, "license bytes didn't match")
_, err = ss.License().Get(context.Background(), "missing")
_, err = ss.License().Get(c, "missing")
require.Error(t, err, "should fail on get license")
}

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

@@ -10,6 +10,8 @@ import (
model "github.com/mattermost/mattermost/server/public/model"
mock "github.com/stretchr/testify/mock"
request "github.com/mattermost/mattermost/server/public/shared/request"
store "github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -270,25 +272,25 @@ func (_m *ChannelStore) CreateDirectChannel(userID *model.User, otherUserID *mod
return r0, r1
}
// CreateInitialSidebarCategories provides a mock function with given fields: userID, opts
func (_m *ChannelStore) CreateInitialSidebarCategories(userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) {
ret := _m.Called(userID, opts)
// CreateInitialSidebarCategories provides a mock function with given fields: c, userID, opts
func (_m *ChannelStore) CreateInitialSidebarCategories(c request.CTX, userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) {
ret := _m.Called(c, userID, opts)
var r0 *model.OrderedSidebarCategories
var r1 error
if rf, ok := ret.Get(0).(func(string, *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error)); ok {
return rf(userID, opts)
if rf, ok := ret.Get(0).(func(request.CTX, string, *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error)); ok {
return rf(c, userID, opts)
}
if rf, ok := ret.Get(0).(func(string, *store.SidebarCategorySearchOpts) *model.OrderedSidebarCategories); ok {
r0 = rf(userID, opts)
if rf, ok := ret.Get(0).(func(request.CTX, string, *store.SidebarCategorySearchOpts) *model.OrderedSidebarCategories); ok {
r0 = rf(c, userID, opts)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.OrderedSidebarCategories)
}
}
if rf, ok := ret.Get(1).(func(string, *store.SidebarCategorySearchOpts) error); ok {
r1 = rf(userID, opts)
if rf, ok := ret.Get(1).(func(request.CTX, string, *store.SidebarCategorySearchOpts) error); ok {
r1 = rf(c, userID, opts)
} else {
r1 = ret.Error(1)
}

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

@@ -5,9 +5,8 @@
package mocks
import (
context "context"
model "github.com/mattermost/mattermost/server/public/model"
request "github.com/mattermost/mattermost/server/public/shared/request"
mock "github.com/stretchr/testify/mock"
)
@@ -101,32 +100,32 @@ func (_m *ComplianceStore) GetAll(offset int, limit int) (model.Compliances, err
return r0, r1
}
// MessageExport provides a mock function with given fields: ctx, cursor, limit
func (_m *ComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
ret := _m.Called(ctx, cursor, limit)
// MessageExport provides a mock function with given fields: c, cursor, limit
func (_m *ComplianceStore) MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
ret := _m.Called(c, cursor, limit)
var r0 []*model.MessageExport
var r1 model.MessageExportCursor
var r2 error
if rf, ok := ret.Get(0).(func(context.Context, model.MessageExportCursor, int) ([]*model.MessageExport, model.MessageExportCursor, error)); ok {
return rf(ctx, cursor, limit)
if rf, ok := ret.Get(0).(func(request.CTX, model.MessageExportCursor, int) ([]*model.MessageExport, model.MessageExportCursor, error)); ok {
return rf(c, cursor, limit)
}
if rf, ok := ret.Get(0).(func(context.Context, model.MessageExportCursor, int) []*model.MessageExport); ok {
r0 = rf(ctx, cursor, limit)
if rf, ok := ret.Get(0).(func(request.CTX, model.MessageExportCursor, int) []*model.MessageExport); ok {
r0 = rf(c, cursor, limit)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).([]*model.MessageExport)
}
}
if rf, ok := ret.Get(1).(func(context.Context, model.MessageExportCursor, int) model.MessageExportCursor); ok {
r1 = rf(ctx, cursor, limit)
if rf, ok := ret.Get(1).(func(request.CTX, model.MessageExportCursor, int) model.MessageExportCursor); ok {
r1 = rf(c, cursor, limit)
} else {
r1 = ret.Get(1).(model.MessageExportCursor)
}
if rf, ok := ret.Get(2).(func(context.Context, model.MessageExportCursor, int) error); ok {
r2 = rf(ctx, cursor, limit)
if rf, ok := ret.Get(2).(func(request.CTX, model.MessageExportCursor, int) error); ok {
r2 = rf(c, cursor, limit)
} else {
r2 = ret.Error(2)
}

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

@@ -29,17 +29,17 @@ func (_m *EmojiStore) Delete(emoji *model.Emoji, timestamp int64) error {
return r0
}
// Get provides a mock function with given fields: ctx, id, allowFromCache
func (_m *EmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error) {
ret := _m.Called(ctx, id, allowFromCache)
// Get provides a mock function with given fields: c, id, allowFromCache
func (_m *EmojiStore) Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error) {
ret := _m.Called(c, id, allowFromCache)
var r0 *model.Emoji
var r1 error
if rf, ok := ret.Get(0).(func(request.CTX, string, bool) (*model.Emoji, error)); ok {
return rf(ctx, id, allowFromCache)
return rf(c, id, allowFromCache)
}
if rf, ok := ret.Get(0).(func(request.CTX, string, bool) *model.Emoji); ok {
r0 = rf(ctx, id, allowFromCache)
r0 = rf(c, id, allowFromCache)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.Emoji)
@@ -47,7 +47,7 @@ func (_m *EmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*mod
}
if rf, ok := ret.Get(1).(func(request.CTX, string, bool) error); ok {
r1 = rf(ctx, id, allowFromCache)
r1 = rf(c, id, allowFromCache)
} else {
r1 = ret.Error(1)
}
@@ -55,17 +55,17 @@ func (_m *EmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*mod
return r0, r1
}
// GetByName provides a mock function with given fields: ctx, name, allowFromCache
func (_m *EmojiStore) GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error) {
ret := _m.Called(ctx, name, allowFromCache)
// GetByName provides a mock function with given fields: c, name, allowFromCache
func (_m *EmojiStore) GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error) {
ret := _m.Called(c, name, allowFromCache)
var r0 *model.Emoji
var r1 error
if rf, ok := ret.Get(0).(func(request.CTX, string, bool) (*model.Emoji, error)); ok {
return rf(ctx, name, allowFromCache)
return rf(c, name, allowFromCache)
}
if rf, ok := ret.Get(0).(func(request.CTX, string, bool) *model.Emoji); ok {
r0 = rf(ctx, name, allowFromCache)
r0 = rf(c, name, allowFromCache)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.Emoji)
@@ -73,7 +73,7 @@ func (_m *EmojiStore) GetByName(ctx request.CTX, name string, allowFromCache boo
}
if rf, ok := ret.Get(1).(func(request.CTX, string, bool) error); ok {
r1 = rf(ctx, name, allowFromCache)
r1 = rf(c, name, allowFromCache)
} else {
r1 = ret.Error(1)
}
@@ -107,17 +107,17 @@ func (_m *EmojiStore) GetList(offset int, limit int, sort string) ([]*model.Emoj
return r0, r1
}
// GetMultipleByName provides a mock function with given fields: ctx, names
func (_m *EmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error) {
ret := _m.Called(ctx, names)
// GetMultipleByName provides a mock function with given fields: c, names
func (_m *EmojiStore) GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error) {
ret := _m.Called(c, names)
var r0 []*model.Emoji
var r1 error
if rf, ok := ret.Get(0).(func(request.CTX, []string) ([]*model.Emoji, error)); ok {
return rf(ctx, names)
return rf(c, names)
}
if rf, ok := ret.Get(0).(func(request.CTX, []string) []*model.Emoji); ok {
r0 = rf(ctx, names)
r0 = rf(c, names)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).([]*model.Emoji)
@@ -125,7 +125,7 @@ func (_m *EmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*mod
}
if rf, ok := ret.Get(1).(func(request.CTX, []string) error); ok {
r1 = rf(ctx, names)
r1 = rf(c, names)
} else {
r1 = ret.Error(1)
}

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

@@ -5,9 +5,8 @@
package mocks
import (
context "context"
model "github.com/mattermost/mattermost/server/public/model"
request "github.com/mattermost/mattermost/server/public/shared/request"
mock "github.com/stretchr/testify/mock"
)
@@ -16,25 +15,25 @@ type LicenseStore struct {
mock.Mock
}
// Get provides a mock function with given fields: ctx, id
func (_m *LicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) {
ret := _m.Called(ctx, id)
// Get provides a mock function with given fields: c, id
func (_m *LicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) {
ret := _m.Called(c, id)
var r0 *model.LicenseRecord
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, string) (*model.LicenseRecord, error)); ok {
return rf(ctx, id)
if rf, ok := ret.Get(0).(func(request.CTX, string) (*model.LicenseRecord, error)); ok {
return rf(c, id)
}
if rf, ok := ret.Get(0).(func(context.Context, string) *model.LicenseRecord); ok {
r0 = rf(ctx, id)
if rf, ok := ret.Get(0).(func(request.CTX, string) *model.LicenseRecord); ok {
r0 = rf(c, id)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.LicenseRecord)
}
}
if rf, ok := ret.Get(1).(func(context.Context, string) error); ok {
r1 = rf(ctx, id)
if rf, ok := ret.Get(1).(func(request.CTX, string) error); ok {
r1 = rf(c, id)
} else {
r1 = ret.Error(1)
}

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

@@ -5,9 +5,8 @@
package mocks
import (
context "context"
model "github.com/mattermost/mattermost/server/public/model"
request "github.com/mattermost/mattermost/server/public/shared/request"
mock "github.com/stretchr/testify/mock"
)
@@ -54,25 +53,25 @@ func (_m *SessionStore) Cleanup(expiryTime int64, batchSize int64) error {
return r0
}
// Get provides a mock function with given fields: ctx, sessionIDOrToken
func (_m *SessionStore) Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error) {
ret := _m.Called(ctx, sessionIDOrToken)
// Get provides a mock function with given fields: c, sessionIDOrToken
func (_m *SessionStore) Get(c request.CTX, sessionIDOrToken string) (*model.Session, error) {
ret := _m.Called(c, sessionIDOrToken)
var r0 *model.Session
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, string) (*model.Session, error)); ok {
return rf(ctx, sessionIDOrToken)
if rf, ok := ret.Get(0).(func(request.CTX, string) (*model.Session, error)); ok {
return rf(c, sessionIDOrToken)
}
if rf, ok := ret.Get(0).(func(context.Context, string) *model.Session); ok {
r0 = rf(ctx, sessionIDOrToken)
if rf, ok := ret.Get(0).(func(request.CTX, string) *model.Session); ok {
r0 = rf(c, sessionIDOrToken)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.Session)
}
}
if rf, ok := ret.Get(1).(func(context.Context, string) error); ok {
r1 = rf(ctx, sessionIDOrToken)
if rf, ok := ret.Get(1).(func(request.CTX, string) error); ok {
r1 = rf(c, sessionIDOrToken)
} else {
r1 = ret.Error(1)
}
@@ -80,25 +79,25 @@ func (_m *SessionStore) Get(ctx context.Context, sessionIDOrToken string) (*mode
return r0, r1
}
// GetSessions provides a mock function with given fields: userID
func (_m *SessionStore) GetSessions(userID string) ([]*model.Session, error) {
ret := _m.Called(userID)
// GetSessions provides a mock function with given fields: c, userID
func (_m *SessionStore) GetSessions(c *request.Context, userID string) ([]*model.Session, error) {
ret := _m.Called(c, userID)
var r0 []*model.Session
var r1 error
if rf, ok := ret.Get(0).(func(string) ([]*model.Session, error)); ok {
return rf(userID)
if rf, ok := ret.Get(0).(func(*request.Context, string) ([]*model.Session, error)); ok {
return rf(c, userID)
}
if rf, ok := ret.Get(0).(func(string) []*model.Session); ok {
r0 = rf(userID)
if rf, ok := ret.Get(0).(func(*request.Context, string) []*model.Session); ok {
r0 = rf(c, userID)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).([]*model.Session)
}
}
if rf, ok := ret.Get(1).(func(string) error); ok {
r1 = rf(userID)
if rf, ok := ret.Get(1).(func(*request.Context, string) error); ok {
r1 = rf(c, userID)
} else {
r1 = ret.Error(1)
}
@@ -200,25 +199,25 @@ func (_m *SessionStore) RemoveAllSessions() error {
return r0
}
// Save provides a mock function with given fields: session
func (_m *SessionStore) Save(session *model.Session) (*model.Session, error) {
ret := _m.Called(session)
// Save provides a mock function with given fields: c, session
func (_m *SessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) {
ret := _m.Called(c, session)
var r0 *model.Session
var r1 error
if rf, ok := ret.Get(0).(func(*model.Session) (*model.Session, error)); ok {
return rf(session)
if rf, ok := ret.Get(0).(func(request.CTX, *model.Session) (*model.Session, error)); ok {
return rf(c, session)
}
if rf, ok := ret.Get(0).(func(*model.Session) *model.Session); ok {
r0 = rf(session)
if rf, ok := ret.Get(0).(func(request.CTX, *model.Session) *model.Session); ok {
r0 = rf(c, session)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.Session)
}
}
if rf, ok := ret.Get(1).(func(*model.Session) error); ok {
r1 = rf(session)
if rf, ok := ret.Get(1).(func(request.CTX, *model.Session) error); ok {
r1 = rf(c, session)
} else {
r1 = ret.Error(1)
}

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

@@ -5,9 +5,8 @@
package mocks
import (
context "context"
model "github.com/mattermost/mattermost/server/public/model"
request "github.com/mattermost/mattermost/server/public/shared/request"
mock "github.com/stretchr/testify/mock"
)
@@ -497,25 +496,25 @@ func (_m *TeamStore) GetMany(ids []string) ([]*model.Team, error) {
return r0, r1
}
// GetMember provides a mock function with given fields: ctx, teamID, userID
func (_m *TeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) {
ret := _m.Called(ctx, teamID, userID)
// GetMember provides a mock function with given fields: c, teamID, userID
func (_m *TeamStore) GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) {
ret := _m.Called(c, teamID, userID)
var r0 *model.TeamMember
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, string, string) (*model.TeamMember, error)); ok {
return rf(ctx, teamID, userID)
if rf, ok := ret.Get(0).(func(request.CTX, string, string) (*model.TeamMember, error)); ok {
return rf(c, teamID, userID)
}
if rf, ok := ret.Get(0).(func(context.Context, string, string) *model.TeamMember); ok {
r0 = rf(ctx, teamID, userID)
if rf, ok := ret.Get(0).(func(request.CTX, string, string) *model.TeamMember); ok {
r0 = rf(c, teamID, userID)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.TeamMember)
}
}
if rf, ok := ret.Get(1).(func(context.Context, string, string) error); ok {
r1 = rf(ctx, teamID, userID)
if rf, ok := ret.Get(1).(func(request.CTX, string, string) error); ok {
r1 = rf(c, teamID, userID)
} else {
r1 = ret.Error(1)
}
@@ -653,25 +652,25 @@ func (_m *TeamStore) GetTeamsByUserId(userID string) ([]*model.Team, error) {
return r0, r1
}
// GetTeamsForUser provides a mock function with given fields: ctx, userID, excludeTeamID, includeDeleted
func (_m *TeamStore) GetTeamsForUser(ctx context.Context, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
ret := _m.Called(ctx, userID, excludeTeamID, includeDeleted)
// GetTeamsForUser provides a mock function with given fields: c, userID, excludeTeamID, includeDeleted
func (_m *TeamStore) GetTeamsForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
ret := _m.Called(c, userID, excludeTeamID, includeDeleted)
var r0 []*model.TeamMember
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, string, string, bool) ([]*model.TeamMember, error)); ok {
return rf(ctx, userID, excludeTeamID, includeDeleted)
if rf, ok := ret.Get(0).(func(request.CTX, string, string, bool) ([]*model.TeamMember, error)); ok {
return rf(c, userID, excludeTeamID, includeDeleted)
}
if rf, ok := ret.Get(0).(func(context.Context, string, string, bool) []*model.TeamMember); ok {
r0 = rf(ctx, userID, excludeTeamID, includeDeleted)
if rf, ok := ret.Get(0).(func(request.CTX, string, string, bool) []*model.TeamMember); ok {
r0 = rf(c, userID, excludeTeamID, includeDeleted)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).([]*model.TeamMember)
}
}
if rf, ok := ret.Get(1).(func(context.Context, string, string, bool) error); ok {
r1 = rf(ctx, userID, excludeTeamID, includeDeleted)
if rf, ok := ret.Get(1).(func(request.CTX, string, string, bool) error); ok {
r1 = rf(c, userID, excludeTeamID, includeDeleted)
} else {
r1 = ret.Error(1)
}

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

@@ -5,9 +5,8 @@
package mocks
import (
context "context"
model "github.com/mattermost/mattermost/server/public/model"
request "github.com/mattermost/mattermost/server/public/shared/request"
mock "github.com/stretchr/testify/mock"
)
@@ -30,25 +29,25 @@ func (_m *UploadSessionStore) Delete(id string) error {
return r0
}
// Get provides a mock function with given fields: ctx, id
func (_m *UploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) {
ret := _m.Called(ctx, id)
// Get provides a mock function with given fields: c, id
func (_m *UploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) {
ret := _m.Called(c, id)
var r0 *model.UploadSession
var r1 error
if rf, ok := ret.Get(0).(func(context.Context, string) (*model.UploadSession, error)); ok {
return rf(ctx, id)
if rf, ok := ret.Get(0).(func(request.CTX, string) (*model.UploadSession, error)); ok {
return rf(c, id)
}
if rf, ok := ret.Get(0).(func(context.Context, string) *model.UploadSession); ok {
r0 = rf(ctx, id)
if rf, ok := ret.Get(0).(func(request.CTX, string) *model.UploadSession); ok {
r0 = rf(c, id)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(*model.UploadSession)
}
}
if rf, ok := ret.Get(1).(func(context.Context, string) error); ok {
r1 = rf(ctx, id)
if rf, ok := ret.Get(1).(func(request.CTX, string) error); ok {
r1 = rf(c, id)
} else {
r1 = ret.Error(1)
}

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

@@ -4,13 +4,13 @@
package storetest
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -357,6 +357,8 @@ func testOAuthGetAccessDataByUserForApp(t *testing.T, ss store.Store) {
}
func testOAuthStoreDeleteApp(t *testing.T, ss store.Store) {
c := request.TestContext(t)
a1 := model.OAuthApp{}
a1.CreatorId = model.NewId()
a1.Name = "TestApp" + model.NewId()
@@ -374,7 +376,7 @@ func testOAuthStoreDeleteApp(t *testing.T, ss store.Store) {
s1.Token = model.NewId()
s1.IsOAuth = true
s1, nErr := ss.Session().Save(s1)
s1, nErr := ss.Session().Save(c, s1)
require.NoError(t, nErr)
ad1 := model.AccessData{}
@@ -390,7 +392,7 @@ func testOAuthStoreDeleteApp(t *testing.T, ss store.Store) {
err = ss.OAuth().DeleteApp(a1.Id)
require.NoError(t, err)
_, nErr = ss.Session().Get(context.Background(), s1.Token)
_, nErr = ss.Session().Get(c, s1.Token)
require.Error(t, nErr, "should error - session should be deleted")
_, err = ss.OAuth().GetAccessData(s1.Token)

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

@@ -4,13 +4,13 @@
package storetest
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -39,34 +39,38 @@ func TestSessionStore(t *testing.T, ss store.Store) {
}
func testSessionStoreSave(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
_, err := ss.Session().Save(s1)
_, err := ss.Session().Save(c, s1)
require.NoError(t, err)
}
func testSessionGet(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
s2 := &model.Session{}
s2.UserId = s1.UserId
_, err = ss.Session().Save(s2)
_, err = ss.Session().Save(c, s2)
require.NoError(t, err)
s3 := &model.Session{}
s3.UserId = s1.UserId
s3.ExpiresAt = 1
_, err = ss.Session().Save(s3)
_, err = ss.Session().Save(c, s3)
require.NoError(t, err)
session, err := ss.Session().Get(context.Background(), s1.Id)
session, err := ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.Equal(t, session.Id, s1.Id, "should match")
@@ -75,21 +79,23 @@ func testSessionGet(t *testing.T, ss store.Store) {
err = ss.Session().UpdateProps(session)
require.NoError(t, err)
session2, err := ss.Session().Get(context.Background(), session.Id)
session2, err := ss.Session().Get(c, session.Id)
require.NoError(t, err)
require.Equal(t, session.Props, session2.Props, "should match")
data, err := ss.Session().GetSessions(s1.UserId)
data, err := ss.Session().GetSessions(c, s1.UserId)
require.NoError(t, err)
require.Len(t, data, 3, "should match len")
}
func testSessionGetWithDeviceId(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1.ExpiresAt = model.GetMillis() + 10000
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
s2 := &model.Session{}
@@ -97,7 +103,7 @@ func testSessionGetWithDeviceId(t *testing.T, ss store.Store) {
s2.DeviceId = model.NewId()
s2.ExpiresAt = model.GetMillis() + 10000
_, err = ss.Session().Save(s2)
_, err = ss.Session().Save(c, s2)
require.NoError(t, err)
s3 := &model.Session{}
@@ -105,7 +111,7 @@ func testSessionGetWithDeviceId(t *testing.T, ss store.Store) {
s3.ExpiresAt = 1
s3.DeviceId = model.NewId()
_, err = ss.Session().Save(s3)
_, err = ss.Session().Save(c, s3)
require.NoError(t, err)
data, err := ss.Session().GetSessionsWithActiveDeviceIds(s1.UserId)
@@ -114,86 +120,96 @@ func testSessionGetWithDeviceId(t *testing.T, ss store.Store) {
}
func testSessionRemove(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
session, err := ss.Session().Get(context.Background(), s1.Id)
session, err := ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.Equal(t, session.Id, s1.Id, "should match")
removeErr := ss.Session().Remove(s1.Id)
require.NoError(t, removeErr)
_, err = ss.Session().Get(context.Background(), s1.Id)
_, err = ss.Session().Get(c, s1.Id)
require.Error(t, err, "should have been removed")
}
func testSessionRemoveAll(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
session, err := ss.Session().Get(context.Background(), s1.Id)
session, err := ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.Equal(t, session.Id, s1.Id, "should match")
removeErr := ss.Session().RemoveAllSessions()
require.NoError(t, removeErr)
_, err = ss.Session().Get(context.Background(), s1.Id)
_, err = ss.Session().Get(c, s1.Id)
require.Error(t, err, "should have been removed")
}
func testSessionRemoveByUser(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
session, err := ss.Session().Get(context.Background(), s1.Id)
session, err := ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.Equal(t, session.Id, s1.Id, "should match")
deleteErr := ss.Session().PermanentDeleteSessionsByUser(s1.UserId)
require.NoError(t, deleteErr)
_, err = ss.Session().Get(context.Background(), s1.Id)
_, err = ss.Session().Get(c, s1.Id)
require.Error(t, err, "should have been removed")
}
func testSessionRemoveToken(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
session, err := ss.Session().Get(context.Background(), s1.Id)
session, err := ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.Equal(t, session.Id, s1.Id, "should match")
removeErr := ss.Session().Remove(s1.Token)
require.NoError(t, removeErr)
_, err = ss.Session().Get(context.Background(), s1.Id)
_, err = ss.Session().Get(c, s1.Id)
require.Error(t, err, "should have been removed")
data, err := ss.Session().GetSessions(s1.UserId)
data, err := ss.Session().GetSessions(c, s1.UserId)
require.NoError(t, err)
require.Empty(t, data, "should match len")
}
func testSessionUpdateDeviceId(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
_, err = ss.Session().UpdateDeviceId(s1.Id, model.PushNotifyApple+":1234567890", s1.ExpiresAt)
@@ -202,7 +218,7 @@ func testSessionUpdateDeviceId(t *testing.T, ss store.Store) {
s2 := &model.Session{}
s2.UserId = model.NewId()
s2, err = ss.Session().Save(s2)
s2, err = ss.Session().Save(c, s2)
require.NoError(t, err)
_, err = ss.Session().UpdateDeviceId(s2.Id, model.PushNotifyApple+":1234567890", s1.ExpiresAt)
@@ -210,10 +226,12 @@ func testSessionUpdateDeviceId(t *testing.T, ss store.Store) {
}
func testSessionUpdateDeviceId2(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
_, err = ss.Session().UpdateDeviceId(s1.Id, model.PushNotifyAppleReactNative+":1234567890", s1.ExpiresAt)
@@ -222,7 +240,7 @@ func testSessionUpdateDeviceId2(t *testing.T, ss store.Store) {
s2 := &model.Session{}
s2.UserId = model.NewId()
s2, err = ss.Session().Save(s2)
s2, err = ss.Session().Save(c, s2)
require.NoError(t, err)
_, err = ss.Session().UpdateDeviceId(s2.Id, model.PushNotifyAppleReactNative+":1234567890", s1.ExpiresAt)
@@ -230,41 +248,47 @@ func testSessionUpdateDeviceId2(t *testing.T, ss store.Store) {
}
func testSessionStoreUpdateExpiresAt(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
err = ss.Session().UpdateExpiresAt(s1.Id, 1234567890)
require.NoError(t, err)
session, err := ss.Session().Get(context.Background(), s1.Id)
session, err := ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.EqualValues(t, session.ExpiresAt, 1234567890, "ExpiresAt not updated correctly")
}
func testSessionStoreUpdateLastActivityAt(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
err = ss.Session().UpdateLastActivityAt(s1.Id, 1234567890)
require.NoError(t, err)
session, err := ss.Session().Get(context.Background(), s1.Id)
session, err := ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.EqualValues(t, session.LastActivityAt, 1234567890, "LastActivityAt not updated correctly")
}
func testSessionCount(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1.ExpiresAt = model.GetMillis() + 100000
_, err := ss.Session().Save(s1)
_, err := ss.Session().Save(c, s1)
require.NoError(t, err)
count, err := ss.Session().AnalyticsSessionCount()
@@ -273,49 +297,51 @@ func testSessionCount(t *testing.T, ss store.Store) {
}
func testSessionCleanup(t *testing.T, ss store.Store) {
c := request.TestContext(t)
now := model.GetMillis()
s1 := &model.Session{}
s1.UserId = model.NewId()
s1.ExpiresAt = 0 // never expires
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
s2 := &model.Session{}
s2.UserId = s1.UserId
s2.ExpiresAt = now + 1000000 // expires in the future
s2, err = ss.Session().Save(s2)
s2, err = ss.Session().Save(c, s2)
require.NoError(t, err)
s3 := &model.Session{}
s3.UserId = model.NewId()
s3.ExpiresAt = 1 // expired
s3, err = ss.Session().Save(s3)
s3, err = ss.Session().Save(c, s3)
require.NoError(t, err)
s4 := &model.Session{}
s4.UserId = model.NewId()
s4.ExpiresAt = 2 // expired
s4, err = ss.Session().Save(s4)
s4, err = ss.Session().Save(c, s4)
require.NoError(t, err)
err = ss.Session().Cleanup(now, 1)
require.NoError(t, err)
_, err = ss.Session().Get(context.Background(), s1.Id)
_, err = ss.Session().Get(c, s1.Id)
assert.NoError(t, err)
_, err = ss.Session().Get(context.Background(), s2.Id)
_, err = ss.Session().Get(c, s2.Id)
assert.NoError(t, err)
_, err = ss.Session().Get(context.Background(), s3.Id)
_, err = ss.Session().Get(c, s3.Id)
assert.Error(t, err)
_, err = ss.Session().Get(context.Background(), s4.Id)
_, err = ss.Session().Get(c, s4.Id)
assert.Error(t, err)
removeErr := ss.Session().Remove(s1.Id)
@@ -326,6 +352,8 @@ func testSessionCleanup(t *testing.T, ss store.Store) {
}
func testGetSessionsExpired(t *testing.T, ss store.Store) {
c := request.TestContext(t)
now := model.GetMillis()
// Clear existing sessions.
@@ -336,34 +364,34 @@ func testGetSessionsExpired(t *testing.T, ss store.Store) {
s1.UserId = model.NewId()
s1.DeviceId = model.NewId()
s1.ExpiresAt = 0 // never expires
_, err = ss.Session().Save(s1)
_, err = ss.Session().Save(c, s1)
require.NoError(t, err)
s2 := &model.Session{}
s2.UserId = model.NewId()
s2.DeviceId = model.NewId()
s2.ExpiresAt = now - TenMinutes // expired within threshold
s2, err = ss.Session().Save(s2)
s2, err = ss.Session().Save(c, s2)
require.NoError(t, err)
s3 := &model.Session{}
s3.UserId = model.NewId()
s3.DeviceId = model.NewId()
s3.ExpiresAt = now - (TenMinutes * 100) // expired outside threshold
_, err = ss.Session().Save(s3)
_, err = ss.Session().Save(c, s3)
require.NoError(t, err)
s4 := &model.Session{}
s4.UserId = model.NewId()
s4.ExpiresAt = now - TenMinutes // expired within threshold, but not mobile
s4, err = ss.Session().Save(s4)
s4, err = ss.Session().Save(c, s4)
require.NoError(t, err)
s5 := &model.Session{}
s5.UserId = model.NewId()
s5.DeviceId = model.NewId()
s5.ExpiresAt = now + (TenMinutes * 100000) // not expired
_, err = ss.Session().Save(s5)
_, err = ss.Session().Save(c, s5)
require.NoError(t, err)
sessions, err := ss.Session().GetSessionsExpired(TenMinutes*2, true, true) // mobile only
@@ -381,26 +409,28 @@ func testGetSessionsExpired(t *testing.T, ss store.Store) {
}
func testUpdateExpiredNotify(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := &model.Session{}
s1.UserId = model.NewId()
s1.DeviceId = model.NewId()
s1.ExpiresAt = model.GetMillis() + TenMinutes
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
session, err := ss.Session().Get(context.Background(), s1.Id)
session, err := ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.False(t, session.ExpiredNotify)
err = ss.Session().UpdateExpiredNotify(session.Id, true)
require.NoError(t, err)
session, err = ss.Session().Get(context.Background(), s1.Id)
session, err = ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.True(t, session.ExpiredNotify)
err = ss.Session().UpdateExpiredNotify(session.Id, false)
require.NoError(t, err)
session, err = ss.Session().Get(context.Background(), s1.Id)
session, err = ss.Session().Get(c, s1.Id)
require.NoError(t, err)
require.False(t, session.ExpiredNotify)
}

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

@@ -14,6 +14,7 @@ import (
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -1299,6 +1300,8 @@ func testGetMembers(t *testing.T, ss store.Store) {
}
func testTeamMembers(t *testing.T, ss store.Store) {
c := request.TestContext(t)
teamId1 := model.NewId()
teamId2 := model.NewId()
@@ -1318,8 +1321,7 @@ func testTeamMembers(t *testing.T, ss store.Store) {
require.Len(t, ms, 1)
require.Equal(t, m3.UserId, ms[0].UserId)
ctx := context.Background()
ms, err = ss.Team().GetTeamsForUser(ctx, m1.UserId, "", true)
ms, err = ss.Team().GetTeamsForUser(c, m1.UserId, "", true)
require.NoError(t, err)
require.Len(t, ms, 1)
require.Equal(t, m1.TeamId, ms[0].TeamId)
@@ -1348,11 +1350,11 @@ func testTeamMembers(t *testing.T, ss store.Store) {
_, nErr = ss.Team().SaveMultipleMembers([]*model.TeamMember{m4, m5}, -1)
require.NoError(t, nErr)
ms, err = ss.Team().GetTeamsForUser(ctx, uid, "", true)
ms, err = ss.Team().GetTeamsForUser(c, uid, "", true)
require.NoError(t, err)
require.Len(t, ms, 2)
ms, err = ss.Team().GetTeamsForUser(ctx, uid, teamId2, true)
ms, err = ss.Team().GetTeamsForUser(c, uid, teamId2, true)
require.NoError(t, err)
require.Len(t, ms, 1)
@@ -1360,18 +1362,18 @@ func testTeamMembers(t *testing.T, ss store.Store) {
_, err = ss.Team().UpdateMember(m4)
require.NoError(t, err)
ms, err = ss.Team().GetTeamsForUser(ctx, uid, "", true)
ms, err = ss.Team().GetTeamsForUser(c, uid, "", true)
require.NoError(t, err)
require.Len(t, ms, 2)
ms, err = ss.Team().GetTeamsForUser(ctx, uid, "", false)
ms, err = ss.Team().GetTeamsForUser(c, uid, "", false)
require.NoError(t, err)
require.Len(t, ms, 1)
nErr = ss.Team().RemoveAllMembersByUser(uid)
require.NoError(t, nErr)
ms, err = ss.Team().GetTeamsForUser(ctx, m1.UserId, "", true)
ms, err = ss.Team().GetTeamsForUser(c, m1.UserId, "", true)
require.NoError(t, err)
require.Empty(t, ms)
}
@@ -2974,6 +2976,8 @@ func testSaveTeamMemberMaxMembers(t *testing.T, ss store.Store) {
}
func testGetTeamMember(t *testing.T, ss store.Store) {
c := request.TestContext(t)
teamId1 := model.NewId()
m1 := &model.TeamMember{TeamId: teamId1, UserId: model.NewId()}
@@ -2981,17 +2985,17 @@ func testGetTeamMember(t *testing.T, ss store.Store) {
require.NoError(t, nErr)
var rm1 *model.TeamMember
rm1, err := ss.Team().GetMember(context.Background(), m1.TeamId, m1.UserId)
rm1, err := ss.Team().GetMember(c, m1.TeamId, m1.UserId)
require.NoError(t, err)
require.Equal(t, rm1.TeamId, m1.TeamId, "bad team id")
require.Equal(t, rm1.UserId, m1.UserId, "bad user id")
_, err = ss.Team().GetMember(context.Background(), m1.TeamId, "")
_, err = ss.Team().GetMember(c, m1.TeamId, "")
require.Error(t, err, "empty user id - should have failed")
_, err = ss.Team().GetMember(context.Background(), "", m1.UserId)
_, err = ss.Team().GetMember(c, "", m1.UserId)
require.Error(t, err, "empty team id - should have failed")
// Test with a custom team scheme.
@@ -3021,7 +3025,7 @@ func testGetTeamMember(t *testing.T, ss store.Store) {
_, nErr = ss.Team().SaveMember(m2, -1)
require.NoError(t, nErr)
m3, err := ss.Team().GetMember(context.Background(), m2.TeamId, m2.UserId)
m3, err := ss.Team().GetMember(c, m2.TeamId, m2.UserId)
require.NoError(t, err)
t.Log(m3)
@@ -3031,7 +3035,7 @@ func testGetTeamMember(t *testing.T, ss store.Store) {
_, nErr = ss.Team().SaveMember(m4, -1)
require.NoError(t, nErr)
m5, err := ss.Team().GetMember(context.Background(), m4.TeamId, m4.UserId)
m5, err := ss.Team().GetMember(c, m4.TeamId, m4.UserId)
require.NoError(t, err)
assert.Equal(t, s2.DefaultTeamGuestRole, m5.Roles)
@@ -3292,6 +3296,8 @@ func testGetTeamsByScheme(t *testing.T, ss store.Store) {
}
func testTeamStoreMigrateTeamMembers(t *testing.T, ss store.Store) {
c := request.TestContext(t)
s1 := model.NewId()
t1 := &model.Team{
DisplayName: "Name",
@@ -3341,19 +3347,19 @@ func testTeamStoreMigrateTeamMembers(t *testing.T, ss store.Store) {
}
}
tm1b, err := ss.Team().GetMember(context.Background(), tm1.TeamId, tm1.UserId)
tm1b, err := ss.Team().GetMember(c, tm1.TeamId, tm1.UserId)
assert.NoError(t, err)
assert.Equal(t, "", tm1b.ExplicitRoles)
assert.True(t, tm1b.SchemeUser)
assert.True(t, tm1b.SchemeAdmin)
tm2b, err := ss.Team().GetMember(context.Background(), tm2.TeamId, tm2.UserId)
tm2b, err := ss.Team().GetMember(c, tm2.TeamId, tm2.UserId)
assert.NoError(t, err)
assert.Equal(t, "", tm2b.ExplicitRoles)
assert.True(t, tm2b.SchemeUser)
assert.False(t, tm2b.SchemeAdmin)
tm3b, err := ss.Team().GetMember(context.Background(), tm3.TeamId, tm3.UserId)
tm3b, err := ss.Team().GetMember(c, tm3.TeamId, tm3.UserId)
assert.NoError(t, err)
assert.Equal(t, "something_else", tm3b.ExplicitRoles)
assert.False(t, tm3b.SchemeUser)
@@ -3408,6 +3414,8 @@ func testResetAllTeamSchemes(t *testing.T, ss store.Store) {
}
func testTeamStoreClearAllCustomRoleAssignments(t *testing.T, ss store.Store) {
c := request.TestContext(t)
m1 := &model.TeamMember{
TeamId: model.NewId(),
UserId: model.NewId(),
@@ -3434,19 +3442,19 @@ func testTeamStoreClearAllCustomRoleAssignments(t *testing.T, ss store.Store) {
require.NoError(t, (ss.Team().ClearAllCustomRoleAssignments()))
r1, err := ss.Team().GetMember(context.Background(), m1.TeamId, m1.UserId)
r1, err := ss.Team().GetMember(c, m1.TeamId, m1.UserId)
require.NoError(t, err)
assert.Equal(t, m1.ExplicitRoles, r1.Roles)
r2, err := ss.Team().GetMember(context.Background(), m2.TeamId, m2.UserId)
r2, err := ss.Team().GetMember(c, m2.TeamId, m2.UserId)
require.NoError(t, err)
assert.Equal(t, "team_user team_admin", r2.Roles)
r3, err := ss.Team().GetMember(context.Background(), m3.TeamId, m3.UserId)
r3, err := ss.Team().GetMember(c, m3.TeamId, m3.UserId)
require.NoError(t, err)
assert.Equal(t, m3.ExplicitRoles, r3.Roles)
r4, err := ss.Team().GetMember(context.Background(), m4.TeamId, m4.UserId)
r4, err := ss.Team().GetMember(c, m4.TeamId, m4.UserId)
require.NoError(t, err)
assert.Equal(t, "", r4.Roles)
}

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

@@ -4,13 +4,13 @@
package storetest
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -22,6 +22,8 @@ func TestUploadSessionStore(t *testing.T, ss store.Store) {
}
func testUploadSessionStoreSaveGet(t *testing.T, ss store.Store) {
c := request.TestContext(t)
var session *model.UploadSession
t.Run("saving nil session should fail", func(t *testing.T) {
@@ -53,13 +55,13 @@ func testUploadSessionStoreSaveGet(t *testing.T, ss store.Store) {
})
t.Run("getting non-existing session should fail", func(t *testing.T) {
us, err := ss.UploadSession().Get(context.Background(), "fake")
us, err := ss.UploadSession().Get(c, "fake")
require.Error(t, err)
require.Nil(t, us)
})
t.Run("getting existing session should succeed", func(t *testing.T) {
us, err := ss.UploadSession().Get(context.Background(), session.Id)
us, err := ss.UploadSession().Get(c, session.Id)
require.NoError(t, err)
require.NotNil(t, us)
require.Equal(t, session, us)
@@ -67,6 +69,8 @@ func testUploadSessionStoreSaveGet(t *testing.T, ss store.Store) {
}
func testUploadSessionStoreUpdate(t *testing.T, ss store.Store) {
c := request.TestContext(t)
session := &model.UploadSession{
Type: model.UploadTypeAttachment,
UserId: model.NewId(),
@@ -101,7 +105,7 @@ func testUploadSessionStoreUpdate(t *testing.T, ss store.Store) {
err = ss.UploadSession().Update(us)
require.NoError(t, err)
updated, err := ss.UploadSession().Get(context.Background(), us.Id)
updated, err := ss.UploadSession().Get(c, us.Id)
require.NoError(t, err)
require.NotNil(t, us)
require.Equal(t, us, updated)
@@ -176,6 +180,8 @@ func testUploadSessionStoreGetForUser(t *testing.T, ss store.Store) {
}
func testUploadSessionStoreDelete(t *testing.T, ss store.Store) {
c := request.TestContext(t)
session := &model.UploadSession{
Id: model.NewId(),
Type: model.UploadTypeAttachment,
@@ -200,7 +206,7 @@ func testUploadSessionStoreDelete(t *testing.T, ss store.Store) {
err = ss.UploadSession().Delete(session.Id)
require.NoError(t, err)
us, err = ss.UploadSession().Get(context.Background(), us.Id)
us, err = ss.UploadSession().Get(c, us.Id)
require.Error(t, err)
require.Nil(t, us)
require.IsType(t, &store.ErrNotFound{}, err)

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

@@ -4,12 +4,12 @@
package storetest
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -20,6 +20,8 @@ func TestUserAccessTokenStore(t *testing.T, ss store.Store) {
}
func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) {
c := request.TestContext(t)
uat := &model.UserAccessToken{
Token: model.NewId(),
UserId: model.NewId(),
@@ -30,7 +32,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) {
s1.UserId = uat.UserId
s1.Token = uat.Token
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
_, nErr := ss.UserAccessToken().Save(uat)
@@ -58,7 +60,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) {
nErr = ss.UserAccessToken().Delete(uat.Id)
require.NoError(t, nErr)
_, err = ss.Session().Get(context.Background(), s1.Token)
_, err = ss.Session().Get(c, s1.Token)
require.Error(t, err, "should error - session should be deleted")
_, nErr = ss.UserAccessToken().GetByToken(s1.Token)
@@ -68,7 +70,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) {
s2.UserId = uat.UserId
s2.Token = uat.Token
s2, err = ss.Session().Save(s2)
s2, err = ss.Session().Save(c, s2)
require.NoError(t, err)
_, nErr = ss.UserAccessToken().Save(uat)
@@ -77,7 +79,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) {
nErr = ss.UserAccessToken().DeleteAllForUser(uat.UserId)
require.NoError(t, nErr)
_, err = ss.Session().Get(context.Background(), s2.Token)
_, err = ss.Session().Get(c, s2.Token)
require.Error(t, err, "should error - session should be deleted")
_, nErr = ss.UserAccessToken().GetByToken(s2.Token)
@@ -85,6 +87,8 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) {
}
func testUserAccessTokenDisableEnable(t *testing.T, ss store.Store) {
c := request.TestContext(t)
uat := &model.UserAccessToken{
Token: model.NewId(),
UserId: model.NewId(),
@@ -95,7 +99,7 @@ func testUserAccessTokenDisableEnable(t *testing.T, ss store.Store) {
s1.UserId = uat.UserId
s1.Token = uat.Token
s1, err := ss.Session().Save(s1)
s1, err := ss.Session().Save(c, s1)
require.NoError(t, err)
_, nErr := ss.UserAccessToken().Save(uat)
@@ -104,14 +108,14 @@ func testUserAccessTokenDisableEnable(t *testing.T, ss store.Store) {
nErr = ss.UserAccessToken().UpdateTokenDisable(uat.Id)
require.NoError(t, nErr)
_, err = ss.Session().Get(context.Background(), s1.Token)
_, err = ss.Session().Get(c, s1.Token)
require.Error(t, err, "should error - session should be deleted")
s2 := &model.Session{}
s2.UserId = uat.UserId
s2.Token = uat.Token
_, err = ss.Session().Save(s2)
_, err = ss.Session().Save(c, s2)
require.NoError(t, err)
nErr = ss.UserAccessToken().UpdateTokenEnable(uat.Id)
@@ -119,6 +123,8 @@ func testUserAccessTokenDisableEnable(t *testing.T, ss store.Store) {
}
func testUserAccessTokenSearch(t *testing.T, ss store.Store) {
c := request.TestContext(t)
u1 := model.User{}
u1.Email = MakeEmail()
u1.Username = model.NewId()
@@ -136,7 +142,7 @@ func testUserAccessTokenSearch(t *testing.T, ss store.Store) {
s1.UserId = uat.UserId
s1.Token = uat.Token
_, nErr := ss.Session().Save(s1)
_, nErr := ss.Session().Save(c, s1)
require.NoError(t, nErr)
_, nErr = ss.UserAccessToken().Save(uat)

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

@@ -14,6 +14,7 @@ import (
"github.com/stretchr/testify/require"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/request"
"github.com/mattermost/mattermost/server/v8/channels/store"
)
@@ -5244,6 +5245,8 @@ func testUserStoreGetChannelGroupUsers(t *testing.T, ss store.Store) {
}
func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
c := request.TestContext(t)
// create users
t.Run("Must do nothing with regular user", func(t *testing.T) {
id := model.NewId()
@@ -5280,7 +5283,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
require.Equal(t, "system_user", updatedUser.Roles)
require.True(t, user.UpdateAt < updatedUser.UpdateAt)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id)
require.NoError(t, nErr)
require.False(t, updatedTeamMember.SchemeGuest)
require.True(t, updatedTeamMember.SchemeUser)
@@ -5325,7 +5328,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_user system_admin", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id)
require.NoError(t, nErr)
require.False(t, updatedTeamMember.SchemeGuest)
require.True(t, updatedTeamMember.SchemeUser)
@@ -5381,7 +5384,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_user", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id)
require.NoError(t, nErr)
require.False(t, updatedTeamMember.SchemeGuest)
require.True(t, updatedTeamMember.SchemeUser)
@@ -5421,7 +5424,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_user", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id)
require.NoError(t, nErr)
require.False(t, updatedTeamMember.SchemeGuest)
require.True(t, updatedTeamMember.SchemeUser)
@@ -5466,7 +5469,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_user custom_role", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id)
require.NoError(t, nErr)
require.False(t, updatedTeamMember.SchemeGuest)
require.True(t, updatedTeamMember.SchemeUser)
@@ -5532,7 +5535,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_user", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId1, user1.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId1, user1.Id)
require.NoError(t, nErr)
require.False(t, updatedTeamMember.SchemeGuest)
require.True(t, updatedTeamMember.SchemeUser)
@@ -5546,7 +5549,7 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_guest", notUpdatedUser.Roles)
notUpdatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId2, user2.Id)
notUpdatedTeamMember, nErr := ss.Team().GetMember(c, teamId2, user2.Id)
require.NoError(t, nErr)
require.True(t, notUpdatedTeamMember.SchemeGuest)
require.False(t, notUpdatedTeamMember.SchemeUser)
@@ -5559,6 +5562,8 @@ func testUserStorePromoteGuestToUser(t *testing.T, ss store.Store) {
}
func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
c := request.TestContext(t)
// create users
t.Run("Must do nothing with guest", func(t *testing.T) {
id := model.NewId()
@@ -5593,7 +5598,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
require.Equal(t, "system_guest", updatedUser.Roles)
require.True(t, user.UpdateAt < updatedUser.UpdateAt)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, updatedUser.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, updatedUser.Id)
require.NoError(t, nErr)
require.True(t, updatedTeamMember.SchemeGuest)
require.False(t, updatedTeamMember.SchemeUser)
@@ -5636,7 +5641,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_guest", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id)
require.NoError(t, nErr)
require.True(t, updatedTeamMember.SchemeGuest)
require.False(t, updatedTeamMember.SchemeUser)
@@ -5688,7 +5693,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_guest", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id)
require.NoError(t, nErr)
require.True(t, updatedTeamMember.SchemeGuest)
require.False(t, updatedTeamMember.SchemeUser)
@@ -5726,7 +5731,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_guest", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id)
require.NoError(t, nErr)
require.True(t, updatedTeamMember.SchemeGuest)
require.False(t, updatedTeamMember.SchemeUser)
@@ -5769,7 +5774,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_guest", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId, user.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId, user.Id)
require.NoError(t, nErr)
require.True(t, updatedTeamMember.SchemeGuest)
require.False(t, updatedTeamMember.SchemeUser)
@@ -5833,7 +5838,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_guest", updatedUser.Roles)
updatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId1, user1.Id)
updatedTeamMember, nErr := ss.Team().GetMember(c, teamId1, user1.Id)
require.NoError(t, nErr)
require.True(t, updatedTeamMember.SchemeGuest)
require.False(t, updatedTeamMember.SchemeUser)
@@ -5847,7 +5852,7 @@ func testUserStoreDemoteUserToGuest(t *testing.T, ss store.Store) {
require.NoError(t, err)
require.Equal(t, "system_user", notUpdatedUser.Roles)
notUpdatedTeamMember, nErr := ss.Team().GetMember(context.Background(), teamId2, user2.Id)
notUpdatedTeamMember, nErr := ss.Team().GetMember(c, teamId2, user2.Id)
require.NoError(t, nErr)
require.False(t, notUpdatedTeamMember.SchemeGuest)
require.True(t, notUpdatedTeamMember.SchemeUser)

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

@@ -779,10 +779,10 @@ func (s *TimerLayerChannelStore) CreateDirectChannel(userID *model.User, otherUs
return result, err
}
func (s *TimerLayerChannelStore) CreateInitialSidebarCategories(userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) {
func (s *TimerLayerChannelStore) CreateInitialSidebarCategories(c request.CTX, userID string, opts *store.SidebarCategorySearchOpts) (*model.OrderedSidebarCategories, error) {
start := time.Now()
result, err := s.ChannelStore.CreateInitialSidebarCategories(userID, opts)
result, err := s.ChannelStore.CreateInitialSidebarCategories(c, userID, opts)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -2947,10 +2947,10 @@ func (s *TimerLayerComplianceStore) GetAll(offset int, limit int) (model.Complia
return result, err
}
func (s *TimerLayerComplianceStore) MessageExport(ctx context.Context, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
func (s *TimerLayerComplianceStore) MessageExport(c request.CTX, cursor model.MessageExportCursor, limit int) ([]*model.MessageExport, model.MessageExportCursor, error) {
start := time.Now()
result, resultVar1, err := s.ComplianceStore.MessageExport(ctx, cursor, limit)
result, resultVar1, err := s.ComplianceStore.MessageExport(c, cursor, limit)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -3155,10 +3155,10 @@ func (s *TimerLayerEmojiStore) Delete(emoji *model.Emoji, timestamp int64) error
return err
}
func (s *TimerLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bool) (*model.Emoji, error) {
func (s *TimerLayerEmojiStore) Get(c request.CTX, id string, allowFromCache bool) (*model.Emoji, error) {
start := time.Now()
result, err := s.EmojiStore.Get(ctx, id, allowFromCache)
result, err := s.EmojiStore.Get(c, id, allowFromCache)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -3171,10 +3171,10 @@ func (s *TimerLayerEmojiStore) Get(ctx request.CTX, id string, allowFromCache bo
return result, err
}
func (s *TimerLayerEmojiStore) GetByName(ctx request.CTX, name string, allowFromCache bool) (*model.Emoji, error) {
func (s *TimerLayerEmojiStore) GetByName(c request.CTX, name string, allowFromCache bool) (*model.Emoji, error) {
start := time.Now()
result, err := s.EmojiStore.GetByName(ctx, name, allowFromCache)
result, err := s.EmojiStore.GetByName(c, name, allowFromCache)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -3203,10 +3203,10 @@ func (s *TimerLayerEmojiStore) GetList(offset int, limit int, sort string) ([]*m
return result, err
}
func (s *TimerLayerEmojiStore) GetMultipleByName(ctx request.CTX, names []string) ([]*model.Emoji, error) {
func (s *TimerLayerEmojiStore) GetMultipleByName(c request.CTX, names []string) ([]*model.Emoji, error) {
start := time.Now()
result, err := s.EmojiStore.GetMultipleByName(ctx, names)
result, err := s.EmojiStore.GetMultipleByName(c, names)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -4721,10 +4721,10 @@ func (s *TimerLayerJobStore) UpdateStatusOptimistically(id string, currentStatus
return result, err
}
func (s *TimerLayerLicenseStore) Get(ctx context.Context, id string) (*model.LicenseRecord, error) {
func (s *TimerLayerLicenseStore) Get(c request.CTX, id string) (*model.LicenseRecord, error) {
start := time.Now()
result, err := s.LicenseStore.Get(ctx, id)
result, err := s.LicenseStore.Get(c, id)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -7471,10 +7471,10 @@ func (s *TimerLayerSessionStore) Cleanup(expiryTime int64, batchSize int64) erro
return err
}
func (s *TimerLayerSessionStore) Get(ctx context.Context, sessionIDOrToken string) (*model.Session, error) {
func (s *TimerLayerSessionStore) Get(c request.CTX, sessionIDOrToken string) (*model.Session, error) {
start := time.Now()
result, err := s.SessionStore.Get(ctx, sessionIDOrToken)
result, err := s.SessionStore.Get(c, sessionIDOrToken)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -7487,10 +7487,10 @@ func (s *TimerLayerSessionStore) Get(ctx context.Context, sessionIDOrToken strin
return result, err
}
func (s *TimerLayerSessionStore) GetSessions(userID string) ([]*model.Session, error) {
func (s *TimerLayerSessionStore) GetSessions(c *request.Context, userID string) ([]*model.Session, error) {
start := time.Now()
result, err := s.SessionStore.GetSessions(userID)
result, err := s.SessionStore.GetSessions(c, userID)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -7583,10 +7583,10 @@ func (s *TimerLayerSessionStore) RemoveAllSessions() error {
return err
}
func (s *TimerLayerSessionStore) Save(session *model.Session) (*model.Session, error) {
func (s *TimerLayerSessionStore) Save(c request.CTX, session *model.Session) (*model.Session, error) {
start := time.Now()
result, err := s.SessionStore.Save(session)
result, err := s.SessionStore.Save(c, session)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -8670,10 +8670,10 @@ func (s *TimerLayerTeamStore) GetMany(ids []string) ([]*model.Team, error) {
return result, err
}
func (s *TimerLayerTeamStore) GetMember(ctx context.Context, teamID string, userID string) (*model.TeamMember, error) {
func (s *TimerLayerTeamStore) GetMember(c request.CTX, teamID string, userID string) (*model.TeamMember, error) {
start := time.Now()
result, err := s.TeamStore.GetMember(ctx, teamID, userID)
result, err := s.TeamStore.GetMember(c, teamID, userID)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -8766,10 +8766,10 @@ func (s *TimerLayerTeamStore) GetTeamsByUserId(userID string) ([]*model.Team, er
return result, err
}
func (s *TimerLayerTeamStore) GetTeamsForUser(ctx context.Context, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
func (s *TimerLayerTeamStore) GetTeamsForUser(c request.CTX, userID string, excludeTeamID string, includeDeleted bool) ([]*model.TeamMember, error) {
start := time.Now()
result, err := s.TeamStore.GetTeamsForUser(ctx, userID, excludeTeamID, includeDeleted)
result, err := s.TeamStore.GetTeamsForUser(c, userID, excludeTeamID, includeDeleted)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {
@@ -9756,10 +9756,10 @@ func (s *TimerLayerUploadSessionStore) Delete(id string) error {
return err
}
func (s *TimerLayerUploadSessionStore) Get(ctx context.Context, id string) (*model.UploadSession, error) {
func (s *TimerLayerUploadSessionStore) Get(c request.CTX, id string) (*model.UploadSession, error) {
start := time.Now()
result, err := s.UploadSessionStore.Get(ctx, id)
result, err := s.UploadSessionStore.Get(c, id)
elapsed := float64(time.Since(start)) / float64(time.Second)
if s.Root.Metrics != nil {