mirror of
https://github.com/nais/wonderwall.git
synced 2026-08-23 21:16:14 +00:00
refactor: clean up errors and reverseproxy logging
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -13,7 +13,6 @@ import (
|
||||
|
||||
var (
|
||||
ErrKeyNotFound = errors.New("key not found")
|
||||
ErrUnexpected = errors.New("unexpected error")
|
||||
)
|
||||
|
||||
type Store interface {
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user