mirror of
https://github.com/nais/wonderwall.git
synced 2026-08-23 21:16:14 +00:00
refactor: only delete fallback session cookies if set
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user