refactor: only delete fallback session cookies if set

This commit is contained in:
Trong Huu Nguyen
2021-11-01 10:56:49 +01:00
parent 325caeac34
commit b85ea7136e
3 changed files with 71 additions and 43 deletions
+7 -5
View File
@@ -37,7 +37,7 @@ func (h *Handler) getSessionFromCookie(w http.ResponseWriter, r *http.Request) (
return nil, fmt.Errorf("decrypting session data: %w", err)
}
h.DeleteSessionFallback(w)
h.DeleteSessionFallback(w, r)
return sessionData, nil
}
@@ -45,12 +45,13 @@ func (h *Handler) getSessionFromCookie(w http.ResponseWriter, r *http.Request) (
return nil, fmt.Errorf("session not found in store: %w", err)
}
log.Warnf("get session: store is unavailable: %+v; using cookie fallback", err)
fallbackSessionData, err := h.GetSessionFallback(r)
if err != nil {
return nil, fmt.Errorf("fallback session not found: %w", err)
}
log.Warnf("get session: store is unavailable: %+v; using cookie fallback", err)
return fallbackSessionData, nil
}
@@ -93,16 +94,17 @@ func (h *Handler) createSession(w http.ResponseWriter, r *http.Request, external
err = h.Sessions.Write(r.Context(), sessionID, encryptedSessionData, sessionLifetime)
if err == nil {
h.DeleteSessionFallback(w)
h.DeleteSessionFallback(w, r)
return nil
}
log.Warnf("create session: store is unavailable: %+v; using cookie fallback", err)
err = h.SetSessionFallback(w, sessionData, sessionLifetime)
if err != nil {
return fmt.Errorf("writing session to fallback store: %w", err)
}
log.Warnf("create session: store is unavailable: %+v; using cookie fallback", err)
return nil
}
@@ -112,6 +114,6 @@ func (h *Handler) destroySession(w http.ResponseWriter, r *http.Request, session
return fmt.Errorf("deleting session from store: %w", err)
}
h.DeleteSessionFallback(w)
h.DeleteSessionFallback(w, r)
return nil
}
+17 -7
View File
@@ -1,6 +1,7 @@
package router
import (
"errors"
"fmt"
"net/http"
"time"
@@ -9,15 +10,15 @@ import (
)
func (h *Handler) SessionFallbackExternalIDCookieName() string {
return h.GetSessionCookieName() + ".eid"
return h.GetSessionCookieName() + ".1"
}
func (h *Handler) SessionFallbackIDTokenCookieName() string {
return h.GetSessionCookieName() + ".id_token"
return h.GetSessionCookieName() + ".2"
}
func (h *Handler) SessionFallbackAccessTokenCookieName() string {
return h.GetSessionCookieName() + ".access_token"
return h.GetSessionCookieName() + ".3"
}
func (h *Handler) SetSessionFallback(w http.ResponseWriter, data *session.Data, expiresIn time.Duration) error {
@@ -58,8 +59,17 @@ func (h *Handler) GetSessionFallback(r *http.Request) (*session.Data, error) {
return session.NewData(externalSessionID, accessToken, idToken), nil
}
func (h *Handler) DeleteSessionFallback(w http.ResponseWriter) {
h.deleteCookie(w, h.SessionFallbackAccessTokenCookieName())
h.deleteCookie(w, h.SessionFallbackExternalIDCookieName())
h.deleteCookie(w, h.SessionFallbackIDTokenCookieName())
func (h *Handler) DeleteSessionFallback(w http.ResponseWriter, r *http.Request) {
deleteIfNotFound := func(h *Handler, w http.ResponseWriter, cookieName string) {
_, err := r.Cookie(cookieName)
if errors.Is(err, http.ErrNoCookie) {
return
}
h.deleteCookie(w, cookieName)
}
deleteIfNotFound(h, w, h.SessionFallbackAccessTokenCookieName())
deleteIfNotFound(h, w, h.SessionFallbackExternalIDCookieName())
deleteIfNotFound(h, w, h.SessionFallbackIDTokenCookieName())
}
+47 -31
View File
@@ -24,33 +24,12 @@ func TestHandler_GetSessionFallback(t *testing.T) {
})
t.Run("request with fallback session cookies", func(t *testing.T) {
// set up fallback session cookies
writer := httptest.NewRecorder()
expiresIn := time.Minute
data := session.NewData("sid", "access_token", "id_token")
err := h.SetSessionFallback(writer, data, expiresIn)
assert.NoError(t, err)
cookies := writer.Result().Cookies()
externalSessionIDCookie := getCookieFromJar(h.SessionFallbackExternalIDCookieName(), cookies)
assert.NotNil(t, externalSessionIDCookie)
idTokenCookie := getCookieFromJar(h.SessionFallbackIDTokenCookieName(), cookies)
assert.NotNil(t, idTokenCookie)
accessTokenCookie := getCookieFromJar(h.SessionFallbackAccessTokenCookieName(), cookies)
assert.NotNil(t, accessTokenCookie)
// make request with fallback session cookies set
r := httptest.NewRequest(http.MethodGet, "/", nil)
r.AddCookie(externalSessionIDCookie)
r.AddCookie(idTokenCookie)
r.AddCookie(accessTokenCookie)
r := makeRequestWithFallbackCookies(t)
sessionData, err := h.GetSessionFallback(r)
assert.NoError(t, err)
assert.Equal(t, data.ExternalSessionID, sessionData.ExternalSessionID)
assert.Equal(t, data.AccessToken, sessionData.AccessToken)
assert.Equal(t, data.IDToken, sessionData.IDToken)
assert.Equal(t, "sid", sessionData.ExternalSessionID)
assert.Equal(t, "access_token", sessionData.AccessToken)
assert.Equal(t, "id_token", sessionData.IDToken)
})
}
@@ -90,17 +69,54 @@ func TestHandler_SetSessionFallback(t *testing.T) {
func TestHandler_DeleteSessionFallback(t *testing.T) {
h := newHandler(mock.NewTestProvider())
t.Run("expire cookies if they are set", func(t *testing.T) {
r := makeRequestWithFallbackCookies(t)
writer := httptest.NewRecorder()
h.DeleteSessionFallback(writer, r)
cookies := writer.Result().Cookies()
assert.NotEmpty(t, cookies)
assert.Len(t, cookies, 3)
assertCookieExpired(t, h.SessionFallbackExternalIDCookieName(), cookies)
assertCookieExpired(t, h.SessionFallbackIDTokenCookieName(), cookies)
assertCookieExpired(t, h.SessionFallbackAccessTokenCookieName(), cookies)
})
t.Run("skip expiring cookies if they are not set", func(t *testing.T) {
writer := httptest.NewRecorder()
r := httptest.NewRequest(http.MethodGet, "/", nil)
h.DeleteSessionFallback(writer, r)
cookies := writer.Result().Cookies()
assert.Empty(t, cookies)
})
}
func makeRequestWithFallbackCookies(t *testing.T) *http.Request {
h := newHandler(mock.NewTestProvider())
writer := httptest.NewRecorder()
h.DeleteSessionFallback(writer)
expiresIn := time.Minute
data := session.NewData("sid", "access_token", "id_token")
err := h.SetSessionFallback(writer, data, expiresIn)
assert.NoError(t, err)
cookies := writer.Result().Cookies()
assert.NotEmpty(t, cookies)
assert.Len(t, cookies, 3)
externalSessionIDCookie := getCookieFromJar(h.SessionFallbackExternalIDCookieName(), cookies)
assert.NotNil(t, externalSessionIDCookie)
idTokenCookie := getCookieFromJar(h.SessionFallbackIDTokenCookieName(), cookies)
assert.NotNil(t, idTokenCookie)
accessTokenCookie := getCookieFromJar(h.SessionFallbackAccessTokenCookieName(), cookies)
assert.NotNil(t, accessTokenCookie)
assertCookieExpired(t, h.SessionFallbackExternalIDCookieName(), cookies)
assertCookieExpired(t, h.SessionFallbackIDTokenCookieName(), cookies)
assertCookieExpired(t, h.SessionFallbackAccessTokenCookieName(), cookies)
// make request with fallback session cookies set
r := httptest.NewRequest(http.MethodGet, "/", nil)
r.AddCookie(externalSessionIDCookie)
r.AddCookie(idTokenCookie)
r.AddCookie(accessTokenCookie)
return r
}
func assertCookieExpired(t *testing.T, cookieName string, cookies []*http.Cookie) {