From 61a7a8f1612aed05888594bb701e1e4df80b3209 Mon Sep 17 00:00:00 2001 From: Trong Huu Nguyen Date: Mon, 9 Jan 2023 16:04:22 +0100 Subject: [PATCH] refactor: clean up errors and reverseproxy logging --- pkg/cookie/cookie.go | 7 +++- pkg/handler/logout.go | 2 +- pkg/handler/logout_frontchannel.go | 2 +- pkg/handler/logout_frontchannel_test.go | 2 +- pkg/handler/reverseproxy.go | 17 ++++++--- pkg/handler/session.go | 13 +++++-- pkg/handler/session_refresh.go | 6 +-- pkg/session/handler.go | 51 +++++++++++-------------- pkg/session/store.go | 1 - pkg/session/store_redis.go | 8 ++-- 10 files changed, 58 insertions(+), 51 deletions(-) diff --git a/pkg/cookie/cookie.go b/pkg/cookie/cookie.go index e674afe..5a96bd9 100644 --- a/pkg/cookie/cookie.go +++ b/pkg/cookie/cookie.go @@ -2,6 +2,7 @@ package cookie import ( "encoding/base64" + "errors" "fmt" "net/http" "time" @@ -16,6 +17,10 @@ const ( Retry = "io.nais.wonderwall.retry" ) +var ( + ErrInvalidValue = errors.New("invalid value") +) + type Cookie struct { *http.Cookie } @@ -35,7 +40,7 @@ func (in *Cookie) Encrypt(crypter crypto.Crypter) (*Cookie, error) { func (in *Cookie) Decrypt(crypter crypto.Crypter) (string, error) { ciphertext, err := base64.StdEncoding.DecodeString(in.Value) if err != nil { - return "", fmt.Errorf("value for cookie '%s' is not base64 encoded: %w", in.Name, err) + return "", fmt.Errorf("%w: named '%s': %+v", ErrInvalidValue, in.Name, err) } plaintext, err := crypter.Decrypt(ciphertext) diff --git a/pkg/handler/logout.go b/pkg/handler/logout.go index 473c7e0..9fcb63c 100644 --- a/pkg/handler/logout.go +++ b/pkg/handler/logout.go @@ -41,7 +41,7 @@ func Logout(src LogoutSource, w http.ResponseWriter, r *http.Request, opts Logou key, err := sessions.GetKey(r) if err == nil { - sessionData, err := sessions.GetForKey(r, key) + sessionData, err := sessions.Get(r, key) if err == nil && sessionData != nil { idToken = sessionData.IDToken diff --git a/pkg/handler/logout_frontchannel.go b/pkg/handler/logout_frontchannel.go index 3e52fdb..2143a64 100644 --- a/pkg/handler/logout_frontchannel.go +++ b/pkg/handler/logout_frontchannel.go @@ -39,7 +39,7 @@ func LogoutFrontChannel(src LogoutFrontChannelSource, w http.ResponseWriter, r * return } - sessionData, err := sessions.GetForKey(r, key) + sessionData, err := sessions.Get(r, key) if err != nil { logger.Debugf("front-channel logout: could not get session (user might already be logged out): %+v", err) w.WriteHeader(http.StatusAccepted) diff --git a/pkg/handler/logout_frontchannel_test.go b/pkg/handler/logout_frontchannel_test.go index 955ec7f..8b90ebc 100644 --- a/pkg/handler/logout_frontchannel_test.go +++ b/pkg/handler/logout_frontchannel_test.go @@ -28,7 +28,7 @@ func TestFrontChannelLogout(t *testing.T) { sessionKey, err := idp.RelyingPartyHandler.GetCrypter().Decrypt(ciphertext) assert.NoError(t, err) - data, err := idp.RelyingPartyHandler.GetSessions().GetForKey(r, string(sessionKey)) + data, err := idp.RelyingPartyHandler.GetSessions().Get(r, string(sessionKey)) assert.NoError(t, err) return data.ExternalSessionID diff --git a/pkg/handler/reverseproxy.go b/pkg/handler/reverseproxy.go index 20e6517..cf31bfd 100644 --- a/pkg/handler/reverseproxy.go +++ b/pkg/handler/reverseproxy.go @@ -9,6 +9,7 @@ import ( "github.com/sirupsen/logrus" + "github.com/nais/wonderwall/pkg/cookie" "github.com/nais/wonderwall/pkg/handler/autologin" "github.com/nais/wonderwall/pkg/handler/url" "github.com/nais/wonderwall/pkg/loginstatus" @@ -71,16 +72,20 @@ func (rp *ReverseProxy) Handler(src ReverseProxySource, w http.ResponseWriter, r isAuthenticated = false logger.Info("default: loginstatus was enabled, but no matching cookie was found; state is now unauthenticated") } - case errors.Is(err, session.ErrUnexpected): - logger.Errorf("default: unauthenticated: %+v", err) - case errors.Is(err, session.ErrInvalidState): + case errors.Is(err, context.Canceled): + logger.Debugf("default: unauthenticated: %+v (client disconnected before we could respond)", err) + case errors.Is(err, session.ErrInvalidIdpState): logger.Warnf("default: unauthenticated: %+v", err) case errors.Is(err, session.ErrKeyNotFound): - logger.Debug("default: unauthenticated: session not found") + logger.Debug("default: unauthenticated: session not found in store") case errors.Is(err, session.ErrCookieNotFound): - logger.Debug("default: unauthenticated: session cookie not found") - default: + logger.Debug("default: unauthenticated: session cookie not found in request") + case errors.Is(err, session.ErrInvalidSession): logger.Infof("default: unauthenticated: %+v", err) + case errors.Is(err, cookie.ErrInvalidValue): + logger.Debugf("default: unauthenticated: %+v", err) + default: + logger.Errorf("default: unauthenticated: unexpected error: %+v", err) } if src.GetAutoLogin().NeedsLogin(r, isAuthenticated) { diff --git a/pkg/handler/session.go b/pkg/handler/session.go index 4d53038..a8cb916 100644 --- a/pkg/handler/session.go +++ b/pkg/handler/session.go @@ -18,15 +18,20 @@ type SessionSource interface { func Session(src SessionSource, w http.ResponseWriter, r *http.Request) { logger := mw.LogEntryFrom(r) - data, err := src.GetSessions().Get(r) + key, err := src.GetSessions().GetKey(r) + if err != nil { + logger.Infof("session/refresh: getting key: %+v", err) + w.WriteHeader(http.StatusUnauthorized) + return + } + + data, err := src.GetSessions().Get(r, key) if err != nil { switch { - case errors.Is(err, session.ErrCookieNotFound), errors.Is(err, session.ErrKeyNotFound): + case errors.Is(err, session.ErrInvalidSession), errors.Is(err, session.ErrKeyNotFound): logger.Infof("session/info: getting session: %+v", err) w.WriteHeader(http.StatusUnauthorized) return - case errors.Is(err, session.ErrSessionInactive): - // do nothing; we want to return metadata even if the session is inactive default: logger.Warnf("session/info: getting session: %+v", err) w.WriteHeader(http.StatusInternalServerError) diff --git a/pkg/handler/session_refresh.go b/pkg/handler/session_refresh.go index 04f312d..25c4be9 100644 --- a/pkg/handler/session_refresh.go +++ b/pkg/handler/session_refresh.go @@ -23,10 +23,10 @@ func SessionRefresh(src SessionRefreshSource, w http.ResponseWriter, r *http.Req return } - data, err := src.GetSessions().GetForKey(r, key) + data, err := src.GetSessions().Get(r, key) if err != nil { switch { - case errors.Is(err, session.ErrKeyNotFound), errors.Is(err, session.ErrSessionInactive): + case errors.Is(err, session.ErrInvalidSession), errors.Is(err, session.ErrKeyNotFound): logger.Infof("session/refresh: getting session: %+v", err) w.WriteHeader(http.StatusUnauthorized) default: @@ -38,7 +38,7 @@ func SessionRefresh(src SessionRefreshSource, w http.ResponseWriter, r *http.Req data, err = src.GetSessions().Refresh(r, key, data) if err != nil { - if errors.Is(err, session.ErrInvalidState) { + if errors.Is(err, session.ErrInvalidIdpState) || errors.Is(err, session.ErrInvalidSession) { logger.Infof("session/refresh: refreshing: %+v", err) w.WriteHeader(http.StatusUnauthorized) return diff --git a/pkg/session/handler.go b/pkg/session/handler.go index da9fe6c..5353189 100644 --- a/pkg/session/handler.go +++ b/pkg/session/handler.go @@ -22,12 +22,9 @@ import ( ) var ( - ErrCookieNotFound = errors.New("cookie not found") - ErrExpiredAccessToken = errors.New("access token is expired") - ErrInvalidState = errors.New("invalid state") - ErrNoSessionData = errors.New("no session data") - ErrNoAccessToken = errors.New("no access token in session data") - ErrSessionInactive = errors.New("session is inactive") + ErrCookieNotFound = errors.New("session cookie not found") + ErrInvalidSession = errors.New("invalid session") + ErrInvalidIdpState = errors.New("invalid state at idp") ) const ( @@ -113,16 +110,6 @@ func (h *Handler) Destroy(r *http.Request, key string) error { return nil } -// Get returns the session data for a given http.Request, matching by the session cookie. -func (h *Handler) Get(r *http.Request) (*Data, error) { - key, err := h.GetKey(r) - if err != nil { - return nil, fmt.Errorf("no session cookie: %w", err) - } - - return h.GetForKey(r, key) -} - // GetAccessToken returns an access token from the session. If the token is empty or expired, an error is returned. func (h *Handler) GetAccessToken(r *http.Request) (string, error) { sessionData, err := h.GetOrRefresh(r) @@ -131,22 +118,22 @@ func (h *Handler) GetAccessToken(r *http.Request) (string, error) { } if sessionData == nil { - return "", ErrNoSessionData + return "", fmt.Errorf("%w: no session data", ErrInvalidSession) } if !sessionData.HasAccessToken() { - return "", ErrNoAccessToken + return "", fmt.Errorf("%w: no access token in session data", ErrInvalidSession) } if sessionData.Metadata.IsExpired() { - return "", ErrExpiredAccessToken + return "", fmt.Errorf("%w: access token is expired", ErrInvalidSession) } return sessionData.AccessToken, nil } -// GetForKey returns the session data for a given session Key. -func (h *Handler) GetForKey(r *http.Request, key string) (*Data, error) { +// Get returns the session data for a given session Key. +func (h *Handler) Get(r *http.Request, key string) (*Data, error) { var encryptedSessionData *EncryptedData var err error @@ -178,8 +165,14 @@ func (h *Handler) GetForKey(r *http.Request, key string) (*Data, error) { // GetKey extracts the session Key from the session cookie found in the request, if any. func (h *Handler) GetKey(r *http.Request) (string, error) { key, err := cookie.GetDecrypted(r, cookie.Session, h.crypter) + if errors.Is(err, http.ErrNoCookie) { + return "", ErrCookieNotFound + } + if errors.Is(err, cookie.ErrInvalidValue) { + return "", err + } if err != nil { - return "", fmt.Errorf("%w: %+v", ErrCookieNotFound, err) + return "", err } return key, nil @@ -192,13 +185,13 @@ func (h *Handler) GetOrRefresh(r *http.Request) (*Data, error) { return nil, err } - sessionData, err := h.GetForKey(r, key) + sessionData, err := h.Get(r, key) if err != nil { return nil, err } if h.isTimedOut(sessionData) { - return nil, ErrSessionInactive + return nil, fmt.Errorf("%w: session is inactive", ErrInvalidSession) } if !h.shouldRefresh(sessionData) { @@ -206,7 +199,7 @@ func (h *Handler) GetOrRefresh(r *http.Request) (*Data, error) { } refreshed, err := h.Refresh(r, key, sessionData) - if errors.Is(err, ErrInvalidState) || errors.Is(err, ErrSessionInactive) { + if errors.Is(err, ErrInvalidIdpState) || errors.Is(err, ErrInvalidSession) { return nil, err } else if err != nil { mw.LogEntryFrom(r).Warnf("session: could not refresh tokens; falling back to existing token: %+v", err) @@ -270,7 +263,7 @@ func (h *Handler) Refresh(r *http.Request, key string, data *Data) (*Data, error } if !errors.Is(err, ErrAcquireLock) { - return fmt.Errorf("unexpected error: %+v", err) + return err } } } @@ -286,7 +279,7 @@ func (h *Handler) Refresh(r *http.Request, key string, data *Data) (*Data, error }(lock, ctx) // Get the latest session state again in case it was changed while acquiring the lock - data, err = h.GetForKey(r, key) + data, err = h.Get(r, key) if err != nil { return nil, err } @@ -297,7 +290,7 @@ func (h *Handler) Refresh(r *http.Request, key string, data *Data) (*Data, error } if h.isTimedOut(data) { - return nil, ErrSessionInactive + return nil, fmt.Errorf("%w: session is inactive", ErrInvalidSession) } logger.Debug("session: performing refresh grant...") @@ -312,7 +305,7 @@ func (h *Handler) Refresh(r *http.Request, key string, data *Data) (*Data, error } if err := retry.Do(ctx, retrypkg.DefaultBackoff, refresh); err != nil { if errors.Is(err, openidclient.ErrOpenIDClient) { - return nil, fmt.Errorf("%w: authorization might be invalid: %+v", ErrInvalidState, err) + return nil, fmt.Errorf("%w: authorization might be invalid: %+v", ErrInvalidIdpState, err) } return nil, fmt.Errorf("performing refresh: %w", err) } diff --git a/pkg/session/store.go b/pkg/session/store.go index a025abe..5b4d2bc 100644 --- a/pkg/session/store.go +++ b/pkg/session/store.go @@ -13,7 +13,6 @@ import ( var ( ErrKeyNotFound = errors.New("key not found") - ErrUnexpected = errors.New("unexpected error") ) type Store interface { diff --git a/pkg/session/store_redis.go b/pkg/session/store_redis.go index 77df8b0..460b1c0 100644 --- a/pkg/session/store_redis.go +++ b/pkg/session/store_redis.go @@ -36,7 +36,7 @@ func (s *redisSessionStore) Read(ctx context.Context, key string) (*EncryptedDat return nil, fmt.Errorf("%w: %s", ErrKeyNotFound, err.Error()) } - return nil, fmt.Errorf("%w: %s", ErrUnexpected, err.Error()) + return nil, err } func (s *redisSessionStore) Write(ctx context.Context, key string, value *EncryptedData, expiration time.Duration) error { @@ -44,7 +44,7 @@ func (s *redisSessionStore) Write(ctx context.Context, key string, value *Encryp return s.client.Set(ctx, key, value, expiration).Err() }) if err != nil { - return fmt.Errorf("%w: %s", ErrUnexpected, err.Error()) + return err } return nil @@ -62,7 +62,7 @@ func (s *redisSessionStore) Delete(ctx context.Context, keys ...string) error { return fmt.Errorf("%w: %s", ErrKeyNotFound, err.Error()) } - return fmt.Errorf("%w: %s", ErrUnexpected, err.Error()) + return err } func (s *redisSessionStore) Update(ctx context.Context, key string, value *EncryptedData) error { @@ -75,7 +75,7 @@ func (s *redisSessionStore) Update(ctx context.Context, key string, value *Encry return s.client.Set(ctx, key, value, redis.KeepTTL).Err() }) if err != nil { - return fmt.Errorf("%w: %s", ErrUnexpected, err.Error()) + return err } return nil