refactor: clean up errors and reverseproxy logging

This commit is contained in:
Trong Huu Nguyen
2023-02-10 14:57:53 +01:00
parent ce177fb4a5
commit 61a7a8f161
10 changed files with 58 additions and 51 deletions
+6 -1
View File
@@ -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)
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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)
+1 -1
View File
@@ -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
+11 -6
View File
@@ -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) {
+9 -4
View File
@@ -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)
+3 -3
View File
@@ -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
+22 -29
View File
@@ -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)
}
-1
View File
@@ -13,7 +13,6 @@ import (
var (
ErrKeyNotFound = errors.New("key not found")
ErrUnexpected = errors.New("unexpected error")
)
type Store interface {
+4 -4
View File
@@ -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