diff --git a/cmd/wonderwall/main.go b/cmd/wonderwall/main.go index 700c3dc..112b3c5 100644 --- a/cmd/wonderwall/main.go +++ b/cmd/wonderwall/main.go @@ -7,6 +7,7 @@ import ( log "github.com/sirupsen/logrus" "github.com/nais/wonderwall/pkg/config" + "github.com/nais/wonderwall/pkg/cookie" "github.com/nais/wonderwall/pkg/crypto" "github.com/nais/wonderwall/pkg/handler" "github.com/nais/wonderwall/pkg/metrics" @@ -35,7 +36,8 @@ func run() error { defer cancel() crypt := crypto.NewCrypter(key) - h, err := handler.NewHandler(ctx, cfg, openidConfig, crypt) + cookieOpts := cookie.DefaultOptions() + h, err := handler.NewHandler(ctx, cfg, cookieOpts, openidConfig, crypt) if err != nil { return fmt.Errorf("initializing routing handler: %w", err) } diff --git a/pkg/handler/api/login/login.go b/pkg/handler/api/login/login.go new file mode 100644 index 0000000..94e4054 --- /dev/null +++ b/pkg/handler/api/login/login.go @@ -0,0 +1,81 @@ +package login + +import ( + "encoding/json" + "errors" + "fmt" + "net/http" + "time" + + log "github.com/sirupsen/logrus" + + "github.com/nais/wonderwall/pkg/cookie" + "github.com/nais/wonderwall/pkg/crypto" + errorhandler "github.com/nais/wonderwall/pkg/handler/error" + "github.com/nais/wonderwall/pkg/loginstatus" + logentry "github.com/nais/wonderwall/pkg/middleware" + "github.com/nais/wonderwall/pkg/openid" + openidclient "github.com/nais/wonderwall/pkg/openid/client" +) + +const ( + CookieLifetime = 1 * time.Hour +) + +type Source interface { + GetClient() openidclient.Client + GetCookieOptsPathAware(r *http.Request) cookie.Options + GetCrypter() crypto.Crypter + GetErrorHandler() errorhandler.Handler + GetLoginstatus() loginstatus.Loginstatus +} + +func Handler(src Source, w http.ResponseWriter, r *http.Request) { + login, err := src.GetClient().Login(r, src.GetLoginstatus()) + if err != nil { + if errors.Is(err, openidclient.InvalidSecurityLevelError) || errors.Is(err, openidclient.InvalidLocaleError) { + src.GetErrorHandler().BadRequest(w, r, err) + } else { + src.GetErrorHandler().InternalError(w, r, err) + } + + return + } + + err = setLoginCookies(src, w, r, login.Cookie()) + if err != nil { + src.GetErrorHandler().InternalError(w, r, fmt.Errorf("login: setting cookie: %w", err)) + return + } + + fields := log.Fields{ + "redirect_after_login": login.CanonicalRedirect(), + } + logentry.LogEntryFrom(r).WithFields(fields).Debug("login: redirecting to identity provider") + http.Redirect(w, r, login.AuthCodeURL(), http.StatusTemporaryRedirect) +} + +func setLoginCookies(src Source, w http.ResponseWriter, r *http.Request, loginCookie *openid.LoginCookie) error { + loginCookieJson, err := json.Marshal(loginCookie) + if err != nil { + return fmt.Errorf("marshalling login cookie: %w", err) + } + + opts := src.GetCookieOptsPathAware(r). + WithExpiresIn(CookieLifetime). + WithSameSite(http.SameSiteNoneMode) + value := string(loginCookieJson) + + err = cookie.EncryptAndSet(w, cookie.Login, value, opts, src.GetCrypter()) + if err != nil { + return err + } + + // set a duplicate cookie without the SameSite value set for user agents that do not properly handle SameSite + err = cookie.EncryptAndSet(w, cookie.LoginLegacy, value, opts.WithSameSite(http.SameSiteDefaultMode), src.GetCrypter()) + if err != nil { + return err + } + + return nil +} diff --git a/pkg/handler/api/logincallback/logincallback.go b/pkg/handler/api/logincallback/logincallback.go new file mode 100644 index 0000000..c2f8503 --- /dev/null +++ b/pkg/handler/api/logincallback/logincallback.go @@ -0,0 +1,151 @@ +package logincallback + +import ( + "context" + "errors" + "fmt" + "net/http" + + "github.com/sethvargo/go-retry" + log "github.com/sirupsen/logrus" + + "github.com/nais/wonderwall/pkg/config" + "github.com/nais/wonderwall/pkg/cookie" + "github.com/nais/wonderwall/pkg/crypto" + errorhandler "github.com/nais/wonderwall/pkg/handler/error" + "github.com/nais/wonderwall/pkg/loginstatus" + "github.com/nais/wonderwall/pkg/metrics" + logentry "github.com/nais/wonderwall/pkg/middleware" + "github.com/nais/wonderwall/pkg/openid" + openidclient "github.com/nais/wonderwall/pkg/openid/client" + openidprovider "github.com/nais/wonderwall/pkg/openid/provider" + retrypkg "github.com/nais/wonderwall/pkg/retry" + "github.com/nais/wonderwall/pkg/session" +) + +type Source interface { + GetClient() openidclient.Client + GetCookieOptions() cookie.Options + GetCookieOptsPathAware(r *http.Request) cookie.Options + GetCrypter() crypto.Crypter + GetErrorHandler() errorhandler.Handler + GetLoginstatus() loginstatus.Loginstatus + GetProvider() openidprovider.Provider + GetSessions() *session.Handler + GetSessionConfig() config.Session +} + +func Handler(src Source, w http.ResponseWriter, r *http.Request) { + // unconditionally clear login cookie + clearLoginCookies(src, w, r) + + loginCookie, err := openid.GetLoginCookie(r, src.GetCrypter()) + if err != nil { + msg := "callback: fetching login cookie" + if errors.Is(err, http.ErrNoCookie) { + msg += ": fallback cookie not found (user might have blocked all cookies, or the callback route was accessed before the login route)" + } + src.GetErrorHandler().Unauthorized(w, r, fmt.Errorf("%s: %w", msg, err)) + return + } + + loginCallback, err := src.GetClient().LoginCallback(r, src.GetProvider(), loginCookie) + if err != nil { + src.GetErrorHandler().InternalError(w, r, err) + return + } + + if err := loginCallback.IdentityProviderError(); err != nil { + src.GetErrorHandler().InternalError(w, r, fmt.Errorf("callback: %w", err)) + return + } + + if err := loginCallback.StateMismatchError(); err != nil { + src.GetErrorHandler().Unauthorized(w, r, fmt.Errorf("callback: %w", err)) + return + } + + tokens, err := redeemValidTokens(r, loginCallback) + if err != nil { + src.GetErrorHandler().InternalError(w, r, fmt.Errorf("callback: redeeming tokens: %w", err)) + return + } + + sessionLifetime := src.GetSessionConfig().MaxLifetime + + key, err := src.GetSessions().Create(r, tokens, sessionLifetime) + if err != nil { + src.GetErrorHandler().InternalError(w, r, fmt.Errorf("callback: creating session: %w", err)) + return + } + + opts := src.GetCookieOptsPathAware(r). + WithExpiresIn(sessionLifetime) + err = cookie.EncryptAndSet(w, cookie.Session, key, opts, src.GetCrypter()) + if err != nil { + src.GetErrorHandler().InternalError(w, r, fmt.Errorf("callback: setting session cookie: %w", err)) + return + } + + if src.GetLoginstatus().Enabled() { + tokenResponse, err := getLoginstatusToken(src, r, tokens) + if err != nil { + src.GetErrorHandler().InternalError(w, r, fmt.Errorf("callback: exchanging loginstatus token: %w", err)) + return + } + + src.GetLoginstatus().SetCookie(w, tokenResponse, src.GetCookieOptions()) + logentry.LogEntryFrom(r).Debug("callback: successfully fetched loginstatus token") + } + + logSuccessfulLogin(r, tokens, loginCookie.Referer) + http.Redirect(w, r, loginCookie.Referer, http.StatusTemporaryRedirect) +} + +func clearLoginCookies(src Source, w http.ResponseWriter, r *http.Request) { + opts := src.GetCookieOptsPathAware(r) + cookie.Clear(w, cookie.Login, opts.WithSameSite(http.SameSiteNoneMode)) + cookie.Clear(w, cookie.LoginLegacy, opts.WithSameSite(http.SameSiteDefaultMode)) +} + +func redeemValidTokens(r *http.Request, loginCallback openidclient.LoginCallback) (*openid.Tokens, error) { + var tokens *openid.Tokens + var err error + + retryable := func(ctx context.Context) error { + tokens, err = loginCallback.RedeemTokens(ctx) + return retry.RetryableError(err) + } + + if err := retry.Do(r.Context(), retrypkg.DefaultBackoff, retryable); err != nil { + return nil, err + } + + return tokens, nil +} + +func getLoginstatusToken(src Source, r *http.Request, tokens *openid.Tokens) (*loginstatus.TokenResponse, error) { + var tokenResponse *loginstatus.TokenResponse + + retryable := func(ctx context.Context) error { + var err error + + tokenResponse, err = src.GetLoginstatus().ExchangeToken(ctx, tokens.AccessToken) + return retry.RetryableError(err) + } + if err := retry.Do(r.Context(), retrypkg.DefaultBackoff, retryable); err != nil { + return nil, err + } + + return tokenResponse, nil +} + +func logSuccessfulLogin(r *http.Request, tokens *openid.Tokens, referer string) { + fields := log.Fields{ + "redirect_to": referer, + "jti": tokens.IDToken.GetJwtID(), + } + + logentry.LogEntryFrom(r).WithFields(fields).Info("callback: successful login") + metrics.ObserveLogin() +} diff --git a/pkg/handler/api/logout/logout.go b/pkg/handler/api/logout/logout.go new file mode 100644 index 0000000..4292120 --- /dev/null +++ b/pkg/handler/api/logout/logout.go @@ -0,0 +1,63 @@ +package logout + +import ( + "errors" + "fmt" + "net/http" + + log "github.com/sirupsen/logrus" + + "github.com/nais/wonderwall/pkg/cookie" + errorhandler "github.com/nais/wonderwall/pkg/handler/error" + "github.com/nais/wonderwall/pkg/loginstatus" + "github.com/nais/wonderwall/pkg/metrics" + logentry "github.com/nais/wonderwall/pkg/middleware" + openidclient "github.com/nais/wonderwall/pkg/openid/client" + "github.com/nais/wonderwall/pkg/session" +) + +type Source interface { + GetClient() openidclient.Client + GetCookieOptions() cookie.Options + GetCookieOptsPathAware(r *http.Request) cookie.Options + GetErrorHandler() errorhandler.Handler + GetLoginstatus() loginstatus.Loginstatus + GetSessions() *session.Handler +} + +func Handler(src Source, w http.ResponseWriter, r *http.Request) { + logger := logentry.LogEntryFrom(r) + logout, err := src.GetClient().Logout(r) + if err != nil { + src.GetErrorHandler().InternalError(w, r, err) + return + } + + idToken := "" + + sessionData, err := src.GetSessions().Get(r) + if err == nil && sessionData != nil { + idToken = sessionData.IDToken + + err = src.GetSessions().DestroyForID(r, sessionData.ExternalSessionID) + if err != nil && !errors.Is(err, session.KeyNotFoundError) { + src.GetErrorHandler().InternalError(w, r, fmt.Errorf("logout: destroying session: %w", err)) + return + } + + fields := log.Fields{ + "jti": sessionData.IDTokenJwtID, + } + logger.WithFields(fields).Info("logout: successful local logout") + } + + cookie.Clear(w, cookie.Session, src.GetCookieOptsPathAware(r)) + + if src.GetLoginstatus().Enabled() { + src.GetLoginstatus().ClearCookie(w, src.GetCookieOptions()) + } + + logger.Debug("logout: redirecting to identity provider") + metrics.ObserveLogout(metrics.LogoutOperationSelfInitiated) + http.Redirect(w, r, logout.SingleLogoutURL(idToken), http.StatusTemporaryRedirect) +} diff --git a/pkg/handler/api/logoutcallback/logoutcallback.go b/pkg/handler/api/logoutcallback/logoutcallback.go new file mode 100644 index 0000000..d83d0a1 --- /dev/null +++ b/pkg/handler/api/logoutcallback/logoutcallback.go @@ -0,0 +1,19 @@ +package logoutcallback + +import ( + "net/http" + + logentry "github.com/nais/wonderwall/pkg/middleware" + openidclient "github.com/nais/wonderwall/pkg/openid/client" +) + +type Source interface { + GetClient() openidclient.Client +} + +func Handler(src Source, w http.ResponseWriter, r *http.Request) { + redirect := src.GetClient().LogoutCallback(r).PostLogoutRedirectURI() + + logentry.LogEntryFrom(r).Debugf("logout/callback: redirecting to %s", redirect) + http.Redirect(w, r, redirect, http.StatusTemporaryRedirect) +} diff --git a/pkg/handler/handler_frontchannellogout.go b/pkg/handler/api/logoutfrontchannel/logoutfrontchannel.go similarity index 55% rename from pkg/handler/handler_frontchannellogout.go rename to pkg/handler/api/logoutfrontchannel/logoutfrontchannel.go index a259c9d..8398d0e 100644 --- a/pkg/handler/handler_frontchannellogout.go +++ b/pkg/handler/api/logoutfrontchannel/logoutfrontchannel.go @@ -1,25 +1,35 @@ -package handler +package logoutfrontchannel import ( "net/http" "github.com/nais/wonderwall/pkg/cookie" + "github.com/nais/wonderwall/pkg/loginstatus" "github.com/nais/wonderwall/pkg/metrics" mw "github.com/nais/wonderwall/pkg/middleware" + openidclient "github.com/nais/wonderwall/pkg/openid/client" + "github.com/nais/wonderwall/pkg/session" ) -// FrontChannelLogout performs a local logout initiated by a third party in the SSO circle-of-trust. -func (h *Handler) FrontChannelLogout(w http.ResponseWriter, r *http.Request) { +type Source interface { + GetClient() openidclient.Client + GetCookieOptions() cookie.Options + GetCookieOptsPathAware(r *http.Request) cookie.Options + GetLoginstatus() loginstatus.Loginstatus + GetSessions() *session.Handler +} + +func Handler(src Source, w http.ResponseWriter, r *http.Request) { logger := mw.LogEntryFrom(r) // Unconditionally destroy all local references to the session. - cookie.Clear(w, cookie.Session, h.CookieOptsPathAware(r)) + cookie.Clear(w, cookie.Session, src.GetCookieOptsPathAware(r)) - if h.Loginstatus.Enabled() { - h.Loginstatus.ClearCookie(w, h.CookieOptions) + if src.GetLoginstatus().Enabled() { + src.GetLoginstatus().ClearCookie(w, src.GetCookieOptions()) } - logoutFrontchannel := h.Client.LogoutFrontchannel(r) + logoutFrontchannel := src.GetClient().LogoutFrontchannel(r) if logoutFrontchannel.MissingSidParameter() { logger.Debug("front-channel logout: sid parameter not set in request; ignoring") w.WriteHeader(http.StatusAccepted) @@ -27,14 +37,14 @@ func (h *Handler) FrontChannelLogout(w http.ResponseWriter, r *http.Request) { } sid := logoutFrontchannel.Sid() - sessionData, err := h.Sessions.GetForID(r, sid) + sessionData, err := src.GetSessions().GetForID(r, sid) if err != nil { logger.Debugf("front-channel logout: could not get session (user might already be logged out): %+v", err) w.WriteHeader(http.StatusAccepted) return } - err = h.Sessions.DestroyForID(r, sid) + err = src.GetSessions().DestroyForID(r, sid) if err != nil { logger.Warnf("front-channel logout: destroying session: %+v", err) w.WriteHeader(http.StatusAccepted) diff --git a/pkg/handler/handler_session_info.go b/pkg/handler/api/session/session.go similarity index 74% rename from pkg/handler/handler_session_info.go rename to pkg/handler/api/session/session.go index 48acbc2..3242768 100644 --- a/pkg/handler/handler_session_info.go +++ b/pkg/handler/api/session/session.go @@ -1,19 +1,24 @@ -package handler +package session import ( "encoding/json" "errors" "net/http" + "github.com/nais/wonderwall/pkg/config" mw "github.com/nais/wonderwall/pkg/middleware" "github.com/nais/wonderwall/pkg/session" ) -// SessionInfo returns metadata for the current user's session. -func (h *Handler) SessionInfo(w http.ResponseWriter, r *http.Request) { +type Source interface { + GetSessions() *session.Handler + GetSessionConfig() config.Session +} + +func Handler(src Source, w http.ResponseWriter, r *http.Request) { logger := mw.LogEntryFrom(r) - data, err := h.Sessions.Get(r) + data, err := src.GetSessions().Get(r) if err != nil { if errors.Is(err, session.CookieNotFoundError) || errors.Is(err, session.KeyNotFoundError) { logger.Infof("session/info: getting session: %+v", err) @@ -28,7 +33,7 @@ func (h *Handler) SessionInfo(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") - if h.Config.Session.Refresh { + if src.GetSessionConfig().Refresh { err = json.NewEncoder(w).Encode(data.Metadata.VerboseWithRefresh()) } else { err = json.NewEncoder(w).Encode(data.Metadata.Verbose()) diff --git a/pkg/handler/handler_session_refresh.go b/pkg/handler/api/sessionrefresh/sessionrefresh.go similarity index 73% rename from pkg/handler/handler_session_refresh.go rename to pkg/handler/api/sessionrefresh/sessionrefresh.go index ba733a9..efc1295 100644 --- a/pkg/handler/handler_session_refresh.go +++ b/pkg/handler/api/sessionrefresh/sessionrefresh.go @@ -1,4 +1,4 @@ -package handler +package sessionrefresh import ( "encoding/json" @@ -9,23 +9,21 @@ import ( "github.com/nais/wonderwall/pkg/session" ) -// SessionRefresh refreshes current user's session and returns the associated updated metadata. -func (h *Handler) SessionRefresh(w http.ResponseWriter, r *http.Request) { - if !h.Config.Session.Refresh { - http.NotFound(w, r) - return - } +type Source interface { + GetSessions() *session.Handler +} +func Handler(src Source, w http.ResponseWriter, r *http.Request) { logger := mw.LogEntryFrom(r) - key, err := h.Sessions.GetKey(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 := h.Sessions.Get(r) + data, err := src.GetSessions().Get(r) if err != nil { if errors.Is(err, session.KeyNotFoundError) { logger.Infof("session/refresh: getting session: %+v", err) @@ -38,7 +36,7 @@ func (h *Handler) SessionRefresh(w http.ResponseWriter, r *http.Request) { return } - data, err = h.Sessions.Refresh(r, key, data) + data, err = src.GetSessions().Refresh(r, key, data) if err != nil { logger.Warnf("session/refresh: refreshing: %+v", err) w.WriteHeader(http.StatusInternalServerError) diff --git a/pkg/handler/handler_error.go b/pkg/handler/error/error.go similarity index 58% rename from pkg/handler/handler_error.go rename to pkg/handler/error/error.go index 469d63f..1a53c62 100644 --- a/pkg/handler/handler_error.go +++ b/pkg/handler/error/error.go @@ -1,8 +1,6 @@ -package handler +package error import ( - _ "embed" - "html/template" "net/http" "net/url" "strconv" @@ -11,32 +9,74 @@ import ( "github.com/go-chi/chi/v5/middleware" log "github.com/sirupsen/logrus" + "github.com/nais/wonderwall/pkg/crypto" + "github.com/nais/wonderwall/pkg/handler/templates" urlpkg "github.com/nais/wonderwall/pkg/handler/url" mw "github.com/nais/wonderwall/pkg/middleware" "github.com/nais/wonderwall/pkg/openid" "github.com/nais/wonderwall/pkg/router/paths" ) -type ErrorPage struct { +type Handler interface { + InternalError(w http.ResponseWriter, r *http.Request, cause error) + BadRequest(w http.ResponseWriter, r *http.Request, cause error) + Unauthorized(w http.ResponseWriter, r *http.Request, cause error) + Retry(r *http.Request, loginCookie *openid.LoginCookie) string +} + +type Source interface { + GetCrypter() crypto.Crypter + GetErrorRedirectURI() string + GetPath(r *http.Request) string +} + +type Page struct { CorrelationID string RetryURI string } -//go:embed templates/error.gohtml -var errorGoHtml string -var errorTemplate *template.Template - -func init() { - var err error - - errorTemplate = template.New("error") - errorTemplate, err = errorTemplate.Parse(errorGoHtml) - if err != nil { - log.Fatalf("parsing error template: %+v", err) - } +type handler struct { + Source } -func (h *Handler) respondError(w http.ResponseWriter, r *http.Request, statusCode int, cause error, level log.Level) { +func New(src Source) Handler { + return handler{src} +} + +func (h handler) InternalError(w http.ResponseWriter, r *http.Request, cause error) { + h.respondError(w, r, http.StatusInternalServerError, cause, log.ErrorLevel) +} + +func (h handler) BadRequest(w http.ResponseWriter, r *http.Request, cause error) { + h.respondError(w, r, http.StatusBadRequest, cause, log.ErrorLevel) +} + +func (h handler) Unauthorized(w http.ResponseWriter, r *http.Request, cause error) { + h.respondError(w, r, http.StatusUnauthorized, cause, log.WarnLevel) +} + +// Retry returns a URI that should retry the desired route that failed. +// It only handles the routes exposed by Wonderwall, i.e. `/oauth2/*`. As these routes +// are related to the authentication flow, we default to redirecting back to the handled +// `/oauth2/login` endpoint unless the original request attempted to reach the logout-flow. +func (h handler) Retry(r *http.Request, loginCookie *openid.LoginCookie) string { + requestPath := r.URL.Path + ingressPath := h.GetPath(r) + + if strings.HasSuffix(requestPath, paths.OAuth2+paths.Logout) || strings.HasSuffix(requestPath, paths.OAuth2+paths.LogoutFrontChannel) { + return requestPath + } + + redirect := urlpkg.CanonicalRedirect(r) + + if loginCookie != nil && len(loginCookie.Referer) > 0 { + redirect = loginCookie.Referer + } + + return urlpkg.LoginURL(ingressPath, redirect) +} + +func (h handler) respondError(w http.ResponseWriter, r *http.Request, statusCode int, cause error, level log.Level) { logger := mw.LogEntryFrom(r) msg := "error in route: %+v" @@ -47,7 +87,7 @@ func (h *Handler) respondError(w http.ResponseWriter, r *http.Request, statusCod logger.Errorf(msg, cause) } - if len(h.Config.ErrorRedirectURI) > 0 { + if len(h.GetErrorRedirectURI()) > 0 { err := h.customErrorRedirect(w, r, statusCode) if err == nil { return @@ -57,26 +97,26 @@ func (h *Handler) respondError(w http.ResponseWriter, r *http.Request, statusCod h.defaultErrorResponse(w, r, statusCode) } -func (h *Handler) defaultErrorResponse(w http.ResponseWriter, r *http.Request, statusCode int) { +func (h handler) defaultErrorResponse(w http.ResponseWriter, r *http.Request, statusCode int) { w.WriteHeader(statusCode) - loginCookie, err := h.getLoginCookie(r) + loginCookie, err := openid.GetLoginCookie(r, h.GetCrypter()) if err != nil { loginCookie = nil } - errorPage := ErrorPage{ + errorPage := Page{ CorrelationID: middleware.GetReqID(r.Context()), RetryURI: h.Retry(r, loginCookie), } - err = errorTemplate.Execute(w, errorPage) + err = templates.ErrorTemplate.Execute(w, errorPage) if err != nil { mw.LogEntryFrom(r).Errorf("executing error template: %+v", err) } } -func (h *Handler) customErrorRedirect(w http.ResponseWriter, r *http.Request, statusCode int) error { - override, err := url.Parse(h.Config.ErrorRedirectURI) +func (h handler) customErrorRedirect(w http.ResponseWriter, r *http.Request, statusCode int) error { + override, err := url.Parse(h.GetErrorRedirectURI()) if err != nil { return err } @@ -94,36 +134,3 @@ func (h *Handler) customErrorRedirect(w http.ResponseWriter, r *http.Request, st http.Redirect(w, r, errorRedirectURI, http.StatusFound) return nil } - -func (h *Handler) InternalError(w http.ResponseWriter, r *http.Request, cause error) { - h.respondError(w, r, http.StatusInternalServerError, cause, log.ErrorLevel) -} - -func (h *Handler) BadRequest(w http.ResponseWriter, r *http.Request, cause error) { - h.respondError(w, r, http.StatusBadRequest, cause, log.ErrorLevel) -} - -func (h *Handler) Unauthorized(w http.ResponseWriter, r *http.Request, cause error) { - h.respondError(w, r, http.StatusUnauthorized, cause, log.WarnLevel) -} - -// Retry returns a URI that should retry the desired route that failed. -// It only handles the routes exposed by Wonderwall, i.e. `/oauth2/*`. As these routes -// are related to the authentication flow, we default to redirecting back to the handled -// `/oauth2/login` endpoint unless the original request attempted to reach the logout-flow. -func (h *Handler) Retry(r *http.Request, loginCookie *openid.LoginCookie) string { - requestPath := r.URL.Path - ingressPath := h.Path(r) - - if strings.HasSuffix(requestPath, paths.OAuth2+paths.Logout) || strings.HasSuffix(requestPath, paths.OAuth2+paths.FrontChannelLogout) { - return requestPath - } - - redirect := urlpkg.CanonicalRedirect(r) - - if loginCookie != nil && len(loginCookie.Referer) > 0 { - redirect = loginCookie.Referer - } - - return urlpkg.LoginURL(ingressPath, redirect) -} diff --git a/pkg/handler/handler_error_test.go b/pkg/handler/error/error_test.go similarity index 98% rename from pkg/handler/handler_error_test.go rename to pkg/handler/error/error_test.go index fbb9284..19c966f 100644 --- a/pkg/handler/handler_error_test.go +++ b/pkg/handler/error/error_test.go @@ -1,4 +1,4 @@ -package handler_test +package error_test import ( "fmt" @@ -19,7 +19,7 @@ func TestHandler_Error(t *testing.T) { idp := mock.NewIdentityProvider(cfg) defer idp.Close() - rpHandler := idp.RelyingPartyHandler + rpHandler := idp.RelyingPartyHandler.GetErrorHandler() for _, test := range []struct { name string @@ -57,7 +57,7 @@ func TestHandler_Retry(t *testing.T) { idp := mock.NewIdentityProvider(cfg) defer idp.Close() - handler := idp.RelyingPartyHandler + handler := idp.RelyingPartyHandler.GetErrorHandler() httpRequest := func(url string, referer ...string) *http.Request { req := httptest.NewRequest(http.MethodGet, url, nil) diff --git a/pkg/handler/handler.go b/pkg/handler/handler.go index 7ae390f..0beec8c 100644 --- a/pkg/handler/handler.go +++ b/pkg/handler/handler.go @@ -3,41 +3,27 @@ package handler import ( "context" "net/http" - "net/http/httputil" "time" "github.com/nais/wonderwall/pkg/config" "github.com/nais/wonderwall/pkg/cookie" "github.com/nais/wonderwall/pkg/crypto" "github.com/nais/wonderwall/pkg/handler/autologin" - "github.com/nais/wonderwall/pkg/ingress" + "github.com/nais/wonderwall/pkg/handler/reverseproxy" "github.com/nais/wonderwall/pkg/loginstatus" - "github.com/nais/wonderwall/pkg/middleware" "github.com/nais/wonderwall/pkg/openid/client" openidconfig "github.com/nais/wonderwall/pkg/openid/config" "github.com/nais/wonderwall/pkg/openid/provider" "github.com/nais/wonderwall/pkg/session" ) -type Handler struct { - AutoLogin *autologin.AutoLogin - Client client.Client - Config *config.Config - CookieOptions cookie.Options - Crypter crypto.Crypter - Loginstatus loginstatus.Loginstatus - OpenIDConfig openidconfig.Config - Provider provider.Provider - ReverseProxy *httputil.ReverseProxy - Sessions *session.Handler -} - func NewHandler( ctx context.Context, cfg *config.Config, + cookieOpts cookie.Options, openidConfig openidconfig.Config, crypter crypto.Crypter, -) (*Handler, error) { +) (*StandardHandler, error) { openidProvider, err := provider.NewProvider(ctx, openidConfig) if err != nil { return nil, err @@ -60,60 +46,16 @@ func NewHandler( return nil, err } - return &Handler{ - AutoLogin: autoLogin, - Client: openidClient, - Config: cfg, - CookieOptions: cookie.DefaultOptions(), - Crypter: crypter, - Loginstatus: loginstatus.NewClient(cfg.Loginstatus, httpClient), - OpenIDConfig: openidConfig, - Provider: openidProvider, - ReverseProxy: newReverseProxy(cfg.UpstreamHost), - Sessions: sessionHandler, + return &StandardHandler{ + autoLogin: autoLogin, + client: openidClient, + config: cfg, + cookieOptions: cookieOpts, + crypter: crypter, + loginstatus: loginstatus.NewClient(cfg.Loginstatus, httpClient), + openidConfig: openidConfig, + provider: openidProvider, + sessions: sessionHandler, + upstreamProxy: reverseproxy.New(cfg.UpstreamHost), }, nil } - -func (h *Handler) CookieOptsPathAware(r *http.Request) cookie.Options { - path := h.Path(r) - return h.CookieOptions.WithPath(path) -} - -func (h *Handler) Ingresses() *ingress.Ingresses { - return h.OpenIDConfig.Client().Ingresses() -} - -func (h *Handler) Path(r *http.Request) string { - path, ok := middleware.PathFrom(r.Context()) - if !ok { - path = h.OpenIDConfig.Client().Ingresses().MatchingPath(r) - } - - return path -} - -func (h *Handler) ProviderName() string { - return h.OpenIDConfig.Provider().Name() -} - -func newReverseProxy(upstreamHost string) *httputil.ReverseProxy { - return &httputil.ReverseProxy{ - Director: func(r *http.Request) { - // Delete incoming authentication - r.Header.Del("authorization") - // Instruct http.ReverseProxy to not modify X-Forwarded-For header - r.Header["X-Forwarded-For"] = nil - // Request should go to correct host - r.URL.Host = upstreamHost - r.URL.Scheme = "http" - - accessToken, ok := middleware.AccessTokenFrom(r.Context()) - if ok { - r.Header.Set("authorization", "Bearer "+accessToken) - } - }, - ErrorHandler: func(w http.ResponseWriter, r *http.Request, err error) { - http.Error(w, err.Error(), http.StatusBadGateway) - }, - } -} diff --git a/pkg/handler/handler_callback.go b/pkg/handler/handler_callback.go deleted file mode 100644 index 43fc744..0000000 --- a/pkg/handler/handler_callback.go +++ /dev/null @@ -1,129 +0,0 @@ -package handler - -import ( - "context" - "errors" - "fmt" - "net/http" - - "github.com/sethvargo/go-retry" - log "github.com/sirupsen/logrus" - - "github.com/nais/wonderwall/pkg/cookie" - "github.com/nais/wonderwall/pkg/loginstatus" - "github.com/nais/wonderwall/pkg/metrics" - logentry "github.com/nais/wonderwall/pkg/middleware" - "github.com/nais/wonderwall/pkg/openid" - "github.com/nais/wonderwall/pkg/openid/client" - retrypkg "github.com/nais/wonderwall/pkg/retry" -) - -// Callback handles the authentication response from the identity provider. -func (h *Handler) Callback(w http.ResponseWriter, r *http.Request) { - // unconditionally clear login cookie - h.clearLoginCookies(w, r) - - loginCookie, err := h.getLoginCookie(r) - if err != nil { - msg := "callback: fetching login cookie" - if errors.Is(err, http.ErrNoCookie) { - msg += ": fallback cookie not found (user might have blocked all cookies, or the callback route was accessed before the login route)" - } - h.Unauthorized(w, r, fmt.Errorf("%s: %w", msg, err)) - return - } - - loginCallback, err := h.Client.LoginCallback(r, h.Provider, loginCookie) - if err != nil { - h.InternalError(w, r, err) - return - } - - if err := loginCallback.IdentityProviderError(); err != nil { - h.InternalError(w, r, fmt.Errorf("callback: %w", err)) - return - } - - if err := loginCallback.StateMismatchError(); err != nil { - h.Unauthorized(w, r, fmt.Errorf("callback: %w", err)) - return - } - - tokens, err := h.redeemValidTokens(r, loginCallback) - if err != nil { - h.InternalError(w, r, fmt.Errorf("callback: redeeming tokens: %w", err)) - return - } - - sessionLifetime := h.Config.Session.MaxLifetime - - key, err := h.Sessions.Create(r, tokens, sessionLifetime) - if err != nil { - h.InternalError(w, r, fmt.Errorf("callback: creating session: %w", err)) - return - } - - opts := h.CookieOptsPathAware(r). - WithExpiresIn(sessionLifetime) - err = cookie.EncryptAndSet(w, cookie.Session, key, opts, h.Crypter) - if err != nil { - h.InternalError(w, r, fmt.Errorf("callback: setting session cookie: %w", err)) - return - } - - if h.Loginstatus.Enabled() { - tokenResponse, err := h.getLoginstatusToken(r, tokens) - if err != nil { - h.InternalError(w, r, fmt.Errorf("callback: exchanging loginstatus token: %w", err)) - return - } - - h.Loginstatus.SetCookie(w, tokenResponse, h.CookieOptions) - logentry.LogEntryFrom(r).Debug("callback: successfully fetched loginstatus token") - } - - logSuccessfulLogin(r, tokens, loginCookie.Referer) - http.Redirect(w, r, loginCookie.Referer, http.StatusTemporaryRedirect) -} - -func (h *Handler) redeemValidTokens(r *http.Request, loginCallback client.LoginCallback) (*openid.Tokens, error) { - var tokens *openid.Tokens - var err error - - retryable := func(ctx context.Context) error { - tokens, err = loginCallback.RedeemTokens(ctx) - return retry.RetryableError(err) - } - - if err := retry.Do(r.Context(), retrypkg.DefaultBackoff, retryable); err != nil { - return nil, err - } - - return tokens, nil -} - -func (h *Handler) getLoginstatusToken(r *http.Request, tokens *openid.Tokens) (*loginstatus.TokenResponse, error) { - var tokenResponse *loginstatus.TokenResponse - - retryable := func(ctx context.Context) error { - var err error - - tokenResponse, err = h.Loginstatus.ExchangeToken(ctx, tokens.AccessToken) - return retry.RetryableError(err) - } - if err := retry.Do(r.Context(), retrypkg.DefaultBackoff, retryable); err != nil { - return nil, err - } - - return tokenResponse, nil -} - -func logSuccessfulLogin(r *http.Request, tokens *openid.Tokens, referer string) { - fields := log.Fields{ - "redirect_to": referer, - "jti": tokens.IDToken.GetJwtID(), - } - - logentry.LogEntryFrom(r).WithFields(fields).Info("callback: successful login") - metrics.ObserveLogin() -} diff --git a/pkg/handler/handler_default.go b/pkg/handler/handler_default.go deleted file mode 100644 index 78b82a5..0000000 --- a/pkg/handler/handler_default.go +++ /dev/null @@ -1,49 +0,0 @@ -package handler - -import ( - "errors" - "net/http" - - "github.com/nais/wonderwall/pkg/handler/url" - mw "github.com/nais/wonderwall/pkg/middleware" - "github.com/nais/wonderwall/pkg/session" -) - -// Default proxies all requests upstream. -func (h *Handler) Default(w http.ResponseWriter, r *http.Request) { - logger := mw.LogEntryFrom(r).WithField("request_path", r.URL.Path) - isAuthenticated := false - - accessToken, err := h.Sessions.GetAccessToken(r) - if err == nil { - // add authentication if session cookie and token checks out - isAuthenticated = true - - // force new authentication if loginstatus is enabled and cookie isn't set - if h.Loginstatus.NeedsLogin(r) { - isAuthenticated = false - logger.Info("default: loginstatus was enabled, but no matching cookie was found; state is now unauthenticated") - } - } else if errors.Is(err, session.UnexpectedError) { - logger.Errorf("default: getting session: %+v", err) - } - - if h.AutoLogin.NeedsLogin(r, isAuthenticated) { - logger.Debug("default: auto-login is enabled; request does not match any configured ignorable paths") - - redirectTarget := r.URL.String() - path := h.Path(r) - - loginUrl := url.LoginURL(path, redirectTarget) - http.Redirect(w, r, loginUrl, http.StatusTemporaryRedirect) - return - } - - ctx := r.Context() - - if isAuthenticated { - ctx = mw.WithAccessToken(ctx, accessToken) - } - - h.ReverseProxy.ServeHTTP(w, r.WithContext(ctx)) -} diff --git a/pkg/handler/handler_login.go b/pkg/handler/handler_login.go deleted file mode 100644 index fd6f53a..0000000 --- a/pkg/handler/handler_login.go +++ /dev/null @@ -1,97 +0,0 @@ -package handler - -import ( - "encoding/json" - "errors" - "fmt" - "net/http" - "time" - - log "github.com/sirupsen/logrus" - - "github.com/nais/wonderwall/pkg/cookie" - logentry "github.com/nais/wonderwall/pkg/middleware" - "github.com/nais/wonderwall/pkg/openid" - "github.com/nais/wonderwall/pkg/openid/client" -) - -const ( - LoginCookieLifetime = 1 * time.Hour -) - -// Login initiates the authorization code flow. -func (h *Handler) Login(w http.ResponseWriter, r *http.Request) { - login, err := h.Client.Login(r, h.Loginstatus) - if err != nil { - if errors.Is(err, client.InvalidSecurityLevelError) || errors.Is(err, client.InvalidLocaleError) { - h.BadRequest(w, r, err) - } else { - h.InternalError(w, r, err) - } - - return - } - - err = h.setLoginCookies(w, r, login.Cookie()) - if err != nil { - h.InternalError(w, r, fmt.Errorf("login: setting cookie: %w", err)) - return - } - - fields := log.Fields{ - "redirect_after_login": login.CanonicalRedirect(), - } - logentry.LogEntryFrom(r).WithFields(fields).Debug("login: redirecting to identity provider") - http.Redirect(w, r, login.AuthCodeURL(), http.StatusTemporaryRedirect) -} - -func (h *Handler) getLoginCookie(r *http.Request) (*openid.LoginCookie, error) { - loginCookieJson, err := cookie.GetDecrypted(r, cookie.Login, h.Crypter) - if err != nil { - logentry.LogEntryFrom(r).Debugf("failed to fetch login cookie: %+v; falling back to legacy cookie", err) - - loginCookieJson, err = cookie.GetDecrypted(r, cookie.LoginLegacy, h.Crypter) - if err != nil { - return nil, err - } - } - - var loginCookie openid.LoginCookie - err = json.Unmarshal([]byte(loginCookieJson), &loginCookie) - if err != nil { - return nil, fmt.Errorf("unmarshalling: %w", err) - } - - return &loginCookie, nil -} - -func (h *Handler) setLoginCookies(w http.ResponseWriter, r *http.Request, loginCookie *openid.LoginCookie) error { - loginCookieJson, err := json.Marshal(loginCookie) - if err != nil { - return fmt.Errorf("marshalling login cookie: %w", err) - } - - opts := h.CookieOptsPathAware(r). - WithExpiresIn(LoginCookieLifetime). - WithSameSite(http.SameSiteNoneMode) - value := string(loginCookieJson) - - err = cookie.EncryptAndSet(w, cookie.Login, value, opts, h.Crypter) - if err != nil { - return err - } - - // set a duplicate cookie without the SameSite value set for user agents that do not properly handle SameSite - err = cookie.EncryptAndSet(w, cookie.LoginLegacy, value, opts.WithSameSite(http.SameSiteDefaultMode), h.Crypter) - if err != nil { - return err - } - - return nil -} - -func (h *Handler) clearLoginCookies(w http.ResponseWriter, r *http.Request) { - opts := h.CookieOptsPathAware(r) - cookie.Clear(w, cookie.Login, opts.WithSameSite(http.SameSiteNoneMode)) - cookie.Clear(w, cookie.LoginLegacy, opts.WithSameSite(http.SameSiteDefaultMode)) -} diff --git a/pkg/handler/handler_logout.go b/pkg/handler/handler_logout.go deleted file mode 100644 index 23b0d4d..0000000 --- a/pkg/handler/handler_logout.go +++ /dev/null @@ -1,52 +0,0 @@ -package handler - -import ( - "errors" - "fmt" - "net/http" - - log "github.com/sirupsen/logrus" - - "github.com/nais/wonderwall/pkg/cookie" - "github.com/nais/wonderwall/pkg/metrics" - logentry "github.com/nais/wonderwall/pkg/middleware" - "github.com/nais/wonderwall/pkg/session" -) - -// Logout triggers self-initiated logout for the current user. -func (h *Handler) Logout(w http.ResponseWriter, r *http.Request) { - logger := logentry.LogEntryFrom(r) - logout, err := h.Client.Logout(r) - if err != nil { - h.InternalError(w, r, err) - return - } - - idToken := "" - - sessionData, err := h.Sessions.Get(r) - if err == nil && sessionData != nil { - idToken = sessionData.IDToken - - err = h.Sessions.DestroyForID(r, sessionData.ExternalSessionID) - if err != nil && !errors.Is(err, session.KeyNotFoundError) { - h.InternalError(w, r, fmt.Errorf("logout: destroying session: %w", err)) - return - } - - fields := log.Fields{ - "jti": sessionData.IDTokenJwtID, - } - logger.WithFields(fields).Info("logout: successful local logout") - } - - cookie.Clear(w, cookie.Session, h.CookieOptsPathAware(r)) - - if h.Loginstatus.Enabled() { - h.Loginstatus.ClearCookie(w, h.CookieOptions) - } - - logger.Debug("logout: redirecting to identity provider") - metrics.ObserveLogout(metrics.LogoutOperationSelfInitiated) - http.Redirect(w, r, logout.SingleLogoutURL(idToken), http.StatusTemporaryRedirect) -} diff --git a/pkg/handler/handler_logout_callback.go b/pkg/handler/handler_logout_callback.go deleted file mode 100644 index 0206b40..0000000 --- a/pkg/handler/handler_logout_callback.go +++ /dev/null @@ -1,15 +0,0 @@ -package handler - -import ( - "net/http" - - logentry "github.com/nais/wonderwall/pkg/middleware" -) - -// LogoutCallback handles the callback initiated by the self-initiated logout after single-logout at the identity provider. -func (h *Handler) LogoutCallback(w http.ResponseWriter, r *http.Request) { - redirect := h.Client.LogoutCallback(r).PostLogoutRedirectURI() - - logentry.LogEntryFrom(r).Debugf("logout/callback: redirecting to %s", redirect) - http.Redirect(w, r, redirect, http.StatusTemporaryRedirect) -} diff --git a/pkg/handler/handler_standard.go b/pkg/handler/handler_standard.go new file mode 100644 index 0000000..5475b16 --- /dev/null +++ b/pkg/handler/handler_standard.go @@ -0,0 +1,141 @@ +package handler + +import ( + "net/http" + + "github.com/nais/wonderwall/pkg/config" + "github.com/nais/wonderwall/pkg/cookie" + "github.com/nais/wonderwall/pkg/crypto" + apilogin "github.com/nais/wonderwall/pkg/handler/api/login" + apilogincallback "github.com/nais/wonderwall/pkg/handler/api/logincallback" + apilogout "github.com/nais/wonderwall/pkg/handler/api/logout" + apilogoutcallback "github.com/nais/wonderwall/pkg/handler/api/logoutcallback" + apilogoutfrontchannel "github.com/nais/wonderwall/pkg/handler/api/logoutfrontchannel" + apisession "github.com/nais/wonderwall/pkg/handler/api/session" + apisessionrefresh "github.com/nais/wonderwall/pkg/handler/api/sessionrefresh" + "github.com/nais/wonderwall/pkg/handler/autologin" + errorhandler "github.com/nais/wonderwall/pkg/handler/error" + "github.com/nais/wonderwall/pkg/handler/reverseproxy" + "github.com/nais/wonderwall/pkg/ingress" + "github.com/nais/wonderwall/pkg/loginstatus" + "github.com/nais/wonderwall/pkg/middleware" + openidclient "github.com/nais/wonderwall/pkg/openid/client" + openidconfig "github.com/nais/wonderwall/pkg/openid/config" + "github.com/nais/wonderwall/pkg/openid/provider" + "github.com/nais/wonderwall/pkg/router" + "github.com/nais/wonderwall/pkg/session" +) + +var _ router.Source = &StandardHandler{} + +type StandardHandler struct { + autoLogin *autologin.AutoLogin + client openidclient.Client + config *config.Config + cookieOptions cookie.Options + crypter crypto.Crypter + loginstatus loginstatus.Loginstatus + openidConfig openidconfig.Config + provider provider.Provider + sessions *session.Handler + upstreamProxy *reverseproxy.ReverseProxy +} + +func (s *StandardHandler) GetAutoLogin() *autologin.AutoLogin { + return s.autoLogin +} + +func (s *StandardHandler) GetClient() openidclient.Client { + return s.client +} + +func (s *StandardHandler) GetCookieOptions() cookie.Options { + return s.cookieOptions +} + +func (s *StandardHandler) GetCookieOptsPathAware(r *http.Request) cookie.Options { + path := s.GetPath(r) + return s.cookieOptions.WithPath(path) +} + +func (s *StandardHandler) GetCrypter() crypto.Crypter { + return s.crypter +} + +func (s *StandardHandler) GetErrorHandler() errorhandler.Handler { + return errorhandler.New(s) +} + +func (s *StandardHandler) GetErrorRedirectURI() string { + return s.config.ErrorRedirectURI +} + +func (s *StandardHandler) GetIngresses() *ingress.Ingresses { + return s.openidConfig.Client().Ingresses() +} + +func (s *StandardHandler) GetLoginstatus() loginstatus.Loginstatus { + return s.loginstatus +} + +func (s *StandardHandler) GetPath(r *http.Request) string { + path, ok := middleware.PathFrom(r.Context()) + if !ok { + path = s.GetIngresses().MatchingPath(r) + } + + return path +} + +func (s *StandardHandler) GetProvider() provider.Provider { + return s.provider +} + +func (s *StandardHandler) GetProviderName() string { + return s.openidConfig.Provider().Name() +} + +func (s *StandardHandler) GetSessions() *session.Handler { + return s.sessions +} + +func (s *StandardHandler) GetSessionConfig() config.Session { + return s.config.Session +} + +func (s *StandardHandler) Login(w http.ResponseWriter, r *http.Request) { + apilogin.Handler(s, w, r) +} + +func (s *StandardHandler) LoginCallback(w http.ResponseWriter, r *http.Request) { + apilogincallback.Handler(s, w, r) +} + +func (s *StandardHandler) Logout(w http.ResponseWriter, r *http.Request) { + apilogout.Handler(s, w, r) +} + +func (s *StandardHandler) LogoutCallback(w http.ResponseWriter, r *http.Request) { + apilogoutcallback.Handler(s, w, r) +} + +func (s *StandardHandler) LogoutFrontChannel(w http.ResponseWriter, r *http.Request) { + apilogoutfrontchannel.Handler(s, w, r) +} + +func (s *StandardHandler) Session(w http.ResponseWriter, r *http.Request) { + apisession.Handler(s, w, r) +} + +func (s *StandardHandler) SessionRefresh(w http.ResponseWriter, r *http.Request) { + if !s.config.Session.Refresh { + http.NotFound(w, r) + return + } + + apisessionrefresh.Handler(s, w, r) +} + +func (s *StandardHandler) ReverseProxy(w http.ResponseWriter, r *http.Request) { + s.upstreamProxy.Handler(s, w, r) +} diff --git a/pkg/handler/handler_test.go b/pkg/handler/handler_test.go index 385b3be..af398cb 100644 --- a/pkg/handler/handler_test.go +++ b/pkg/handler/handler_test.go @@ -31,9 +31,9 @@ func TestHandler_Login(t *testing.T) { resp := localLogin(t, rpClient, idp) loginURL := resp.Location - req := idp.GetRequest(idp.RelyingPartyServer.URL + "/oauth2/logout") + req := idp.GetRequest(idp.RelyingPartyServer.URL + "/oauth2/login") - expectedCallbackURL, err := urlpkg.CallbackURL(req) + expectedCallbackURL, err := urlpkg.LoginCallbackURL(req) assert.NoError(t, err) assert.Equal(t, idp.ProviderServer.URL, fmt.Sprintf("%s://%s", loginURL.Scheme, loginURL.Host)) @@ -116,10 +116,10 @@ func TestHandler_FrontChannelLogout(t *testing.T) { ciphertext, err := base64.StdEncoding.DecodeString(sessionCookie.Value) assert.NoError(t, err) - sessionKey, err := idp.RelyingPartyHandler.Crypter.Decrypt(ciphertext) + sessionKey, err := idp.RelyingPartyHandler.GetCrypter().Decrypt(ciphertext) assert.NoError(t, err) - data, err := idp.RelyingPartyHandler.Sessions.GetForKey(r, string(sessionKey)) + data, err := idp.RelyingPartyHandler.GetSessions().GetForKey(r, string(sessionKey)) assert.NoError(t, err) return data.ExternalSessionID @@ -399,7 +399,7 @@ func TestHandler_Default(t *testing.T) { callbackEndpoint.RawQuery = "" req := idp.GetRequest(callbackLocation.String()) - expectedCallbackURL, err := urlpkg.CallbackURL(req) + expectedCallbackURL, err := urlpkg.LoginCallbackURL(req) assert.NoError(t, err) assert.Equal(t, expectedCallbackURL, callbackEndpoint.String()) diff --git a/pkg/handler/reverseproxy/reverseproxy.go b/pkg/handler/reverseproxy/reverseproxy.go new file mode 100644 index 0000000..45f412d --- /dev/null +++ b/pkg/handler/reverseproxy/reverseproxy.go @@ -0,0 +1,85 @@ +package reverseproxy + +import ( + "errors" + "net/http" + "net/http/httputil" + + "github.com/nais/wonderwall/pkg/handler/autologin" + "github.com/nais/wonderwall/pkg/handler/url" + "github.com/nais/wonderwall/pkg/loginstatus" + mw "github.com/nais/wonderwall/pkg/middleware" + "github.com/nais/wonderwall/pkg/session" +) + +type Source interface { + GetAutoLogin() *autologin.AutoLogin + GetLoginstatus() loginstatus.Loginstatus + GetPath(r *http.Request) string + GetSessions() *session.Handler +} + +type ReverseProxy struct { + *httputil.ReverseProxy +} + +func New(upstreamHost string) *ReverseProxy { + rp := &httputil.ReverseProxy{ + Director: func(r *http.Request) { + // Delete incoming authentication + r.Header.Del("authorization") + // Instruct http.ReverseProxy to not modify X-Forwarded-For header + r.Header["X-Forwarded-For"] = nil + // Request should go to correct host + r.URL.Host = upstreamHost + r.URL.Scheme = "http" + + accessToken, ok := mw.AccessTokenFrom(r.Context()) + if ok { + r.Header.Set("authorization", "Bearer "+accessToken) + } + }, + ErrorHandler: func(w http.ResponseWriter, r *http.Request, err error) { + http.Error(w, err.Error(), http.StatusBadGateway) + }, + } + return &ReverseProxy{rp} +} + +func (rp *ReverseProxy) Handler(src Source, w http.ResponseWriter, r *http.Request) { + logger := mw.LogEntryFrom(r).WithField("request_path", r.URL.Path) + isAuthenticated := false + + accessToken, err := src.GetSessions().GetAccessToken(r) + if err == nil { + // add authentication if session cookie and token checks out + isAuthenticated = true + + // force new authentication if loginstatus is enabled and cookie isn't set + if src.GetLoginstatus().NeedsLogin(r) { + isAuthenticated = false + logger.Info("default: loginstatus was enabled, but no matching cookie was found; state is now unauthenticated") + } + } else if errors.Is(err, session.UnexpectedError) { + logger.Errorf("default: getting session: %+v", err) + } + + if src.GetAutoLogin().NeedsLogin(r, isAuthenticated) { + logger.Debug("default: auto-login is enabled; request does not match any configured ignorable paths") + + redirectTarget := r.URL.String() + path := src.GetPath(r) + + loginUrl := url.LoginURL(path, redirectTarget) + http.Redirect(w, r, loginUrl, http.StatusTemporaryRedirect) + return + } + + ctx := r.Context() + + if isAuthenticated { + ctx = mw.WithAccessToken(ctx, accessToken) + } + + rp.ServeHTTP(w, r.WithContext(ctx)) +} diff --git a/pkg/handler/templates/error.go b/pkg/handler/templates/error.go new file mode 100644 index 0000000..4c64867 --- /dev/null +++ b/pkg/handler/templates/error.go @@ -0,0 +1,22 @@ +package templates + +import ( + _ "embed" + "html/template" + + log "github.com/sirupsen/logrus" +) + +//go:embed error.gohtml +var errorGoHtml string +var ErrorTemplate *template.Template + +func init() { + var err error + + ErrorTemplate = template.New("error") + ErrorTemplate, err = ErrorTemplate.Parse(errorGoHtml) + if err != nil { + log.Fatalf("parsing error template: %+v", err) + } +} diff --git a/pkg/handler/url/url.go b/pkg/handler/url/url.go index 6a58597..d23a8cc 100644 --- a/pkg/handler/url/url.go +++ b/pkg/handler/url/url.go @@ -68,8 +68,8 @@ func LoginURL(prefix, redirectTarget string) string { return loginPath + redirectParam } -func CallbackURL(r *http.Request) (string, error) { - return makeCallbackURL(r, paths.Callback) +func LoginCallbackURL(r *http.Request) (string, error) { + return makeCallbackURL(r, paths.LoginCallback) } func LogoutCallbackURL(r *http.Request) (string, error) { diff --git a/pkg/handler/url/url_test.go b/pkg/handler/url/url_test.go index 512c51b..c2518e4 100644 --- a/pkg/handler/url/url_test.go +++ b/pkg/handler/url/url_test.go @@ -187,7 +187,7 @@ func TestLoginURL(t *testing.T) { } } -func TestClient_CallbackURL(t *testing.T) { +func TestLoginCallbackURL(t *testing.T) { cfg := mock.Config() cfg.Ingresses = []string{ "https://nav.no", @@ -226,7 +226,7 @@ func TestClient_CallbackURL(t *testing.T) { t.Run(test.input, func(t *testing.T) { req := mock.NewGetRequest(test.input, openidConfig) - actual, err := urlpkg.CallbackURL(req) + actual, err := urlpkg.LoginCallbackURL(req) if test.err != nil { assert.EqualError(t, err, test.err.Error()) } else { @@ -237,7 +237,7 @@ func TestClient_CallbackURL(t *testing.T) { } } -func TestClient_LogoutCallbackURL(t *testing.T) { +func TestLogoutCallbackURL(t *testing.T) { cfg := mock.Config() cfg.Ingresses = []string{ "https://nav.no", diff --git a/pkg/middleware/ingress.go b/pkg/middleware/ingress.go index 8726ee2..1f1d6a3 100644 --- a/pkg/middleware/ingress.go +++ b/pkg/middleware/ingress.go @@ -7,7 +7,7 @@ import ( ) type IngressSource interface { - Ingresses() *ingress.Ingresses + GetIngresses() *ingress.Ingresses } type IngressMiddleware struct { @@ -20,7 +20,7 @@ func Ingress(source IngressSource) IngressMiddleware { func (i *IngressMiddleware) Handler(next http.Handler) http.Handler { fn := func(w http.ResponseWriter, r *http.Request) { - ingresses := i.Ingresses() + ingresses := i.GetIngresses() ctx := r.Context() path := ingresses.MatchingPath(r) diff --git a/pkg/mock/openid.go b/pkg/mock/openid.go index 4295798..57b7457 100644 --- a/pkg/mock/openid.go +++ b/pkg/mock/openid.go @@ -18,6 +18,7 @@ import ( "github.com/lestrrat-go/jwx/v2/jwt" "github.com/nais/wonderwall/pkg/config" + "github.com/nais/wonderwall/pkg/cookie" "github.com/nais/wonderwall/pkg/crypto" handlerpkg "github.com/nais/wonderwall/pkg/handler" "github.com/nais/wonderwall/pkg/openid" @@ -34,7 +35,7 @@ type IdentityProvider struct { Provider *TestProvider ProviderHandler *IdentityProviderHandler ProviderServer *httptest.Server - RelyingPartyHandler *handlerpkg.Handler + RelyingPartyHandler *handlerpkg.StandardHandler RelyingPartyServer *httptest.Server } @@ -84,12 +85,12 @@ func NewIdentityProvider(cfg *config.Config) *IdentityProvider { crypter := crypto.NewCrypter([]byte(cfg.EncryptionKey)) ctx, cancel := context.WithCancel(context.Background()) - rpHandler, err := handlerpkg.NewHandler(ctx, cfg, openidConfig, crypter) + cookieOpts := cookie.DefaultOptions().WithSecure(false) + rpHandler, err := handlerpkg.NewHandler(ctx, cfg, cookieOpts, openidConfig, crypter) if err != nil { panic(err) } - rpHandler.CookieOptions = rpHandler.CookieOptions.WithSecure(false) rpRouter := router.New(rpHandler) rpServer := httptest.NewServer(rpRouter) diff --git a/pkg/openid/client/login.go b/pkg/openid/client/login.go index 213b9f7..c28d8ba 100644 --- a/pkg/openid/client/login.go +++ b/pkg/openid/client/login.go @@ -53,7 +53,7 @@ func NewLogin(c Client, r *http.Request, loginstatus loginstatus.Loginstatus) (L return nil, fmt.Errorf("generating parameters: %w", err) } - callbackURL, err := urlpkg.CallbackURL(r) + callbackURL, err := urlpkg.LoginCallbackURL(r) if err != nil { return nil, fmt.Errorf("generating callback url: %w", err) } diff --git a/pkg/openid/client/login_callback_test.go b/pkg/openid/client/login_callback_test.go index 2425235..c4a7fd7 100644 --- a/pkg/openid/client/login_callback_test.go +++ b/pkg/openid/client/login_callback_test.go @@ -117,7 +117,7 @@ func newLoginCallback(t *testing.T, url string) (*mock.IdentityProvider, client. cfg := idp.OpenIDConfig - redirect, err := urlpkg.CallbackURL(req) + redirect, err := urlpkg.LoginCallbackURL(req) assert.NoError(t, err) idp.ProviderHandler.Codes = map[string]*mock.AuthorizeRequest{ diff --git a/pkg/openid/client/login_test.go b/pkg/openid/client/login_test.go index 8c22685..ab7055d 100644 --- a/pkg/openid/client/login_test.go +++ b/pkg/openid/client/login_test.go @@ -90,7 +90,7 @@ func TestLogin_URL(t *testing.T) { assert.Contains(t, query, "code_challenge_method") assert.NotContains(t, query, "resource") - callbackURL, err := urlpkg.CallbackURL(req) + callbackURL, err := urlpkg.LoginCallbackURL(req) assert.NoError(t, err) assert.ElementsMatch(t, query["response_type"], []string{"code"}) diff --git a/pkg/openid/cookies.go b/pkg/openid/cookies.go index 939ee2c..d7c3699 100644 --- a/pkg/openid/cookies.go +++ b/pkg/openid/cookies.go @@ -1,5 +1,15 @@ package openid +import ( + "encoding/json" + "fmt" + "net/http" + + "github.com/nais/wonderwall/pkg/cookie" + "github.com/nais/wonderwall/pkg/crypto" + "github.com/nais/wonderwall/pkg/middleware" +) + type LoginCookie struct { State string `json:"state"` Nonce string `json:"nonce"` @@ -7,3 +17,23 @@ type LoginCookie struct { Referer string `json:"referer"` RedirectURI string `json:"redirect_uri"` } + +func GetLoginCookie(r *http.Request, crypter crypto.Crypter) (*LoginCookie, error) { + loginCookieJson, err := cookie.GetDecrypted(r, cookie.Login, crypter) + if err != nil { + middleware.LogEntryFrom(r).Debugf("failed to fetch login cookie: %+v; falling back to legacy cookie", err) + + loginCookieJson, err = cookie.GetDecrypted(r, cookie.LoginLegacy, crypter) + if err != nil { + return nil, err + } + } + + var loginCookie LoginCookie + err = json.Unmarshal([]byte(loginCookieJson), &loginCookie) + if err != nil { + return nil, fmt.Errorf("unmarshalling: %w", err) + } + + return &loginCookie, nil +} diff --git a/pkg/router/paths/paths.go b/pkg/router/paths/paths.go index 8a22e61..85b88e4 100644 --- a/pkg/router/paths/paths.go +++ b/pkg/router/paths/paths.go @@ -3,10 +3,10 @@ package paths const ( OAuth2 = "/oauth2" Login = "/login" - Callback = "/callback" + LoginCallback = "/callback" Logout = "/logout" LogoutCallback = "/logout/callback" - FrontChannelLogout = "/logout/frontchannel" + LogoutFrontChannel = "/logout/frontchannel" Session = "/session" SessionRefresh = "/session/refresh" ) diff --git a/pkg/router/router.go b/pkg/router/router.go index 12239ea..af3b4c2 100644 --- a/pkg/router/router.go +++ b/pkg/router/router.go @@ -11,35 +11,46 @@ import ( "github.com/nais/wonderwall/pkg/router/paths" ) -type Handler interface { - Login(http.ResponseWriter, *http.Request) - Callback(http.ResponseWriter, *http.Request) - - Logout(http.ResponseWriter, *http.Request) - LogoutCallback(http.ResponseWriter, *http.Request) - - FrontChannelLogout(http.ResponseWriter, *http.Request) - - SessionInfo(http.ResponseWriter, *http.Request) - SessionRefresh(http.ResponseWriter, *http.Request) - - Default(http.ResponseWriter, *http.Request) - - Ingresses() *ingress.Ingresses - ProviderName() string +type Source interface { + Handlers + Config } -func New(handler Handler) chi.Router { - ingressMw := middleware.Ingress(handler) - prometheus := middleware.Prometheus(handler.ProviderName()) - logentry := middleware.LogEntry(handler.ProviderName()) +type Handlers interface { + // Login initiates the authorization code flow. + Login(http.ResponseWriter, *http.Request) + // LoginCallback handles the authentication response from the identity provider. + LoginCallback(http.ResponseWriter, *http.Request) + // Logout triggers self-initiated logout for the current user. + Logout(http.ResponseWriter, *http.Request) + // LogoutCallback handles the callback initiated by the self-initiated logout after single-logout at the identity provider. + LogoutCallback(http.ResponseWriter, *http.Request) + // LogoutFrontChannel performs a local logout initiated by a third party in the SSO circle-of-trust. + LogoutFrontChannel(http.ResponseWriter, *http.Request) + // Session returns metadata for the current user's session. + Session(http.ResponseWriter, *http.Request) + // SessionRefresh refreshes current user's session and returns the associated updated metadata. + SessionRefresh(http.ResponseWriter, *http.Request) + // ReverseProxy proxies all requests upstream. + ReverseProxy(http.ResponseWriter, *http.Request) +} + +type Config interface { + GetIngresses() *ingress.Ingresses + GetProviderName() string +} + +func New(src Source) chi.Router { + ingressMw := middleware.Ingress(src) + prometheus := middleware.Prometheus(src.GetProviderName()) + logentry := middleware.LogEntry(src.GetProviderName()) r := chi.NewRouter() r.Use(middleware.CorrelationIDHandler) r.Use(chi_middleware.Recoverer) r.Use(ingressMw.Handler) - prefixes := handler.Ingresses().Paths() + prefixes := src.GetIngresses().Paths() r.Group(func(r chi.Router) { r.Use(logentry.Handler) @@ -48,17 +59,17 @@ func New(handler Handler) chi.Router { for _, prefix := range prefixes { r.Route(prefix+paths.OAuth2, func(r chi.Router) { - r.Get(paths.Login, handler.Login) - r.Get(paths.Callback, handler.Callback) - r.Get(paths.Logout, handler.Logout) - r.Get(paths.FrontChannelLogout, handler.FrontChannelLogout) - r.Get(paths.LogoutCallback, handler.LogoutCallback) - r.Get(paths.Session, handler.SessionInfo) - r.Get(paths.SessionRefresh, handler.SessionRefresh) + r.Get(paths.Login, src.Login) + r.Get(paths.LoginCallback, src.LoginCallback) + r.Get(paths.Logout, src.Logout) + r.Get(paths.LogoutFrontChannel, src.LogoutFrontChannel) + r.Get(paths.LogoutCallback, src.LogoutCallback) + r.Get(paths.Session, src.Session) + r.Get(paths.SessionRefresh, src.SessionRefresh) }) } }) - r.HandleFunc("/*", handler.Default) + r.HandleFunc("/*", src.ReverseProxy) return r }