From b85ea7136e735c96d2a187d60a4bddfc9c7269d6 Mon Sep 17 00:00:00 2001 From: Trong Huu Nguyen Date: Thu, 21 Oct 2021 09:56:20 +0200 Subject: [PATCH] refactor: only delete fallback session cookies if set --- pkg/router/session.go | 12 +++-- pkg/router/session_fallback.go | 24 ++++++--- pkg/router/session_fallback_test.go | 78 +++++++++++++++++------------ 3 files changed, 71 insertions(+), 43 deletions(-) diff --git a/pkg/router/session.go b/pkg/router/session.go index 2a326d0..34d7c37 100644 --- a/pkg/router/session.go +++ b/pkg/router/session.go @@ -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 } diff --git a/pkg/router/session_fallback.go b/pkg/router/session_fallback.go index d578501..1a71586 100644 --- a/pkg/router/session_fallback.go +++ b/pkg/router/session_fallback.go @@ -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()) } diff --git a/pkg/router/session_fallback_test.go b/pkg/router/session_fallback_test.go index 5edf45f..c08ef04 100644 --- a/pkg/router/session_fallback_test.go +++ b/pkg/router/session_fallback_test.go @@ -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) {