* users: add cache to service

* reflect review comments
Этот коммит содержится в:
Ibrahim Serdar Acikgoz
2021-06-14 18:08:00 +03:00
коммит произвёл GitHub
родитель 0ae307808a
Коммит 24fb0033f4
26 изменённых файлов: 364 добавлений и 271 удалений

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

@@ -168,7 +168,7 @@ func (s *Server) InvalidateAllCaches() *model.AppError {
func (s *Server) InvalidateAllCachesSkipSend() { func (s *Server) InvalidateAllCachesSkipSend() {
mlog.Info("Purging all caches") mlog.Info("Purging all caches")
s.sessionCache.Purge() s.userService.ClearAllUsersSessionCacheLocal()
s.statusCache.Purge() s.statusCache.Purge()
s.Store.Team().ClearCaches() s.Store.Team().ClearCaches()
s.Store.Channel().ClearCaches() s.Store.Channel().ClearCaches()

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

@@ -926,6 +926,7 @@ type AppIface interface {
RestoreTeam(teamID string) *model.AppError RestoreTeam(teamID string) *model.AppError
RestrictUsersGetByPermissions(userID string, options *model.UserGetOptions) (*model.UserGetOptions, *model.AppError) RestrictUsersGetByPermissions(userID string, options *model.UserGetOptions) (*model.UserGetOptions, *model.AppError)
RestrictUsersSearchByPermissions(userID string, options *model.UserSearchOptions) (*model.UserSearchOptions, *model.AppError) RestrictUsersSearchByPermissions(userID string, options *model.UserSearchOptions) (*model.UserSearchOptions, *model.AppError)
ReturnSessionToPool(session *model.Session)
RevokeAccessToken(token string) *model.AppError RevokeAccessToken(token string) *model.AppError
RevokeAllSessions(userID string) *model.AppError RevokeAllSessions(userID string) *model.AppError
RevokeSession(session *model.Session) *model.AppError RevokeSession(session *model.Session) *model.AppError
@@ -976,7 +977,6 @@ type AppIface interface {
SendPasswordReset(email string, siteURL string) (bool, *model.AppError) SendPasswordReset(email string, siteURL string) (bool, *model.AppError)
SendPaymentFailedEmail(failedPayment *model.FailedPayment) *model.AppError SendPaymentFailedEmail(failedPayment *model.FailedPayment) *model.AppError
ServeInterPluginRequest(w http.ResponseWriter, r *http.Request, sourcePluginId, destinationPluginId string) ServeInterPluginRequest(w http.ResponseWriter, r *http.Request, sourcePluginId, destinationPluginId string)
SessionCacheLength() int
SessionHasPermissionTo(session model.Session, permission *model.Permission) bool SessionHasPermissionTo(session model.Session, permission *model.Permission) bool
SessionHasPermissionToAny(session model.Session, permissions []*model.Permission) bool SessionHasPermissionToAny(session model.Session, permissions []*model.Permission) bool
SessionHasPermissionToCategory(session model.Session, userID, teamID, categoryId string) bool SessionHasPermissionToCategory(session model.Session, userID, teamID, categoryId string) bool

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

@@ -1968,14 +1968,21 @@ func TestMarkChannelsAsViewedPanic(t *testing.T) {
"userID": 1, "userID": 1,
} }
mockChannelStore.On("UpdateLastViewedAt", []string{"channelID"}, "userID", false).Return(times, nil) mockChannelStore.On("UpdateLastViewedAt", []string{"channelID"}, "userID", false).Return(times, nil)
th.App.srv.userService = users.New(&mockUserStore, th.App.srv.Config) mockSessionStore := mocks.SessionStore{}
var err error
th.App.srv.userService, err = users.New(users.ServiceInitializer{
UserStore: &mockUserStore,
SessionStore: &mockSessionStore,
ConfigFn: th.App.srv.Config,
})
require.NoError(t, err)
mockPreferenceStore := mocks.PreferenceStore{} mockPreferenceStore := mocks.PreferenceStore{}
mockPreferenceStore.On("Get", mock.AnythingOfType("string"), mock.AnythingOfType("string"), mock.AnythingOfType("string")).Return(&model.Preference{Value: "test"}, nil) mockPreferenceStore.On("Get", mock.AnythingOfType("string"), mock.AnythingOfType("string"), mock.AnythingOfType("string")).Return(&model.Preference{Value: "test"}, nil)
mockStore.On("Channel").Return(&mockChannelStore) mockStore.On("Channel").Return(&mockChannelStore)
mockStore.On("Preference").Return(&mockPreferenceStore) mockStore.On("Preference").Return(&mockPreferenceStore)
_, err := th.App.MarkChannelsAsViewed([]string{"channelID"}, "userID", th.Context.Session().Id, false) _, appErr := th.App.MarkChannelsAsViewed([]string{"channelID"}, "userID", th.Context.Session().Id, false)
require.Nil(t, err) require.Nil(t, appErr)
} }
func TestClearChannelMembersCache(t *testing.T) { func TestClearChannelMembersCache(t *testing.T) {

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

@@ -102,26 +102,13 @@ func (s *Server) clusterInvalidateCacheForUserTeamsHandler(msg *model.ClusterMes
} }
func (s *Server) clearSessionCacheForUserSkipClusterSend(userID string) { func (s *Server) clearSessionCacheForUserSkipClusterSend(userID string) {
if keys, err := s.sessionCache.Keys(); err == nil { s.userService.ClearUserSessionCacheLocal(userID)
var session *model.Session
for _, key := range keys {
if err := s.sessionCache.Get(key, &session); err == nil {
if session.UserId == userID {
s.sessionCache.Remove(key)
if s.Metrics != nil {
s.Metrics.IncrementMemCacheInvalidationCounterSession()
}
}
}
}
}
s.invalidateWebConnSessionCacheForUser(userID) s.invalidateWebConnSessionCacheForUser(userID)
} }
func (s *Server) clearSessionCacheForAllUsersSkipClusterSend() { func (s *Server) clearSessionCacheForAllUsersSkipClusterSend() {
mlog.Info("Purging sessions cache") mlog.Info("Purging sessions cache")
s.sessionCache.Purge() s.userService.ClearAllUsersSessionCacheLocal()
} }
func (s *Server) clusterClearSessionCacheForUserHandler(msg *model.ClusterMessage) { func (s *Server) clusterClearSessionCacheForUserHandler(msg *model.ClusterMessage) {

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

@@ -350,7 +350,7 @@ func (s *Server) ClientConfigWithComputed() map[string]string {
// These properties are not configurable, but nevertheless represent configuration expected // These properties are not configurable, but nevertheless represent configuration expected
// by the client. // by the client.
respCfg["NoAccounts"] = strconv.FormatBool(s.IsFirstUserAccount()) respCfg["NoAccounts"] = strconv.FormatBool(s.userService.IsFirstUserAccount())
respCfg["MaxPostSize"] = strconv.Itoa(s.MaxPostSize()) respCfg["MaxPostSize"] = strconv.Itoa(s.MaxPostSize())
respCfg["UpgradedFromTE"] = strconv.FormatBool(s.isUpgradedFromTE()) respCfg["UpgradedFromTE"] = strconv.FormatBool(s.isUpgradedFromTE())
respCfg["InstallationDate"] = "" respCfg["InstallationDate"] = ""

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

@@ -379,7 +379,7 @@ func (a *App) newSession(appName string, user *model.User) (*model.Session, *mod
return nil, model.NewAppError("newSession", "api.oauth.get_access_token.internal_session.app_error", nil, "", http.StatusInternalServerError) return nil, model.NewAppError("newSession", "api.oauth.get_access_token.internal_session.app_error", nil, "", http.StatusInternalServerError)
} }
a.AddSessionToCache(session) a.srv.userService.AddSessionToCache(session)
return session, nil return session, nil
} }
@@ -520,7 +520,7 @@ func (a *App) RegenerateOAuthAppSecret(app *model.OAuthApp) (*model.OAuthApp, *m
func (a *App) RevokeAccessToken(token string) *model.AppError { func (a *App) RevokeAccessToken(token string) *model.AppError {
session, _ := a.GetSession(token) session, _ := a.GetSession(token)
defer ReturnSessionToPool(session) defer a.srv.userService.ReturnSessionToPool(session)
schan := make(chan error, 1) schan := make(chan error, 1)
go func() { go func() {

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

@@ -13202,6 +13202,21 @@ func (a *OpenTracingAppLayer) RestrictUsersSearchByPermissions(userID string, op
return resultVar0, resultVar1 return resultVar0, resultVar1
} }
func (a *OpenTracingAppLayer) ReturnSessionToPool(session *model.Session) {
origCtx := a.ctx
span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.ReturnSessionToPool")
a.ctx = newCtx
a.app.Srv().Store.SetContext(newCtx)
defer func() {
a.app.Srv().Store.SetContext(origCtx)
a.ctx = origCtx
}()
defer span.Finish()
a.app.ReturnSessionToPool(session)
}
func (a *OpenTracingAppLayer) RevokeAccessToken(token string) *model.AppError { func (a *OpenTracingAppLayer) RevokeAccessToken(token string) *model.AppError {
origCtx := a.ctx origCtx := a.ctx
span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RevokeAccessToken") span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.RevokeAccessToken")
@@ -14361,23 +14376,6 @@ func (a *OpenTracingAppLayer) ServeInterPluginRequest(w http.ResponseWriter, r *
a.app.ServeInterPluginRequest(w, r, sourcePluginId, destinationPluginId) a.app.ServeInterPluginRequest(w, r, sourcePluginId, destinationPluginId)
} }
func (a *OpenTracingAppLayer) SessionCacheLength() int {
origCtx := a.ctx
span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.SessionCacheLength")
a.ctx = newCtx
a.app.Srv().Store.SetContext(newCtx)
defer func() {
a.app.Srv().Store.SetContext(origCtx)
a.ctx = origCtx
}()
defer span.Finish()
resultVar0 := a.app.SessionCacheLength()
return resultVar0
}
func (a *OpenTracingAppLayer) SessionHasPermissionTo(session model.Session, permission *model.Permission) bool { func (a *OpenTracingAppLayer) SessionHasPermissionTo(session model.Session, permission *model.Permission) bool {
origCtx := a.ctx origCtx := a.ctx
span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.SessionHasPermissionTo") span, newCtx := tracing.StartSpanWithParentByContext(a.ctx, "app.SessionHasPermissionTo")

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

@@ -143,7 +143,7 @@ func (s *Server) servePluginRequest(w http.ResponseWriter, r *http.Request, hand
r.Header.Del("Mattermost-User-Id") r.Header.Del("Mattermost-User-Id")
if token != "" { if token != "" {
session, err := New(ServerConnector(s)).GetSession(token) session, err := New(ServerConnector(s)).GetSession(token)
defer ReturnSessionToPool(session) defer s.userService.ReturnSessionToPool(session)
csrfCheckPassed := false csrfCheckPassed := false

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

@@ -135,7 +135,6 @@ type Server struct {
newStore func() (store.Store, error) newStore func() (store.Store, error)
htmlTemplateWatcher *templates.Container htmlTemplateWatcher *templates.Container
sessionCache cache.Cache
seenPendingPostIdsCache cache.Cache seenPendingPostIdsCache cache.Cache
statusCache cache.Cache statusCache cache.Cache
configListenerId string configListenerId string
@@ -328,13 +327,6 @@ func NewServer(options ...Option) (*Server, error) {
} }
var err error var err error
if s.sessionCache, err = s.CacheProvider.NewCache(&cache.CacheOptions{
Size: model.SESSION_CACHE_SIZE,
Striped: true,
StripedBuckets: maxInt(runtime.NumCPU()-1, 1),
}); err != nil {
return nil, errors.Wrap(err, "Unable to create session cache")
}
if s.seenPendingPostIdsCache, err = s.CacheProvider.NewCache(&cache.CacheOptions{ if s.seenPendingPostIdsCache, err = s.CacheProvider.NewCache(&cache.CacheOptions{
Size: PendingPostIDsCacheSize, Size: PendingPostIDsCacheSize,
}); err != nil { }); err != nil {
@@ -413,7 +405,16 @@ func NewServer(options ...Option) (*Server, error) {
return nil, errors.Wrap(err, "cannot create store") return nil, errors.Wrap(err, "cannot create store")
} }
s.userService = users.New(s.Store.User(), s.Config) s.userService, err = users.New(users.ServiceInitializer{
UserStore: s.Store.User(),
SessionStore: s.Store.Session(),
ConfigFn: s.Config,
Metrics: s.Metrics,
Cluster: s.Cluster,
})
if err != nil {
return nil, errors.Wrapf(err, "unable to create users service")
}
s.configListenerId = s.AddConfigListener(func(_, _ *model.Config) { s.configListenerId = s.AddConfigListener(func(_, _ *model.Config) {
s.configOrLicenseListener() s.configOrLicenseListener()

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

@@ -10,8 +10,6 @@ import (
"math" "math"
"net/http" "net/http"
"os" "os"
"sync"
"time"
"github.com/mattermost/mattermost-server/v5/audit" "github.com/mattermost/mattermost-server/v5/audit"
"github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/model"
@@ -21,9 +19,7 @@ import (
) )
func (a *App) CreateSession(session *model.Session) (*model.Session, *model.AppError) { func (a *App) CreateSession(session *model.Session) (*model.Session, *model.AppError) {
session.Token = "" session, err := a.srv.userService.CreateSession(session)
session, err := a.Srv().Store.Session().Save(session)
if err != nil { if err != nil {
var invErr *store.ErrInvalidInput var invErr *store.ErrInvalidInput
switch { switch {
@@ -34,24 +30,9 @@ func (a *App) CreateSession(session *model.Session) (*model.Session, *model.AppE
} }
} }
a.AddSessionToCache(session)
return session, nil return session, nil
} }
func ReturnSessionToPool(session *model.Session) {
if session != nil {
session.Id = ""
userSessionPool.Put(session)
}
}
var userSessionPool = sync.Pool{
New: func() interface{} {
return &model.Session{}
},
}
func (a *App) GetCloudSession(token string) (*model.Session, *model.AppError) { func (a *App) GetCloudSession(token string) (*model.Session, *model.AppError) {
apiKey := os.Getenv("MM_CLOUD_API_KEY") apiKey := os.Getenv("MM_CLOUD_API_KEY")
if apiKey != "" && apiKey == token { if apiKey != "" && apiKey == token {
@@ -83,19 +64,9 @@ func (a *App) GetRemoteClusterSession(token string, remoteId string) (*model.Ses
} }
func (a *App) GetSession(token string) (*model.Session, *model.AppError) { func (a *App) GetSession(token string) (*model.Session, *model.AppError) {
metrics := a.Metrics() session, err := a.srv.userService.GetSession(token)
if err != nil {
var session = userSessionPool.Get().(*model.Session) return nil, model.NewAppError("GetSession", "app.session.get.app_error", nil, err.Error(), http.StatusInternalServerError)
var err *model.AppError
if err := a.Srv().sessionCache.Get(token, session); err == nil {
if metrics != nil {
metrics.IncrementMemCacheHitCounterSession()
}
} else {
if metrics != nil {
metrics.IncrementMemCacheMissCounterSession()
}
} }
if session.Id == "" { if session.Id == "" {
@@ -107,7 +78,7 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) {
} }
if !session.IsExpired() { if !session.IsExpired() {
a.AddSessionToCache(session) a.srv.userService.AddSessionToCache(session)
} }
} }
} else if nfErr := new(store.ErrNotFound); !errors.As(nErr, &nfErr) { } else if nfErr := new(store.ErrNotFound); !errors.As(nErr, &nfErr) {
@@ -115,16 +86,17 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) {
} }
} }
var appErr *model.AppError
if session == nil || session.Id == "" { if session == nil || session.Id == "" {
session, err = a.createSessionForUserAccessToken(token) session, appErr = a.createSessionForUserAccessToken(token)
if err != nil { if appErr != nil {
detailedError := "" detailedError := ""
statusCode := http.StatusUnauthorized statusCode := http.StatusUnauthorized
if err.Id != "app.user_access_token.invalid_or_missing" { if appErr.Id != "app.user_access_token.invalid_or_missing" {
detailedError = err.Error() detailedError = appErr.Error()
statusCode = err.StatusCode statusCode = appErr.StatusCode
} else { } else {
mlog.Warn("Error while creating session for user access token", mlog.Err(err)) mlog.Warn("Error while creating session for user access token", mlog.Err(appErr))
} }
return nil, model.NewAppError("GetSession", "api.context.invalid_token.error", map[string]interface{}{"Token": token, "Error": detailedError}, "", statusCode) return nil, model.NewAppError("GetSession", "api.context.invalid_token.error", map[string]interface{}{"Token": token, "Error": detailedError}, "", statusCode)
} }
@@ -162,7 +134,6 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) {
} }
func (a *App) GetSessions(userID string) ([]*model.Session, *model.AppError) { func (a *App) GetSessions(userID string) ([]*model.Session, *model.AppError) {
sessions, err := a.Srv().Store.Session().GetSessions(userID) sessions, err := a.Srv().Store.Session().GetSessions(userID)
if err != nil { if err != nil {
return nil, model.NewAppError("GetSessions", "app.session.get_sessions.app_error", nil, err.Error(), http.StatusInternalServerError) return nil, model.NewAppError("GetSessions", "app.session.get_sessions.app_error", nil, err.Error(), http.StatusInternalServerError)
@@ -185,7 +156,7 @@ func (a *App) UpdateSessionsIsGuest(userID string, isGuest bool) {
mlog.Warn("Unable to update isGuest session", mlog.Err(err)) mlog.Warn("Unable to update isGuest session", mlog.Err(err))
continue continue
} }
a.AddSessionToCache(session) a.srv.userService.AddSessionToCache(session)
} }
} }
@@ -209,6 +180,10 @@ func (a *App) RevokeAllSessions(userID string) *model.AppError {
return nil return nil
} }
func (a *App) AddSessionToCache(session *model.Session) {
a.srv.userService.AddSessionToCache(session)
}
// RevokeSessionsFromAllUsers will go through all the sessions active // RevokeSessionsFromAllUsers will go through all the sessions active
// in the server and revoke them // in the server and revoke them
func (a *App) RevokeSessionsFromAllUsers() *model.AppError { func (a *App) RevokeSessionsFromAllUsers() *model.AppError {
@@ -226,29 +201,16 @@ func (a *App) RevokeSessionsFromAllUsers() *model.AppError {
return nil return nil
} }
func (a *App) ClearSessionCacheForUser(userID string) { func (a *App) ReturnSessionToPool(session *model.Session) {
a.ClearSessionCacheForUserSkipClusterSend(userID) a.srv.userService.ReturnSessionToPool(session)
}
if a.Cluster() != nil { func (a *App) ClearSessionCacheForUser(userID string) {
msg := &model.ClusterMessage{ a.srv.userService.ClearUserSessionCache(userID)
Event: model.CLUSTER_EVENT_CLEAR_SESSION_CACHE_FOR_USER,
SendType: model.CLUSTER_SEND_RELIABLE,
Data: userID,
}
a.Cluster().SendClusterMessage(msg)
}
} }
func (a *App) ClearSessionCacheForAllUsers() { func (a *App) ClearSessionCacheForAllUsers() {
a.ClearSessionCacheForAllUsersSkipClusterSend() a.srv.userService.ClearAllUsersSessionCache()
if a.Cluster() != nil {
msg := &model.ClusterMessage{
Event: model.CLUSTER_EVENT_CLEAR_SESSION_CACHE_FOR_ALL_USERS,
SendType: model.CLUSTER_SEND_RELIABLE,
}
a.Cluster().SendClusterMessage(msg)
}
} }
func (a *App) ClearSessionCacheForUserSkipClusterSend(userID string) { func (a *App) ClearSessionCacheForUserSkipClusterSend(userID string) {
@@ -259,17 +221,6 @@ func (a *App) ClearSessionCacheForAllUsersSkipClusterSend() {
a.Srv().clearSessionCacheForAllUsersSkipClusterSend() a.Srv().clearSessionCacheForAllUsersSkipClusterSend()
} }
func (a *App) AddSessionToCache(session *model.Session) {
a.Srv().sessionCache.SetWithExpiry(session.Token, session, time.Duration(int64(*a.Config().ServiceSettings.SessionCacheInMinutes))*time.Minute)
}
func (a *App) SessionCacheLength() int {
if l, err := a.Srv().sessionCache.Len(); err == nil {
return l
}
return 0
}
func (a *App) RevokeSessionsForDeviceId(userID string, deviceID string, currentSessionId string) *model.AppError { func (a *App) RevokeSessionsForDeviceId(userID string, deviceID string, currentSessionId string) *model.AppError {
sessions, err := a.Srv().Store.Session().GetSessions(userID) sessions, err := a.Srv().Store.Session().GetSessions(userID)
if err != nil { if err != nil {
@@ -344,7 +295,7 @@ func (a *App) UpdateLastActivityAtIfNeeded(session model.Session) {
} }
session.LastActivityAt = now session.LastActivityAt = now
a.AddSessionToCache(&session) a.srv.userService.AddSessionToCache(&session)
} }
// ExtendSessionExpiryIfNeeded extends Session.ExpiresAt based on session lengths in config. // ExtendSessionExpiryIfNeeded extends Session.ExpiresAt based on session lengths in config.
@@ -392,7 +343,7 @@ func (a *App) ExtendSessionExpiryIfNeeded(session *model.Session) bool {
// ensures each node will get an extended expiry within the next 10 minutes. // ensures each node will get an extended expiry within the next 10 minutes.
// Worst case is another node may generate a redundant expiry update. // Worst case is another node may generate a redundant expiry update.
session.ExpiresAt = newExpiry session.ExpiresAt = newExpiry
a.AddSessionToCache(session) a.srv.userService.AddSessionToCache(session)
mlog.Debug("Session extended", mlog.String("user_id", session.UserId), mlog.String("session_id", session.Id), mlog.Debug("Session extended", mlog.String("user_id", session.UserId), mlog.String("session_id", session.Id),
mlog.Int64("newExpiry", newExpiry), mlog.Int64("session_length", sessionLength)) mlog.Int64("newExpiry", newExpiry), mlog.Int64("session_length", sessionLength))
@@ -531,7 +482,7 @@ func (a *App) createSessionForUserAccessToken(tokenString string) (*model.Sessio
} }
} }
a.AddSessionToCache(session) a.srv.userService.AddSessionToCache(session)
return session, nil return session, nil

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

@@ -8,7 +8,6 @@ import (
"fmt" "fmt"
"os" "os"
"testing" "testing"
"time"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -16,43 +15,6 @@ import (
"github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/model"
) )
func TestCache(t *testing.T) {
th := Setup(t)
defer th.TearDown()
session := &model.Session{
Id: model.NewId(),
Token: model.NewId(),
UserId: model.NewId(),
}
session2 := &model.Session{
Id: model.NewId(),
Token: model.NewId(),
UserId: model.NewId(),
}
th.App.Srv().sessionCache.SetWithExpiry(session.Token, session, 5*time.Minute)
th.App.Srv().sessionCache.SetWithExpiry(session2.Token, session2, 5*time.Minute)
keys, err := th.App.Srv().sessionCache.Keys()
require.NoError(t, err)
require.NotEmpty(t, keys)
th.App.ClearSessionCacheForUser(session.UserId)
rkeys, err := th.App.Srv().sessionCache.Keys()
require.NoError(t, err)
require.Lenf(t, rkeys, len(keys)-1, "should have one less: %d - %d != 1", len(keys), len(rkeys))
require.NotEmpty(t, rkeys)
th.App.ClearSessionCacheForAllUsers()
rkeys, err = th.App.Srv().sessionCache.Keys()
require.NoError(t, err)
require.Empty(t, rkeys)
}
func TestGetSessionIdleTimeoutInMinutes(t *testing.T) { func TestGetSessionIdleTimeoutInMinutes(t *testing.T) {
th := Setup(t) th := Setup(t)
defer th.TearDown() defer th.TearDown()
@@ -355,8 +317,7 @@ func TestApp_ExtendExpiryIfNeeded(t *testing.T) {
require.False(t, session.IsExpired()) require.False(t, session.IsExpired())
// check cache was updated // check cache was updated
var cachedSession *model.Session cachedSession, errGet := th.App.srv.userService.GetSession(session.Token)
errGet := th.App.Srv().sessionCache.Get(session.Token, &cachedSession)
require.NoError(t, errGet) require.NoError(t, errGet)
require.Equal(t, session.ExpiresAt, cachedSession.ExpiresAt) require.Equal(t, session.ExpiresAt, cachedSession.ExpiresAt)

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

@@ -19,6 +19,7 @@ import (
"github.com/mattermost/mattermost-server/v5/app/request" "github.com/mattermost/mattermost-server/v5/app/request"
"github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/model"
"github.com/mattermost/mattermost-server/v5/plugin" "github.com/mattermost/mattermost-server/v5/plugin"
"github.com/mattermost/mattermost-server/v5/services/users"
"github.com/mattermost/mattermost-server/v5/shared/i18n" "github.com/mattermost/mattermost-server/v5/shared/i18n"
"github.com/mattermost/mattermost-server/v5/shared/mlog" "github.com/mattermost/mattermost-server/v5/shared/mlog"
"github.com/mattermost/mattermost-server/v5/store" "github.com/mattermost/mattermost-server/v5/store"
@@ -1521,7 +1522,7 @@ func (a *App) InviteGuestsToChannelsGracefully(teamID string, guestsInvite *mode
Email: email, Email: email,
Error: nil, Error: nil,
} }
if !CheckEmailDomain(email, *a.Config().GuestAccountsSettings.RestrictCreationToDomains) { if !users.CheckEmailDomain(email, *a.Config().GuestAccountsSettings.RestrictCreationToDomains) {
invite.Error = model.NewAppError("InviteGuestsToChannelsGracefully", "api.team.invite_members.invalid_email.app_error", map[string]interface{}{"Addresses": email}, "", http.StatusBadRequest) invite.Error = model.NewAppError("InviteGuestsToChannelsGracefully", "api.team.invite_members.invalid_email.app_error", map[string]interface{}{"Addresses": email}, "", http.StatusBadRequest)
} else { } else {
goodEmails = append(goodEmails, email) goodEmails = append(goodEmails, email)
@@ -1594,7 +1595,7 @@ func (a *App) InviteGuestsToChannels(teamID string, guestsInvite *model.GuestsIn
var invalidEmailList []string var invalidEmailList []string
for _, email := range guestsInvite.Emails { for _, email := range guestsInvite.Emails {
if !CheckEmailDomain(email, *a.Config().GuestAccountsSettings.RestrictCreationToDomains) { if !users.CheckEmailDomain(email, *a.Config().GuestAccountsSettings.RestrictCreationToDomains) {
invalidEmailList = append(invalidEmailList, email) invalidEmailList = append(invalidEmailList, email)
} }
} }

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

@@ -138,7 +138,7 @@ func (a *App) CreateUserWithInviteId(c *request.Context, user *model.User, invit
return nil, model.NewAppError("CreateUserWithInviteId", "app.team.invite_id.group_constrained.error", nil, "", http.StatusForbidden) return nil, model.NewAppError("CreateUserWithInviteId", "app.team.invite_id.group_constrained.error", nil, "", http.StatusForbidden)
} }
if !CheckUserDomain(user, team.AllowedDomains) { if !users.CheckUserDomain(user, team.AllowedDomains) {
return nil, model.NewAppError("CreateUserWithInviteId", "api.team.invite_members.invalid_email.app_error", map[string]interface{}{"Addresses": team.AllowedDomains}, "", http.StatusForbidden) return nil, model.NewAppError("CreateUserWithInviteId", "api.team.invite_members.invalid_email.app_error", map[string]interface{}{"Addresses": team.AllowedDomains}, "", http.StatusForbidden)
} }
@@ -207,27 +207,8 @@ func (a *App) IsUserSignUpAllowed() *model.AppError {
return nil return nil
} }
func (s *Server) IsFirstUserAccount() bool {
cachedSessions, err := s.sessionCache.Len()
if err != nil {
return false
}
if cachedSessions == 0 {
count, err := s.Store.User().Count(model.UserCountOptions{IncludeDeleted: true})
if err != nil {
mlog.Debug("There was an error fetching if first user account", mlog.Err(err))
return false
}
if count <= 0 {
return true
}
}
return false
}
func (a *App) IsFirstUserAccount() bool { func (a *App) IsFirstUserAccount() bool {
return a.Srv().IsFirstUserAccount() return a.srv.userService.IsFirstUserAccount()
} }
// CreateUser creates a user and sets several fields of the returned User struct to // CreateUser creates a user and sets several fields of the returned User struct to
@@ -381,28 +362,6 @@ func (a *App) CreateOAuthUser(c *request.Context, service string, userData io.Re
return ruser, nil return ruser, nil
} }
// CheckEmailDomain checks that an email domain matches a list of space-delimited domains as a string.
func CheckEmailDomain(email string, domains string) bool {
if domains == "" {
return true
}
domainArray := strings.Fields(strings.TrimSpace(strings.ToLower(strings.Replace(strings.Replace(domains, "@", " ", -1), ",", " ", -1))))
for _, d := range domainArray {
if strings.HasSuffix(strings.ToLower(email), "@"+d) {
return true
}
}
return false
}
// CheckUserDomain checks that a user's email domain matches a list of space-delimited domains as a string.
func CheckUserDomain(user *model.User, domains string) bool {
return CheckEmailDomain(user.Email, domains)
}
func (a *App) GetUser(userID string) (*model.User, *model.AppError) { func (a *App) GetUser(userID string) (*model.User, *model.AppError) {
user, err := a.srv.userService.GetUser(userID) user, err := a.srv.userService.GetUser(userID)
if err != nil { if err != nil {
@@ -1190,13 +1149,13 @@ func (a *App) UpdateUser(user *model.User, sendNotifications bool) (*model.User,
var newEmail string var newEmail string
if user.Email != prev.Email { if user.Email != prev.Email {
if !CheckUserDomain(user, *a.Config().TeamSettings.RestrictCreationToDomains) { if !users.CheckUserDomain(user, *a.Config().TeamSettings.RestrictCreationToDomains) {
if !prev.IsGuest() && !prev.IsLDAPUser() && !prev.IsSAMLUser() { if !prev.IsGuest() && !prev.IsLDAPUser() && !prev.IsSAMLUser() {
return nil, model.NewAppError("UpdateUser", "api.user.update_user.accepted_domain.app_error", nil, "", http.StatusBadRequest) return nil, model.NewAppError("UpdateUser", "api.user.update_user.accepted_domain.app_error", nil, "", http.StatusBadRequest)
} }
} }
if !CheckUserDomain(user, *a.Config().GuestAccountsSettings.RestrictCreationToDomains) { if !users.CheckUserDomain(user, *a.Config().GuestAccountsSettings.RestrictCreationToDomains) {
if prev.IsGuest() && !prev.IsLDAPUser() && !prev.IsSAMLUser() { if prev.IsGuest() && !prev.IsLDAPUser() && !prev.IsSAMLUser() {
return nil, model.NewAppError("UpdateUser", "api.user.update_user.accepted_guest_domain.app_error", nil, "", http.StatusBadRequest) return nil, model.NewAppError("UpdateUser", "api.user.update_user.accepted_guest_domain.app_error", nil, "", http.StatusBadRequest)
} }

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

@@ -27,34 +27,6 @@ import (
"github.com/mattermost/mattermost-server/v5/utils/testutils" "github.com/mattermost/mattermost-server/v5/utils/testutils"
) )
func TestCheckUserDomain(t *testing.T) {
th := Setup(t).InitBasic()
defer th.TearDown()
user := th.BasicUser
cases := []struct {
domains string
matched bool
}{
{"simulator.amazonses.com", true},
{"gmail.com", false},
{"", true},
{"gmail.com simulator.amazonses.com", true},
}
for _, c := range cases {
matched := CheckUserDomain(user, c.domains)
if matched != c.matched {
if c.matched {
t.Logf("'%v' should have matched '%v'", user.Email, c.domains)
} else {
t.Logf("'%v' should not have matched '%v'", user.Email, c.domains)
}
t.FailNow()
}
}
}
func TestCreateOAuthUser(t *testing.T) { func TestCreateOAuthUser(t *testing.T) {
th := Setup(t).InitBasic() th := Setup(t).InitBasic()
defer th.TearDown() defer th.TearDown()

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

@@ -253,7 +253,7 @@ func (wc *WebConn) SetSession(v *model.Session) {
// Pump starts the WebConn instance. After this, the websocket // Pump starts the WebConn instance. After this, the websocket
// is ready to send/receive messages. // is ready to send/receive messages.
func (wc *WebConn) Pump() { func (wc *WebConn) Pump() {
defer ReturnSessionToPool(wc.GetSession()) defer wc.App.srv.userService.ReturnSessionToPool(wc.GetSession())
var wg sync.WaitGroup var wg sync.WaitGroup
wg.Add(1) wg.Add(1)

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

@@ -255,17 +255,7 @@ func (a *App) invalidateCacheForChannelPosts(channelID string) {
func (a *App) InvalidateCacheForUser(userID string) { func (a *App) InvalidateCacheForUser(userID string) {
a.Srv().invalidateCacheForUserSkipClusterSend(userID) a.Srv().invalidateCacheForUserSkipClusterSend(userID)
a.Srv().Store.User().InvalidateProfilesInChannelCacheByUser(userID) a.srv.userService.InvalidateCacheForUser(userID)
a.Srv().Store.User().InvalidateProfileCacheForUser(userID)
if a.Cluster() != nil {
msg := &model.ClusterMessage{
Event: model.CLUSTER_EVENT_INVALIDATE_CACHE_FOR_USER,
SendType: model.CLUSTER_SEND_BEST_EFFORT,
Data: userID,
}
a.Cluster().SendClusterMessage(msg)
}
} }
func (a *App) invalidateCacheForUserTeams(userID string) { func (a *App) invalidateCacheForUserTeams(userID string) {

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

@@ -16,6 +16,7 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/model"
"github.com/mattermost/mattermost-server/v5/services/users"
"github.com/mattermost/mattermost-server/v5/shared/i18n" "github.com/mattermost/mattermost-server/v5/shared/i18n"
"github.com/mattermost/mattermost-server/v5/store/storetest/mocks" "github.com/mattermost/mattermost-server/v5/store/storetest/mocks"
) )
@@ -158,6 +159,16 @@ func TestHubSessionRevokeRace(t *testing.T) {
mockStore.On("Post").Return(&mockPostStore) mockStore.On("Post").Return(&mockPostStore)
mockStore.On("System").Return(&mockSystemStore) mockStore.On("System").Return(&mockSystemStore)
userService, err := users.New(users.ServiceInitializer{
UserStore: &mockUserStore,
SessionStore: &mockSessionStore,
ConfigFn: th.App.srv.Config,
Metrics: th.App.Metrics(),
Cluster: th.App.Cluster(),
})
require.NoError(t, err)
th.App.srv.userService = userService
// This needs to be false for the condition to trigger // This needs to be false for the condition to trigger
th.App.UpdateConfig(func(cfg *model.Config) { th.App.UpdateConfig(func(cfg *model.Config) {
*cfg.ServiceSettings.ExtendSessionLengthWithActivity = false *cfg.ServiceSettings.ExtendSessionLengthWithActivity = false
@@ -174,7 +185,7 @@ func TestHubSessionRevokeRace(t *testing.T) {
time.Sleep(time.Second) time.Sleep(time.Second)
// We override the LastActivityAt which happens in NewWebConn. // We override the LastActivityAt which happens in NewWebConn.
// This is needed to call RevokeSessionById which triggers the race. // This is needed to call RevokeSessionById which triggers the race.
th.App.AddSessionToCache(sess1) th.App.srv.userService.AddSessionToCache(sess1)
go func() { go func() {
for i := 0; i <= broadcastQueueSize; i++ { for i := 0; i <= broadcastQueueSize; i++ {

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

@@ -16,6 +16,7 @@ import (
"github.com/mattermost/mattermost-server/v5/app/request" "github.com/mattermost/mattermost-server/v5/app/request"
"github.com/mattermost/mattermost-server/v5/audit" "github.com/mattermost/mattermost-server/v5/audit"
"github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/model"
"github.com/mattermost/mattermost-server/v5/services/users"
) )
var UserCmd = &cobra.Command{ var UserCmd = &cobra.Command{
@@ -480,7 +481,7 @@ func getUpdatedUserModel(command *cobra.Command, a *app.App, user *model.User) (
user.Locale = locale user.Locale = locale
} }
if !user.IsLDAPUser() && !user.IsSAMLUser() && !app.CheckUserDomain(user, *a.Config().TeamSettings.RestrictCreationToDomains) { if !user.IsLDAPUser() && !user.IsSAMLUser() && !users.CheckUserDomain(user, *a.Config().TeamSettings.RestrictCreationToDomains) {
return nil, errors.New("The email does not belong to an accepted domain.") return nil, errors.New("The email does not belong to an accepted domain.")
} }

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

@@ -8,12 +8,14 @@ import (
"io/ioutil" "io/ioutil"
"os" "os"
"path/filepath" "path/filepath"
"runtime"
"sync" "sync"
"testing" "testing"
"github.com/mattermost/mattermost-server/v5/app/request" "github.com/mattermost/mattermost-server/v5/app/request"
"github.com/mattermost/mattermost-server/v5/config" "github.com/mattermost/mattermost-server/v5/config"
"github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/model"
"github.com/mattermost/mattermost-server/v5/services/cache"
"github.com/mattermost/mattermost-server/v5/store" "github.com/mattermost/mattermost-server/v5/store"
) )
@@ -72,9 +74,22 @@ func setupTestHelper(s store.Store, includeCacheLayer bool, tb testing.TB) *Test
configStore.Set(config) configStore.Set(config)
buffer := &bytes.Buffer{} buffer := &bytes.Buffer{}
provider := cache.NewProvider()
cache, err := provider.NewCache(&cache.CacheOptions{
Size: model.SESSION_CACHE_SIZE,
Striped: true,
StripedBuckets: maxInt(runtime.NumCPU()-1, 1),
})
if err != nil {
panic(err)
}
return &TestHelper{ return &TestHelper{
service: &UserService{store: s.User(), config: configStore.Get}, service: &UserService{
store: s.User(),
sessionStore: s.Session(),
sessionCache: cache,
config: configStore.Get,
},
Context: &request.Context{}, Context: &request.Context{},
configStore: configStore, configStore: configStore,
dbStore: s, dbStore: s,

100
services/users/session.go Обычный файл
Просмотреть файл

@@ -0,0 +1,100 @@
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
// See LICENSE.txt for license information.
package users
import (
"time"
"github.com/mattermost/mattermost-server/v5/model"
)
func (us *UserService) ReturnSessionToPool(session *model.Session) {
if session != nil {
session.Id = ""
us.sessionPool.Put(session)
}
}
func (us *UserService) CreateSession(session *model.Session) (*model.Session, error) {
session.Token = ""
session, err := us.sessionStore.Save(session)
if err != nil {
return nil, err
}
us.AddSessionToCache(session)
return session, nil
}
func (us *UserService) GetSession(token string) (*model.Session, error) {
var session = us.sessionPool.Get().(*model.Session)
if err := us.sessionCache.Get(token, session); err == nil {
if us.metrics != nil {
us.metrics.IncrementMemCacheHitCounterSession()
}
} else {
if us.metrics != nil {
us.metrics.IncrementMemCacheMissCounterSession()
}
}
return session, nil
}
func (us *UserService) AddSessionToCache(session *model.Session) {
us.sessionCache.SetWithExpiry(session.Token, session, time.Duration(int64(*us.config().ServiceSettings.SessionCacheInMinutes))*time.Minute)
}
func (us *UserService) SessionCacheLength() int {
if l, err := us.sessionCache.Len(); err == nil {
return l
}
return 0
}
func (us *UserService) ClearUserSessionCacheLocal(userID string) {
if keys, err := us.sessionCache.Keys(); err == nil {
var session *model.Session
for _, key := range keys {
if err := us.sessionCache.Get(key, &session); err == nil {
if session.UserId == userID {
us.sessionCache.Remove(key)
if us.metrics != nil {
us.metrics.IncrementMemCacheInvalidationCounterSession()
}
}
}
}
}
}
func (us *UserService) ClearAllUsersSessionCacheLocal() {
us.sessionCache.Purge()
}
func (us *UserService) ClearUserSessionCache(userID string) {
us.ClearUserSessionCacheLocal(userID)
if us.cluster != nil {
msg := &model.ClusterMessage{
Event: model.CLUSTER_EVENT_CLEAR_SESSION_CACHE_FOR_USER,
SendType: model.CLUSTER_SEND_RELIABLE,
Data: userID,
}
us.cluster.SendClusterMessage(msg)
}
}
func (us *UserService) ClearAllUsersSessionCache() {
us.ClearAllUsersSessionCacheLocal()
if us.cluster != nil {
msg := &model.ClusterMessage{
Event: model.CLUSTER_EVENT_CLEAR_SESSION_CACHE_FOR_ALL_USERS,
SendType: model.CLUSTER_SEND_RELIABLE,
}
us.cluster.SendClusterMessage(msg)
}
}

49
services/users/session_test.go Обычный файл
Просмотреть файл

@@ -0,0 +1,49 @@
// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved.
// See LICENSE.txt for license information.
package users
import (
"testing"
"time"
"github.com/mattermost/mattermost-server/v5/model"
"github.com/stretchr/testify/require"
)
func TestCache(t *testing.T) {
th := Setup(t)
defer th.TearDown()
session := &model.Session{
Id: model.NewId(),
Token: model.NewId(),
UserId: model.NewId(),
}
session2 := &model.Session{
Id: model.NewId(),
Token: model.NewId(),
UserId: model.NewId(),
}
th.service.sessionCache.SetWithExpiry(session.Token, session, 5*time.Minute)
th.service.sessionCache.SetWithExpiry(session2.Token, session2, 5*time.Minute)
keys, err := th.service.sessionCache.Keys()
require.NoError(t, err)
require.NotEmpty(t, keys)
th.service.ClearUserSessionCache(session.UserId)
rkeys, err := th.service.sessionCache.Keys()
require.NoError(t, err)
require.Lenf(t, rkeys, len(keys)-1, "should have one less: %d - %d != 1", len(keys), len(rkeys))
require.NotEmpty(t, rkeys)
th.service.ClearAllUsersSessionCache()
rkeys, err = th.service.sessionCache.Keys()
require.NoError(t, err)
require.Empty(t, rkeys)
}

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

@@ -5,17 +5,27 @@ package users
import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"runtime"
"sync"
"github.com/mattermost/mattermost-server/v5/einterfaces"
"github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/model"
"github.com/mattermost/mattermost-server/v5/services/cache"
"github.com/mattermost/mattermost-server/v5/shared/i18n" "github.com/mattermost/mattermost-server/v5/shared/i18n"
"github.com/mattermost/mattermost-server/v5/shared/mlog" "github.com/mattermost/mattermost-server/v5/shared/mlog"
"github.com/mattermost/mattermost-server/v5/store" "github.com/mattermost/mattermost-server/v5/store"
) )
type UserService struct { type UserService struct {
store store.UserStore store store.UserStore
config func() *model.Config sessionStore store.SessionStore
sessionCache cache.Cache
sessionPool sync.Pool
metrics einterfaces.MetricsInterface
cluster einterfaces.ClusterInterface
config func() *model.Config
} }
type UserCreateOptions struct { type UserCreateOptions struct {
@@ -23,11 +33,49 @@ type UserCreateOptions struct {
FromImport bool FromImport bool
} }
func New(s store.UserStore, cfgFn func() *model.Config) *UserService { // ServiceInitializer is used to initialize the UserService.
return &UserService{ type ServiceInitializer struct {
store: s, // Mandatory fields
config: cfgFn, UserStore store.UserStore
SessionStore store.SessionStore
ConfigFn func() *model.Config
// Optional fields
Metrics einterfaces.MetricsInterface
Cluster einterfaces.ClusterInterface
}
func New(initializer ServiceInitializer) (*UserService, error) {
cacheProvider := cache.NewProvider()
if err := cacheProvider.Connect(); err != nil {
return nil, fmt.Errorf("could not create cache provider: %w", err)
} }
sessionCache, err := cacheProvider.NewCache(&cache.CacheOptions{
Size: model.SESSION_CACHE_SIZE,
Striped: true,
StripedBuckets: maxInt(runtime.NumCPU()-1, 1),
})
if err != nil {
return nil, fmt.Errorf("could not create session cache: %w", err)
}
if initializer.ConfigFn == nil || initializer.UserStore == nil || initializer.SessionStore == nil {
return nil, errors.New("required parameters are not provided")
}
return &UserService{
store: initializer.UserStore,
sessionStore: initializer.SessionStore,
config: initializer.ConfigFn,
metrics: initializer.Metrics,
cluster: initializer.Cluster,
sessionCache: sessionCache,
sessionPool: sync.Pool{
New: func() interface{} {
return &model.Session{}
},
},
}, nil
} }
// CreateUser creates a user // CreateUser creates a user
@@ -41,11 +89,11 @@ func (us *UserService) CreateUser(user *model.User, opts UserCreateOptions) (*mo
user.Roles = model.SYSTEM_GUEST_ROLE_ID user.Roles = model.SYSTEM_GUEST_ROLE_ID
} }
if !user.IsLDAPUser() && !user.IsSAMLUser() && !user.IsGuest() && !checkUserDomain(user, *us.config().TeamSettings.RestrictCreationToDomains) { if !user.IsLDAPUser() && !user.IsSAMLUser() && !user.IsGuest() && !CheckUserDomain(user, *us.config().TeamSettings.RestrictCreationToDomains) {
return nil, AcceptedDomainError return nil, AcceptedDomainError
} }
if !user.IsLDAPUser() && !user.IsSAMLUser() && user.IsGuest() && !checkUserDomain(user, *us.config().GuestAccountsSettings.RestrictCreationToDomains) { if !user.IsLDAPUser() && !user.IsSAMLUser() && user.IsGuest() && !CheckUserDomain(user, *us.config().GuestAccountsSettings.RestrictCreationToDomains) {
return nil, AcceptedDomainError return nil, AcceptedDomainError
} }
@@ -194,3 +242,17 @@ func (us *UserService) GetUsersWithoutTeam(options *model.UserGetOptions) ([]*mo
return users, nil return users, nil
} }
func (us *UserService) InvalidateCacheForUser(userID string) {
us.store.InvalidateProfilesInChannelCacheByUser(userID)
us.store.InvalidateProfileCacheForUser(userID)
if us.cluster != nil {
msg := &model.ClusterMessage{
Event: model.CLUSTER_EVENT_INVALIDATE_CACHE_FOR_USER,
SendType: model.CLUSTER_SEND_BEST_EFFORT,
Data: userID,
}
us.cluster.SendClusterMessage(msg)
}
}

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

@@ -9,13 +9,38 @@ import (
"github.com/mattermost/mattermost-server/v5/model" "github.com/mattermost/mattermost-server/v5/model"
) )
// checkUserDomain checks that a user's email domain matches a list of space-delimited domains as a string. func maxInt(a, b int) int {
func checkUserDomain(user *model.User, domains string) bool { if a > b {
return checkEmailDomain(user.Email, domains) return a
}
return b
} }
// checkEmailDomain checks that an email domain matches a list of space-delimited domains as a string. func (us *UserService) IsFirstUserAccount() bool {
func checkEmailDomain(email string, domains string) bool { cachedSessions, err := us.sessionCache.Len()
if err != nil {
return false
}
if cachedSessions == 0 {
count, err := us.store.Count(model.UserCountOptions{IncludeDeleted: true})
if err != nil {
return false
}
if count <= 0 {
return true
}
}
return false
}
// CheckUserDomain checks that a user's email domain matches a list of space-delimited domains as a string.
func CheckUserDomain(user *model.User, domains string) bool {
return CheckEmailDomain(user.Email, domains)
}
// CheckEmailDomain checks that an email domain matches a list of space-delimited domains as a string.
func CheckEmailDomain(email string, domains string) bool {
if domains == "" { if domains == "" {
return true return true
} }

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

@@ -87,6 +87,8 @@ func GetMockStoreForSetupFunctions() *mocks.Store {
roleStore := mocks.RoleStore{} roleStore := mocks.RoleStore{}
roleStore.On("GetAll").Return([]*model.Role{}, nil) roleStore.On("GetAll").Return([]*model.Role{}, nil)
sessionStore := mocks.SessionStore{}
mockStore.On("System").Return(&systemStore) mockStore.On("System").Return(&systemStore)
mockStore.On("User").Return(&userStore) mockStore.On("User").Return(&userStore)
mockStore.On("Post").Return(&postStore) mockStore.On("Post").Return(&postStore)
@@ -98,5 +100,6 @@ func GetMockStoreForSetupFunctions() *mocks.Store {
mockStore.On("Close").Return(nil) mockStore.On("Close").Return(nil)
mockStore.On("DropAllTables").Return(nil) mockStore.On("DropAllTables").Return(nil)
mockStore.On("MarkSystemRanUnitTests").Return(nil) mockStore.On("MarkSystemRanUnitTests").Return(nil)
mockStore.On("Session").Return(&sessionStore)
return &mockStore return &mockStore
} }

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

@@ -206,7 +206,7 @@ func (h Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
if token != "" && tokenLocation != app.TokenLocationCloudHeader && tokenLocation != app.TokenLocationRemoteClusterHeader { if token != "" && tokenLocation != app.TokenLocationCloudHeader && tokenLocation != app.TokenLocationRemoteClusterHeader {
session, err := c.App.GetSession(token) session, err := c.App.GetSession(token)
defer app.ReturnSessionToPool(session) defer c.App.ReturnSessionToPool(session)
if err != nil { if err != nil {
c.Logger.Info("Invalid session", mlog.Err(err)) c.Logger.Info("Invalid session", mlog.Err(err))

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

@@ -29,7 +29,7 @@ func (wh webSocketHandler) ServeWebSocket(conn *app.WebConn, r *model.WebSocketR
return return
} }
session, sessionErr := wh.app.GetSession(conn.GetSessionToken()) session, sessionErr := wh.app.GetSession(conn.GetSessionToken())
defer app.ReturnSessionToPool(session) defer wh.app.ReturnSessionToPool(session)
if sessionErr != nil { if sessionErr != nil {
mlog.Error( mlog.Error(