Migrate store methods to use request.Context instead of context.Context (#24836)
Этот коммит содержится в:
коммит произвёл
GitHub
родитель
0d5a8b8841
Коммит
13c05a571f
@@ -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)
|
||||
}
|
||||
|
||||
Ссылка в новой задаче
Block a user