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 удалений

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

@@ -4,13 +4,12 @@
package platform
import (
"context"
"fmt"
"time"
"github.com/mattermost/mattermost/server/public/model"
"github.com/mattermost/mattermost/server/public/shared/mlog"
"github.com/mattermost/mattermost/server/v8/channels/store/sqlstore"
"github.com/mattermost/mattermost/server/public/shared/request"
)
func (ps *PlatformService) ReturnSessionToPool(session *model.Session) {
@@ -20,10 +19,10 @@ func (ps *PlatformService) ReturnSessionToPool(session *model.Session) {
}
}
func (ps *PlatformService) CreateSession(session *model.Session) (*model.Session, error) {
func (ps *PlatformService) CreateSession(c *request.Context, session *model.Session) (*model.Session, error) {
session.Token = ""
session, err := ps.Store.Session().Save(session)
session, err := ps.Store.Session().Save(c, session)
if err != nil {
return nil, err
}
@@ -33,12 +32,12 @@ func (ps *PlatformService) CreateSession(session *model.Session) (*model.Session
return session, nil
}
func (ps *PlatformService) GetSessionContext(ctx context.Context, token string) (*model.Session, error) {
return ps.Store.Session().Get(ctx, token)
func (ps *PlatformService) GetSessionContext(c *request.Context, token string) (*model.Session, error) {
return ps.Store.Session().Get(c, token)
}
func (ps *PlatformService) GetSessions(userID string) ([]*model.Session, error) {
return ps.Store.Session().GetSessions(userID)
func (ps *PlatformService) GetSessions(c *request.Context, userID string) ([]*model.Session, error) {
return ps.Store.Session().GetSessions(c, userID)
}
func (ps *PlatformService) AddSessionToCache(session *model.Session) {
@@ -97,7 +96,7 @@ func (ps *PlatformService) ClearAllUsersSessionCache() {
}
}
func (ps *PlatformService) GetSession(token string) (*model.Session, error) {
func (ps *PlatformService) GetSession(c *request.Context, token string) (*model.Session, error) {
var session = ps.sessionPool.Get().(*model.Session)
if err := ps.sessionCache.Get(token, session); err == nil {
if m := ps.metricsIFace; m != nil {
@@ -113,11 +112,11 @@ func (ps *PlatformService) GetSession(token string) (*model.Session, error) {
return session, nil
}
return ps.GetSessionContext(sqlstore.WithMaster(context.Background()), token)
return ps.GetSessionContext(c, token)
}
func (ps *PlatformService) GetSessionByID(sessionID string) (*model.Session, error) {
return ps.Store.Session().Get(context.Background(), sessionID)
func (ps *PlatformService) GetSessionByID(c *request.Context, sessionID string) (*model.Session, error) {
return ps.Store.Session().Get(c, sessionID)
}
func (ps *PlatformService) RevokeSessionsFromAllUsers() error {
@@ -135,16 +134,16 @@ func (ps *PlatformService) RevokeSessionsFromAllUsers() error {
return nil
}
func (ps *PlatformService) RevokeSessionsForDeviceId(userID string, deviceID string, currentSessionId string) error {
sessions, err := ps.Store.Session().GetSessions(userID)
func (ps *PlatformService) RevokeSessionsForDeviceId(c *request.Context, userID string, deviceID string, currentSessionId string) error {
sessions, err := ps.Store.Session().GetSessions(c, userID)
if err != nil {
return err
}
for _, session := range sessions {
if session.DeviceId == deviceID && session.Id != currentSessionId {
mlog.Debug("Revoking sessionId for userId. Re-login with the same device Id", mlog.String("session_id", session.Id), mlog.String("user_id", userID))
if err := ps.RevokeSession(session); err != nil {
mlog.Warn("Could not revoke session for device", mlog.String("device_id", deviceID), mlog.Err(err))
c.Logger().Debug("Revoking sessionId for userId. Re-login with the same device Id", mlog.String("session_id", session.Id), mlog.String("user_id", userID))
if err := ps.RevokeSession(c, session); err != nil {
c.Logger().Warn("Could not revoke session for device", mlog.String("device_id", deviceID), mlog.Err(err))
}
}
}
@@ -152,9 +151,9 @@ func (ps *PlatformService) RevokeSessionsForDeviceId(userID string, deviceID str
return nil
}
func (ps *PlatformService) RevokeSession(session *model.Session) error {
func (ps *PlatformService) RevokeSession(c *request.Context, session *model.Session) error {
if session.IsOAuth {
if err := ps.RevokeAccessToken(session.Token); err != nil {
if err := ps.RevokeAccessToken(c, session.Token); err != nil {
return err
}
} else {
@@ -168,8 +167,8 @@ func (ps *PlatformService) RevokeSession(session *model.Session) error {
return nil
}
func (ps *PlatformService) RevokeAccessToken(token string) error {
session, _ := ps.GetSession(token)
func (ps *PlatformService) RevokeAccessToken(c *request.Context, token string) error {
session, _ := ps.GetSession(c, token)
defer ps.ReturnSessionToPool(session)
@@ -223,8 +222,8 @@ func (ps *PlatformService) ExtendSessionExpiry(session *model.Session, newExpiry
return nil
}
func (ps *PlatformService) UpdateSessionsIsGuest(userID string, isGuest bool) error {
sessions, err := ps.GetSessions(userID)
func (ps *PlatformService) UpdateSessionsIsGuest(c *request.Context, userID string, isGuest bool) error {
sessions, err := ps.GetSessions(c, userID)
if err != nil {
return err
}
@@ -241,14 +240,14 @@ func (ps *PlatformService) UpdateSessionsIsGuest(userID string, isGuest bool) er
return nil
}
func (ps *PlatformService) RevokeAllSessions(userID string) error {
sessions, err := ps.Store.Session().GetSessions(userID)
func (ps *PlatformService) RevokeAllSessions(c *request.Context, userID string) error {
sessions, err := ps.Store.Session().GetSessions(c, userID)
if err != nil {
return fmt.Errorf("%s: %w", err.Error(), GetSessionError)
}
for _, session := range sessions {
if session.IsOAuth {
ps.RevokeAccessToken(session.Token)
ps.RevokeAccessToken(c, session.Token)
} else {
if err := ps.Store.Session().Remove(session.Id); err != nil {
return fmt.Errorf("%s: %w", err.Error(), DeleteSessionError)