* use sync.pool for session

* added back to sync pool

* reverted change

* added a new line

* added back session object

* added back session object

* added back session object

* revert

* refactored into function

* added the session object back into the pool

* work in progress

* work in progress

* work in progress

* code review comments

Co-authored-by: Arjuna Marambe <arjunam@buildxact.com>
Co-authored-by: Mattermod <mattermod@users.noreply.github.com>
Этот коммит содержится в:
Arjuna Marambe
2021-01-28 05:58:24 +11:00
коммит произвёл GitHub
родитель 7f8850398c
Коммит 5f16fc644a
8 изменённых файлов: 33 добавлений и 13 удалений

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

@@ -517,6 +517,8 @@ func (a *App) RegenerateOAuthAppSecret(app *model.OAuthApp) (*model.OAuthApp, *m
func (a *App) RevokeAccessToken(token string) *model.AppError {
session, _ := a.GetSession(token)
defer ReturnSessionToPool(session)
schan := make(chan error, 1)
go func() {
schan <- a.Srv().Store.Session().Remove(token)

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

@@ -130,6 +130,8 @@ func (a *App) servePluginRequest(w http.ResponseWriter, r *http.Request, handler
r.Header.Del("Mattermost-User-Id")
if token != "" {
session, err := a.GetSession(token)
defer ReturnSessionToPool(session)
csrfCheckPassed := false
if err == nil && cookieAuth && r.Method != "GET" {
@@ -180,7 +182,7 @@ func (a *App) servePluginRequest(w http.ResponseWriter, r *http.Request, handler
csrfCheckPassed = true
}
if session != nil && err == nil && csrfCheckPassed {
if (session != nil && session.Id != "") && err == nil && csrfCheckPassed {
r.Header.Set("Mattermost-User-Id", session.UserId)
context.SessionId = session.Id
}

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

@@ -9,6 +9,7 @@ import (
"math"
"net/http"
"os"
"sync"
"time"
"github.com/mattermost/mattermost-server/v5/audit"
@@ -36,6 +37,19 @@ func (a *App) CreateSession(session *model.Session) (*model.Session, *model.AppE
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) {
apiKey := os.Getenv("MM_CLOUD_API_KEY")
if apiKey != "" && apiKey == token {
@@ -54,9 +68,10 @@ func (a *App) GetCloudSession(token string) (*model.Session, *model.AppError) {
func (a *App) GetSession(token string) (*model.Session, *model.AppError) {
metrics := a.Metrics()
var session *model.Session
var session = userSessionPool.Get().(*model.Session)
var err *model.AppError
if err := a.Srv().sessionCache.Get(token, &session); err == nil {
if err := a.Srv().sessionCache.Get(token, session); err == nil {
if metrics != nil {
metrics.IncrementMemCacheHitCounterSession()
}
@@ -66,7 +81,7 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) {
}
}
if session == nil {
if session.Id == "" {
var nErr error
if session, nErr = a.Srv().Store.Session().Get(token); nErr == nil {
if session != nil {
@@ -83,7 +98,7 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) {
}
}
if session == nil {
if session == nil || session.Id == "" {
session, err = a.createSessionForUserAccessToken(token)
if err != nil {
detailedError := ""
@@ -98,7 +113,7 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) {
}
}
if session == nil || session.IsExpired() {
if session.Id == "" || session.IsExpired() {
return nil, model.NewAppError("GetSession", "api.context.invalid_token.error", map[string]interface{}{"Token": token, "Error": ""}, "session is either nil or expired", http.StatusUnauthorized)
}

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

@@ -134,6 +134,8 @@ func (wc *WebConn) Pump() {
wg.Wait()
wc.App.HubUnregister(wc)
close(wc.pumpFinished)
defer ReturnSessionToPool(wc.GetSession())
}
func (wc *WebConn) readPump() {

5
services/cache/lru.go поставляемый
Просмотреть файл

@@ -210,11 +210,6 @@ func (l *LRU) get(key string, value interface{}) error {
_, err := u.UnmarshalMsg(val)
*v = &u
return err
case **model.Session:
var s model.Session
_, err := s.UnmarshalMsg(val)
*v = &s
return err
case *map[string]*model.User:
var u model.UserMap
_, err := u.UnmarshalMsg(val)

4
services/cache/lru_test.go поставляемый
Просмотреть файл

@@ -227,9 +227,9 @@ func TestLRUMarshalUnMarshal(t *testing.T) {
err = l.Set("session", session)
require.Nil(t, err)
var s = &model.Session{}
err = l.Get("session", s)
var s *model.Session
err = l.Get("session", &s)
require.Nil(t, err)
require.Equal(t, session, s)

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

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

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

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