refactor(handler): split up request handlers into separate modules

This commit is contained in:
Trong Huu Nguyen
2022-09-02 14:53:11 +02:00
parent 5d00d132dd
commit 9144056e28
30 changed files with 774 additions and 548 deletions
+3 -1
View File
@@ -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)
}
+81
View File
@@ -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
}
@@ -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()
}
+63
View File
@@ -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)
}
@@ -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)
}
@@ -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)
@@ -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())
@@ -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)
@@ -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)
}
@@ -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)
+14 -72
View File
@@ -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)
},
}
}
-129
View File
@@ -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()
}
-49
View File
@@ -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))
}
-97
View File
@@ -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))
}
-52
View File
@@ -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)
}
-15
View File
@@ -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)
}
+141
View File
@@ -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)
}
+5 -5
View File
@@ -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())
+85
View File
@@ -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))
}
+22
View File
@@ -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)
}
}
+2 -2
View File
@@ -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) {
+3 -3
View File
@@ -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",
+2 -2
View File
@@ -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)
+4 -3
View File
@@ -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)
+1 -1
View File
@@ -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)
}
+1 -1
View File
@@ -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{
+1 -1
View File
@@ -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"})
+30
View File
@@ -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
}
+2 -2
View File
@@ -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"
)
+40 -29
View File
@@ -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
}