From ad3201fbfb443407d242c2c1c83ea2faf3bd5dc7 Mon Sep 17 00:00:00 2001 From: Trong Huu Nguyen Date: Mon, 11 Jul 2022 13:09:10 +0200 Subject: [PATCH] refactor(handler/logout): extract to openid client --- pkg/openid/client/client.go | 12 +++-- pkg/openid/client/login.go | 2 +- pkg/openid/client/login_callback.go | 2 +- pkg/openid/client/logout.go | 64 +++++++++++++++++++++++---- pkg/router/handler_logout.go | 56 +++++------------------ pkg/router/handler_logout_callback.go | 5 --- 6 files changed, 78 insertions(+), 63 deletions(-) diff --git a/pkg/openid/client/client.go b/pkg/openid/client/client.go index 0020f31..4fb4b30 100644 --- a/pkg/openid/client/client.go +++ b/pkg/openid/client/client.go @@ -21,7 +21,7 @@ type Client interface { Login(r *http.Request) (Login, error) LoginCallback(r *http.Request, p provider.Provider, cookie *openid.LoginCookie) LoginCallback - Logout(r *http.Request) error + Logout() (Logout, error) LogoutCallback(r *http.Request) error LogoutFrontchannel(r *http.Request) LogoutFrontchannel @@ -74,9 +74,13 @@ func (c client) LoginCallback(r *http.Request, p provider.Provider, cookie *open return NewLoginCallback(c, r, p, cookie) } -func (c client) Logout(r *http.Request) error { - //TODO implement me - panic("implement me") +func (c client) Logout() (Logout, error) { + logout, err := NewLogout(c) + if err != nil { + return nil, fmt.Errorf("logout: %w", err) + } + + return logout, nil } func (c client) LogoutCallback(r *http.Request) error { diff --git a/pkg/openid/client/login.go b/pkg/openid/client/login.go index 9e47f27..4d44535 100644 --- a/pkg/openid/client/login.go +++ b/pkg/openid/client/login.go @@ -56,7 +56,7 @@ func NewLogin(c Client, r *http.Request) (Login, error) { redirect := request.CanonicalRedirectURL(r, c.config().Wonderwall().Ingress) cookie := params.cookie(redirect) - return login{ + return &login{ authCodeURL: url, canonicalRedirect: redirect, cookie: cookie, diff --git a/pkg/openid/client/login_callback.go b/pkg/openid/client/login_callback.go index 4620375..233d9e7 100644 --- a/pkg/openid/client/login_callback.go +++ b/pkg/openid/client/login_callback.go @@ -30,7 +30,7 @@ type loginCallback struct { } func NewLoginCallback(c Client, r *http.Request, p provider.Provider, cookie *openid.LoginCookie) LoginCallback { - return loginCallback{ + return &loginCallback{ client: c, cookie: cookie, provider: p, diff --git a/pkg/openid/client/logout.go b/pkg/openid/client/logout.go index 7e12e0c..0610017 100644 --- a/pkg/openid/client/logout.go +++ b/pkg/openid/client/logout.go @@ -1,17 +1,65 @@ package client -import "github.com/nais/wonderwall/pkg/openid" +import ( + "fmt" + "net/url" -type Logout struct { + "github.com/nais/wonderwall/pkg/openid" + "github.com/nais/wonderwall/pkg/strings" +) + +type Logout interface { + CanonicalRedirect() string + Cookie() *openid.LogoutCookie + SingleLogoutURL(idToken string) string +} + +type logout struct { Client + cookie *openid.LogoutCookie + endSessionEndpoint *url.URL } -func (in Logout) URL() string { - // TODO - panic("not implemented") +func NewLogout(c Client) (Logout, error) { + state, err := strings.GenerateBase64(32) + if err != nil { + return nil, fmt.Errorf("generating state: %w", err) + } + + cookie := &openid.LogoutCookie{ + State: state, + RedirectTo: c.config().Client().GetPostLogoutRedirectURI(), + } + + endSessionEndpoint, err := url.Parse(c.config().Provider().EndSessionEndpoint) + if err != nil { + return nil, fmt.Errorf("parsing end session endpoint: %w", err) + } + + return &logout{ + Client: c, + cookie: cookie, + endSessionEndpoint: endSessionEndpoint, + }, nil } -func (in Logout) Cookie() openid.LogoutCookie { - // TODO - panic("not implemented") +func (in logout) CanonicalRedirect() string { + return in.cookie.RedirectTo +} + +func (in logout) Cookie() *openid.LogoutCookie { + return in.cookie +} + +func (in logout) SingleLogoutURL(idToken string) string { + v := in.endSessionEndpoint.Query() + v.Add("post_logout_redirect_uri", in.config().Client().GetLogoutCallbackURI()) + v.Add("state", in.cookie.State) + + if len(idToken) > 0 { + v.Add("id_token_hint", idToken) + } + + in.endSessionEndpoint.RawQuery = v.Encode() + return in.endSessionEndpoint.String() } diff --git a/pkg/router/handler_logout.go b/pkg/router/handler_logout.go index ed37a9c..7adccb9 100644 --- a/pkg/router/handler_logout.go +++ b/pkg/router/handler_logout.go @@ -5,14 +5,17 @@ import ( "errors" "fmt" "net/http" - "net/url" + "time" "github.com/go-redis/redis/v8" "github.com/nais/wonderwall/pkg/cookie" "github.com/nais/wonderwall/pkg/openid" logentry "github.com/nais/wonderwall/pkg/router/middleware" - "github.com/nais/wonderwall/pkg/strings" +) + +const ( + LogoutCookieLifetime = 5 * time.Minute ) // Logout triggers self-initiated for the current user @@ -41,53 +44,24 @@ func (h *Handler) Logout(w http.ResponseWriter, r *http.Request) { h.Loginstatus.ClearCookie(w, h.CookieOptions) } - u, err := url.Parse(h.Cfg.Provider().EndSessionEndpoint) + logout, err := h.Client.Logout() if err != nil { - h.InternalError(w, r, fmt.Errorf("logout: parsing end session endpoint: %w", err)) - return + h.InternalError(w, r, err) } - logoutCookie, err := h.logoutCookie() - if err != nil { - h.InternalError(w, r, fmt.Errorf("logout: generating logout cookie: %w", err)) - return - } - - err = h.setLogoutCookie(w, logoutCookie) + err = h.setLogoutCookie(w, logout.Cookie()) if err != nil { h.InternalError(w, r, fmt.Errorf("logout: setting logout cookie: %w", err)) return } - v := u.Query() - v.Add("post_logout_redirect_uri", h.Cfg.Client().GetLogoutCallbackURI()) - v.Add("state", logoutCookie.State) - - if len(idToken) > 0 { - v.Add("id_token_hint", idToken) - } - - u.RawQuery = v.Encode() - fields := map[string]interface{}{ - "redirect_to": logoutCookie.RedirectTo, + "redirect_to": logout.CanonicalRedirect(), } logger := logentry.LogEntryWithFields(r.Context(), fields) logger.Info().Msg("logout: redirecting to identity provider") - http.Redirect(w, r, u.String(), http.StatusTemporaryRedirect) -} - -func (h *Handler) logoutCookie() (*openid.LogoutCookie, error) { - state, err := strings.GenerateBase64(32) - if err != nil { - return nil, fmt.Errorf("generating state: %w", err) - } - - return &openid.LogoutCookie{ - State: state, - RedirectTo: h.Cfg.Client().GetPostLogoutRedirectURI(), - }, nil + http.Redirect(w, r, logout.SingleLogoutURL(idToken), http.StatusTemporaryRedirect) } func (h *Handler) setLogoutCookie(w http.ResponseWriter, logoutCookie *openid.LogoutCookie) error { @@ -96,14 +70,8 @@ func (h *Handler) setLogoutCookie(w http.ResponseWriter, logoutCookie *openid.Lo return fmt.Errorf("marshalling login cookie: %w", err) } - opts := h.CookieOptions. - WithExpiresIn(LogoutCookieLifetime) + opts := h.CookieOptions.WithExpiresIn(LogoutCookieLifetime) value := string(logoutCookieJson) - err = cookie.EncryptAndSet(w, cookie.Logout, value, opts, h.Crypter) - if err != nil { - return err - } - - return nil + return cookie.EncryptAndSet(w, cookie.Logout, value, opts, h.Crypter) } diff --git a/pkg/router/handler_logout_callback.go b/pkg/router/handler_logout_callback.go index 275d69f..8e18e47 100644 --- a/pkg/router/handler_logout_callback.go +++ b/pkg/router/handler_logout_callback.go @@ -4,17 +4,12 @@ import ( "encoding/json" "fmt" "net/http" - "time" "github.com/nais/wonderwall/pkg/cookie" "github.com/nais/wonderwall/pkg/openid" logentry "github.com/nais/wonderwall/pkg/router/middleware" ) -const ( - LogoutCookieLifetime = 5 * time.Minute -) - // LogoutCallback handles the callback from the self-initiated logout for the current user func (h *Handler) LogoutCallback(w http.ResponseWriter, r *http.Request) { cookie.Clear(w, cookie.Logout, h.CookieOptions)