From 5f16fc644a856e5a2534d6d8acc74363589888a1 Mon Sep 17 00:00:00 2001 From: Arjuna Marambe Date: Thu, 28 Jan 2021 05:58:24 +1100 Subject: [PATCH] 15249 sync.pool (#16103) * 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 Co-authored-by: Mattermod --- app/oauth.go | 2 ++ app/plugin_requests.go | 4 +++- app/session.go | 25 ++++++++++++++++++++----- app/web_conn.go | 2 ++ services/cache/lru.go | 5 ----- services/cache/lru_test.go | 4 ++-- web/handlers.go | 2 ++ wsapi/websocket_handler.go | 2 ++ 8 files changed, 33 insertions(+), 13 deletions(-) diff --git a/app/oauth.go b/app/oauth.go index 9126247371..720d09d88f 100644 --- a/app/oauth.go +++ b/app/oauth.go @@ -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) diff --git a/app/plugin_requests.go b/app/plugin_requests.go index 7e0a64ed14..e517d435df 100644 --- a/app/plugin_requests.go +++ b/app/plugin_requests.go @@ -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 } diff --git a/app/session.go b/app/session.go index ef7bb47a5f..97016f3712 100644 --- a/app/session.go +++ b/app/session.go @@ -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) } diff --git a/app/web_conn.go b/app/web_conn.go index 60b69e9eb2..930847b876 100644 --- a/app/web_conn.go +++ b/app/web_conn.go @@ -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() { diff --git a/services/cache/lru.go b/services/cache/lru.go index 75a0152f39..ac13dddfdd 100644 --- a/services/cache/lru.go +++ b/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) diff --git a/services/cache/lru_test.go b/services/cache/lru_test.go index 7eab51ffc5..429846b5e9 100644 --- a/services/cache/lru_test.go +++ b/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) diff --git a/web/handlers.go b/web/handlers.go index c2a718558c..b733f914a6 100644 --- a/web/handlers.go +++ b/web/handlers.go @@ -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 { diff --git a/wsapi/websocket_handler.go b/wsapi/websocket_handler.go index 46ed07f67b..d8b4bb7693 100644 --- a/wsapi/websocket_handler.go +++ b/wsapi/websocket_handler.go @@ -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",