diff --git a/app/session.go b/app/session.go index b70a7bae8e..fb1fdf6255 100644 --- a/app/session.go +++ b/app/session.go @@ -28,6 +28,7 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) { metrics := a.Metrics var session *model.Session + var err *model.AppError if ts, ok := a.Srv.sessionCache.Get(token); ok { session = ts.(*model.Session) if metrics != nil { @@ -40,9 +41,7 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) { } if session == nil { - if sessionResult := <-a.Srv.Store.Session().Get(token); sessionResult.Err == nil { - session = sessionResult.Data.(*model.Session) - + if session, err = a.Srv.Store.Session().Get(token); err == nil { if session != nil { if session.Token != token { return nil, model.NewAppError("GetSession", "api.context.invalid_token.error", map[string]interface{}{"Token": token, "Error": ""}, "", http.StatusUnauthorized) @@ -52,13 +51,12 @@ func (a *App) GetSession(token string) (*model.Session, *model.AppError) { a.AddSessionToCache(session) } } - } else if sessionResult.Err.StatusCode == http.StatusInternalServerError { - return nil, sessionResult.Err + } else if err.StatusCode == http.StatusInternalServerError { + return nil, err } } if session == nil { - var err *model.AppError session, err = a.createSessionForUserAccessToken(token) if err != nil { detailedError := "" @@ -180,21 +178,21 @@ func (a *App) RevokeSessionsForDeviceId(userId string, deviceId string, currentS } func (a *App) GetSessionById(sessionId string) (*model.Session, *model.AppError) { - result := <-a.Srv.Store.Session().Get(sessionId) - if result.Err != nil { - result.Err.StatusCode = http.StatusBadRequest - return nil, result.Err + session, err := a.Srv.Store.Session().Get(sessionId) + if err != nil { + err.StatusCode = http.StatusBadRequest + return nil, err } - return result.Data.(*model.Session), nil + return session, nil } func (a *App) RevokeSessionById(sessionId string) *model.AppError { - result := <-a.Srv.Store.Session().Get(sessionId) - if result.Err != nil { - result.Err.StatusCode = http.StatusBadRequest - return result.Err + session, err := a.Srv.Store.Session().Get(sessionId) + if err != nil { + err.StatusCode = http.StatusBadRequest + return err } - return a.RevokeSession(result.Data.(*model.Session)) + return a.RevokeSession(session) } @@ -315,9 +313,7 @@ func (a *App) createSessionForUserAccessToken(tokenString string) (*model.Sessio func (a *App) RevokeUserAccessToken(token *model.UserAccessToken) *model.AppError { var session *model.Session - if result := <-a.Srv.Store.Session().Get(token.Token); result.Err == nil { - session = result.Data.(*model.Session) - } + session, _ = a.Srv.Store.Session().Get(token.Token) if result := <-a.Srv.Store.UserAccessToken().Delete(token.Id); result.Err != nil { return result.Err @@ -332,9 +328,7 @@ func (a *App) RevokeUserAccessToken(token *model.UserAccessToken) *model.AppErro func (a *App) DisableUserAccessToken(token *model.UserAccessToken) *model.AppError { var session *model.Session - if result := <-a.Srv.Store.Session().Get(token.Token); result.Err == nil { - session = result.Data.(*model.Session) - } + session, _ = a.Srv.Store.Session().Get(token.Token) if result := <-a.Srv.Store.UserAccessToken().UpdateTokenDisable(token.Id); result.Err != nil { return result.Err @@ -349,9 +343,7 @@ func (a *App) DisableUserAccessToken(token *model.UserAccessToken) *model.AppErr func (a *App) EnableUserAccessToken(token *model.UserAccessToken) *model.AppError { var session *model.Session - if result := <-a.Srv.Store.Session().Get(token.Token); result.Err == nil { - session = result.Data.(*model.Session) - } + session, _ = a.Srv.Store.Session().Get(token.Token) if result := <-a.Srv.Store.UserAccessToken().UpdateTokenEnable(token.Id); result.Err != nil { return result.Err diff --git a/store/sqlstore/session_store.go b/store/sqlstore/session_store.go index 23d771e54a..427e669d92 100644 --- a/store/sqlstore/session_store.go +++ b/store/sqlstore/session_store.go @@ -75,32 +75,28 @@ func (me SqlSessionStore) Save(session *model.Session) (*model.Session, *model.A return session, nil } -func (me SqlSessionStore) Get(sessionIdOrToken string) store.StoreChannel { - return store.Do(func(result *store.StoreResult) { - var sessions []*model.Session +func (me SqlSessionStore) Get(sessionIdOrToken string) (*model.Session, *model.AppError) { + var sessions []*model.Session - if _, err := me.GetReplica().Select(&sessions, "SELECT * FROM Sessions WHERE Token = :Token OR Id = :Id LIMIT 1", map[string]interface{}{"Token": sessionIdOrToken, "Id": sessionIdOrToken}); err != nil { - result.Err = model.NewAppError("SqlSessionStore.Get", "store.sql_session.get.app_error", nil, "sessionIdOrToken="+sessionIdOrToken+", "+err.Error(), http.StatusInternalServerError) - } else if len(sessions) == 0 { - result.Err = model.NewAppError("SqlSessionStore.Get", "store.sql_session.get.app_error", nil, "sessionIdOrToken="+sessionIdOrToken, http.StatusNotFound) - } else { - result.Data = sessions[0] + if _, err := me.GetReplica().Select(&sessions, "SELECT * FROM Sessions WHERE Token = :Token OR Id = :Id LIMIT 1", map[string]interface{}{"Token": sessionIdOrToken, "Id": sessionIdOrToken}); err != nil { + return nil, model.NewAppError("SqlSessionStore.Get", "store.sql_session.get.app_error", nil, "sessionIdOrToken="+sessionIdOrToken+", "+err.Error(), http.StatusInternalServerError) + } else if len(sessions) == 0 { + return nil, model.NewAppError("SqlSessionStore.Get", "store.sql_session.get.app_error", nil, "sessionIdOrToken="+sessionIdOrToken, http.StatusNotFound) + } + session := sessions[0] - tcs := me.Team().GetTeamsForUser(sessions[0].UserId) - if rtcs := <-tcs; rtcs.Err != nil { - result.Err = model.NewAppError("SqlSessionStore.Get", "store.sql_session.get.app_error", nil, "sessionIdOrToken="+sessionIdOrToken+", "+rtcs.Err.Error(), http.StatusInternalServerError) - return - } else { - tempMembers := rtcs.Data.([]*model.TeamMember) - sessions[0].TeamMembers = make([]*model.TeamMember, 0, len(tempMembers)) - for _, tm := range tempMembers { - if tm.DeleteAt == 0 { - sessions[0].TeamMembers = append(sessions[0].TeamMembers, tm) - } - } - } + rtcs := <-me.Team().GetTeamsForUser(sessions[0].UserId) + if rtcs.Err != nil { + return nil, model.NewAppError("SqlSessionStore.Get", "store.sql_session.get.app_error", nil, "sessionIdOrToken="+sessionIdOrToken+", "+rtcs.Err.Error(), http.StatusInternalServerError) + } + tempMembers := rtcs.Data.([]*model.TeamMember) + sessions[0].TeamMembers = make([]*model.TeamMember, 0, len(tempMembers)) + for _, tm := range tempMembers { + if tm.DeleteAt == 0 { + sessions[0].TeamMembers = append(sessions[0].TeamMembers, tm) } - }) + } + return session, nil } func (me SqlSessionStore) GetSessions(userId string) store.StoreChannel { diff --git a/store/store.go b/store/store.go index df776130c1..141f99c7be 100644 --- a/store/store.go +++ b/store/store.go @@ -312,8 +312,8 @@ type BotStore interface { } type SessionStore interface { + Get(sessionIdOrToken string) (*model.Session, *model.AppError) Save(session *model.Session) (*model.Session, *model.AppError) - Get(sessionIdOrToken string) StoreChannel GetSessions(userId string) StoreChannel GetSessionsWithActiveDeviceIds(userId string) ([]*model.Session, *model.AppError) Remove(sessionIdOrToken string) StoreChannel diff --git a/store/storetest/mocks/SessionStore.go b/store/storetest/mocks/SessionStore.go index 4e94fa28b8..2549a419a0 100644 --- a/store/storetest/mocks/SessionStore.go +++ b/store/storetest/mocks/SessionStore.go @@ -42,19 +42,28 @@ func (_m *SessionStore) Cleanup(expiryTime int64, batchSize int64) { } // Get provides a mock function with given fields: sessionIdOrToken -func (_m *SessionStore) Get(sessionIdOrToken string) store.StoreChannel { +func (_m *SessionStore) Get(sessionIdOrToken string) (*model.Session, *model.AppError) { ret := _m.Called(sessionIdOrToken) - var r0 store.StoreChannel - if rf, ok := ret.Get(0).(func(string) store.StoreChannel); ok { + var r0 *model.Session + if rf, ok := ret.Get(0).(func(string) *model.Session); ok { r0 = rf(sessionIdOrToken) } else { if ret.Get(0) != nil { - r0 = ret.Get(0).(store.StoreChannel) + r0 = ret.Get(0).(*model.Session) } } - return r0 + var r1 *model.AppError + if rf, ok := ret.Get(1).(func(string) *model.AppError); ok { + r1 = rf(sessionIdOrToken) + } else { + if ret.Get(1) != nil { + r1 = ret.Get(1).(*model.AppError) + } + } + + return r0, r1 } // GetSessions provides a mock function with given fields: userId diff --git a/store/storetest/oauth_store.go b/store/storetest/oauth_store.go index eb77f381e4..9ac49332bc 100644 --- a/store/storetest/oauth_store.go +++ b/store/storetest/oauth_store.go @@ -429,7 +429,7 @@ func testOAuthStoreDeleteApp(t *testing.T, ss store.Store) { t.Fatal(err) } - if err := (<-ss.Session().Get(s1.Token)).Err; err == nil { + if _, err := ss.Session().Get(s1.Token); err == nil { t.Fatal("should error - session should be deleted") } diff --git a/store/storetest/session_store.go b/store/storetest/session_store.go index 2bcee6ebf5..72a77cfde0 100644 --- a/store/storetest/session_store.go +++ b/store/storetest/session_store.go @@ -59,10 +59,10 @@ func testSessionGet(t *testing.T, ss store.Store) { s3, err = ss.Session().Save(s3) require.Nil(t, err) - if rs1 := (<-ss.Session().Get(s1.Id)); rs1.Err != nil { - t.Fatal(rs1.Err) + if session, err := ss.Session().Get(s1.Id); err != nil { + t.Fatal(err) } else { - if rs1.Data.(*model.Session).Id != s1.Id { + if session.Id != s1.Id { t.Fatal("should match") } } @@ -116,17 +116,17 @@ func testSessionRemove(t *testing.T, ss store.Store) { s1, err := ss.Session().Save(s1) require.Nil(t, err) - if rs1 := (<-ss.Session().Get(s1.Id)); rs1.Err != nil { - t.Fatal(rs1.Err) + if session, err := ss.Session().Get(s1.Id); err != nil { + t.Fatal(err) } else { - if rs1.Data.(*model.Session).Id != s1.Id { + if session.Id != s1.Id { t.Fatal("should match") } } store.Must(ss.Session().Remove(s1.Id)) - if rs2 := (<-ss.Session().Get(s1.Id)); rs2.Err == nil { + if _, err := ss.Session().Get(s1.Id); err == nil { t.Fatal("should have been removed") } } @@ -138,17 +138,17 @@ func testSessionRemoveAll(t *testing.T, ss store.Store) { s1, err := ss.Session().Save(s1) require.Nil(t, err) - if rs1 := (<-ss.Session().Get(s1.Id)); rs1.Err != nil { - t.Fatal(rs1.Err) + if session, err := ss.Session().Get(s1.Id); err != nil { + t.Fatal(err) } else { - if rs1.Data.(*model.Session).Id != s1.Id { + if session.Id != s1.Id { t.Fatal("should match") } } store.Must(ss.Session().RemoveAllSessions()) - if rs2 := (<-ss.Session().Get(s1.Id)); rs2.Err == nil { + if _, err := ss.Session().Get(s1.Id); err == nil { t.Fatal("should have been removed") } } @@ -160,17 +160,17 @@ func testSessionRemoveByUser(t *testing.T, ss store.Store) { s1, err := ss.Session().Save(s1) require.Nil(t, err) - if rs1 := (<-ss.Session().Get(s1.Id)); rs1.Err != nil { - t.Fatal(rs1.Err) + if session, err := ss.Session().Get(s1.Id); err != nil { + t.Fatal(err) } else { - if rs1.Data.(*model.Session).Id != s1.Id { + if session.Id != s1.Id { t.Fatal("should match") } } store.Must(ss.Session().PermanentDeleteSessionsByUser(s1.UserId)) - if rs2 := (<-ss.Session().Get(s1.Id)); rs2.Err == nil { + if _, err := ss.Session().Get(s1.Id); err == nil { t.Fatal("should have been removed") } } @@ -182,17 +182,17 @@ func testSessionRemoveToken(t *testing.T, ss store.Store) { s1, err := ss.Session().Save(s1) require.Nil(t, err) - if rs1 := (<-ss.Session().Get(s1.Id)); rs1.Err != nil { - t.Fatal(rs1.Err) + if session, err := ss.Session().Get(s1.Id); err != nil { + t.Fatal(err) } else { - if rs1.Data.(*model.Session).Id != s1.Id { + if session.Id != s1.Id { t.Fatal("should match") } } store.Must(ss.Session().Remove(s1.Token)) - if rs2 := (<-ss.Session().Get(s1.Id)); rs2.Err == nil { + if _, err := ss.Session().Get(s1.Id); err == nil { t.Fatal("should have been removed") } @@ -260,10 +260,10 @@ func testSessionStoreUpdateLastActivityAt(t *testing.T, ss store.Store) { t.Fatal(err) } - if r1 := <-ss.Session().Get(s1.Id); r1.Err != nil { - t.Fatal(r1.Err) + if session, err := ss.Session().Get(s1.Id); err != nil { + t.Fatal(err) } else { - if r1.Data.(*model.Session).LastActivityAt != 1234567890 { + if session.LastActivityAt != 1234567890 { t.Fatal("LastActivityAt not updated correctly") } } @@ -320,16 +320,16 @@ func testSessionCleanup(t *testing.T, ss store.Store) { ss.Session().Cleanup(now, 1) - err = (<-ss.Session().Get(s1.Id)).Err + _, err = ss.Session().Get(s1.Id) assert.Nil(t, err) - err = (<-ss.Session().Get(s2.Id)).Err + _, err = ss.Session().Get(s2.Id) assert.Nil(t, err) - err = (<-ss.Session().Get(s3.Id)).Err + _, err = ss.Session().Get(s3.Id) assert.NotNil(t, err) - err = (<-ss.Session().Get(s4.Id)).Err + _, err = ss.Session().Get(s4.Id) assert.NotNil(t, err) store.Must(ss.Session().Remove(s1.Id)) diff --git a/store/storetest/user_access_token_store.go b/store/storetest/user_access_token_store.go index cc39ccd8c3..3af4a621fc 100644 --- a/store/storetest/user_access_token_store.go +++ b/store/storetest/user_access_token_store.go @@ -67,7 +67,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) { t.Fatal(result.Err) } - if err = (<-ss.Session().Get(s1.Token)).Err; err == nil { + if _, err = ss.Session().Get(s1.Token); err == nil { t.Fatal("should error - session should be deleted") } @@ -90,7 +90,7 @@ func testUserAccessTokenSaveGetDelete(t *testing.T, ss store.Store) { t.Fatal(result.Err) } - if err := (<-ss.Session().Get(s2.Token)).Err; err == nil { + if _, err := ss.Session().Get(s2.Token); err == nil { t.Fatal("should error - session should be deleted") } @@ -121,7 +121,7 @@ func testUserAccessTokenDisableEnable(t *testing.T, ss store.Store) { t.Fatal(err) } - if err = (<-ss.Session().Get(s1.Token)).Err; err == nil { + if _, err = ss.Session().Get(s1.Token); err == nil { t.Fatal("should error - session should be deleted") }