refactor(openid/client): remove indirection layer for login callback

Co-authored-by: sindrerh2 <sindre.rodseth.hansen@nav.no>
This commit is contained in:
Trong Huu Nguyen
2025-01-21 09:39:23 +01:00
co-authored by sindrerh2
parent 75f98debc5
commit 64e9167e05
4 changed files with 131 additions and 213 deletions
+1 -12
View File
@@ -189,28 +189,17 @@ func (s *Standalone) LoginCallback(w http.ResponseWriter, r *http.Request) {
return
}
loginCallback, err := s.Client.LoginCallback(r, loginCookie)
tokens, err := s.Client.LoginCallback(r, loginCookie)
if err != nil {
if errors.Is(err, openidclient.ErrCallbackInvalidState) || errors.Is(err, openidclient.ErrCallbackInvalidIssuer) {
s.Unauthorized(w, r, err)
return
}
if errors.Is(err, openidclient.ErrCallbackIdentityProvider) {
s.InternalError(w, r, err)
return
}
s.InternalError(w, r, err)
return
}
tokens, err := loginCallback.RedeemTokens(r.Context())
if err != nil {
s.InternalError(w, r, fmt.Errorf("callback: redeeming tokens: %w", err))
return
}
sessionLifetime := s.Config.Session.MaxLifetime
sess, err := s.SessionManager.Create(r, tokens, sessionLifetime)
-9
View File
@@ -75,15 +75,6 @@ func (c *Client) Login(r *http.Request) (*Login, error) {
return login, nil
}
func (c *Client) LoginCallback(r *http.Request, cookie *openid.LoginCookie) (*LoginCallback, error) {
loginCallback, err := NewLoginCallback(c, r, cookie)
if err != nil {
return nil, fmt.Errorf("callback: %w", err)
}
return loginCallback, nil
}
func (c *Client) Logout(r *http.Request) (*Logout, error) {
logout, err := NewLogout(c, r)
if err != nil {
+43 -49
View File
@@ -5,44 +5,28 @@ import (
"errors"
"fmt"
"net/http"
"net/url"
"golang.org/x/oauth2"
"github.com/nais/wonderwall/pkg/openid"
urlpkg "github.com/nais/wonderwall/pkg/url"
)
var (
ErrCallbackIdentityProvider = errors.New("identity provider error")
ErrCallbackInvalidState = errors.New("invalid state")
ErrCallbackInvalidIssuer = errors.New("invalid issuer")
ErrCallbackIdentityProvider = errors.New("callback: identity provider error")
ErrCallbackInvalidCookie = errors.New("callback: invalid cookie")
ErrCallbackInvalidState = errors.New("callback: invalid state")
ErrCallbackInvalidIssuer = errors.New("callback: invalid issuer")
ErrCallbackRedeemTokens = errors.New("callback: redeeming tokens")
)
type LoginCallback struct {
*Client
cookie *openid.LoginCookie
query url.Values
}
func NewLoginCallback(c *Client, r *http.Request, cookie *openid.LoginCookie) (*LoginCallback, error) {
func (c *Client) LoginCallback(r *http.Request, cookie *openid.LoginCookie) (*openid.Tokens, error) {
if cookie == nil {
return nil, fmt.Errorf("cookie is nil")
}
// redirect_uri not set in cookie (e.g. login initiated at instance running older version, callback handled at newer version)
if len(cookie.RedirectURI) == 0 {
callbackURL, err := urlpkg.LoginCallback(r)
if err != nil {
return nil, fmt.Errorf("generating callback url: %w", err)
}
cookie.RedirectURI = callbackURL
return nil, fmt.Errorf("%w: %s", ErrCallbackInvalidCookie, "cookie is nil")
}
query := r.URL.Query()
if query.Get("error") != "" {
oauthError := query.Get("error")
if oauthError := query.Get("error"); len(oauthError) > 0 {
oauthErrorDescription := query.Get("error_description")
return nil, fmt.Errorf("%w: %s: %s", ErrCallbackIdentityProvider, oauthError, oauthErrorDescription)
}
@@ -51,49 +35,59 @@ func NewLoginCallback(c *Client, r *http.Request, cookie *openid.LoginCookie) (*
return nil, fmt.Errorf("%w: %s", ErrCallbackInvalidState, err)
}
if c.cfg.Provider().AuthorizationResponseIssParameterSupported() {
iss := query.Get("iss")
expectedIss := c.cfg.Provider().Issuer()
if len(iss) == 0 {
return nil, fmt.Errorf("%w: missing issuer parameter", ErrCallbackInvalidIssuer)
}
if iss != expectedIss {
return nil, fmt.Errorf("%w: issuer mismatch: expected %s, got %s", ErrCallbackInvalidIssuer, expectedIss, iss)
}
if err := c.authorizationServerIssuerIdentification(query.Get("iss")); err != nil {
return nil, fmt.Errorf("%w: %s", ErrCallbackInvalidIssuer, err)
}
return &LoginCallback{
Client: c,
cookie: cookie,
query: query,
}, nil
tokens, err := c.redeemTokens(r.Context(), query.Get("code"), cookie)
if err != nil {
return nil, fmt.Errorf("%w: %s", ErrCallbackRedeemTokens, err)
}
return tokens, nil
}
func (in *LoginCallback) RedeemTokens(ctx context.Context) (*openid.Tokens, error) {
params, err := in.AuthParams()
// Verify iss parameter if provider supports RFC 9207 - OAuth 2.0 Authorization Server Issuer Identification
func (c *Client) authorizationServerIssuerIdentification(iss string) error {
if !c.cfg.Provider().AuthorizationResponseIssParameterSupported() {
return nil
}
if len(iss) == 0 {
return fmt.Errorf("missing issuer parameter")
}
expectedIss := c.cfg.Provider().Issuer()
if iss != expectedIss {
return fmt.Errorf("issuer mismatch: expected %q, got %q", expectedIss, iss)
}
return nil
}
func (c *Client) redeemTokens(ctx context.Context, code string, cookie *openid.LoginCookie) (*openid.Tokens, error) {
params, err := c.AuthParams()
if err != nil {
return nil, err
}
rawTokens, err := in.AuthCodeGrant(ctx, in.query.Get("code"), params.AuthCodeOptions([]oauth2.AuthCodeOption{
openid.RedirectURIOption(in.cookie.RedirectURI),
oauth2.VerifierOption(in.cookie.CodeVerifier),
rawTokens, err := c.AuthCodeGrant(ctx, code, params.AuthCodeOptions([]oauth2.AuthCodeOption{
openid.RedirectURIOption(cookie.RedirectURI),
oauth2.VerifierOption(cookie.CodeVerifier),
}))
if err != nil {
return nil, fmt.Errorf("exchanging authorization code for token: %w", err)
}
jwkSet, err := in.jwksProvider.GetPublicJwkSet(ctx)
jwkSet, err := c.jwksProvider.GetPublicJwkSet(ctx)
if err != nil {
return nil, fmt.Errorf("getting jwks: %w", err)
}
tokens, err := openid.NewTokens(rawTokens, jwkSet, in.cfg, in.cookie)
tokens, err := openid.NewTokens(rawTokens, jwkSet, c.cfg, cookie)
if err != nil {
// JWKS might not be up to date, so we'll want to force a refresh for the next attempt
_, _ = in.jwksProvider.RefreshPublicJwkSet(ctx)
_, _ = c.jwksProvider.RefreshPublicJwkSet(ctx)
return nil, fmt.Errorf("parsing tokens: %w", err)
}
+87 -143
View File
@@ -1,7 +1,6 @@
package client_test
import (
"context"
"testing"
"time"
@@ -16,99 +15,11 @@ import (
)
func TestLoginCallback(t *testing.T) {
t.Run("invalid state", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?state=some-other-state"
idp := mock.NewIdentityProvider(mock.Config())
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
assert.Nil(t, lc)
assert.ErrorIs(t, err, client.ErrCallbackInvalidState)
})
t.Run("missing state", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback"
idp := mock.NewIdentityProvider(mock.Config())
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
assert.Nil(t, lc)
assert.ErrorIs(t, err, client.ErrCallbackInvalidState)
})
t.Run("identity provider error", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?error=invalid_client&error_description=client%20authenticaion%20failed"
idp := mock.NewIdentityProvider(mock.Config())
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
assert.Nil(t, lc)
assert.ErrorIs(t, err, client.ErrCallbackIdentityProvider)
})
t.Run("supports authorization response with iss parameter", func(t *testing.T) {
idp := mock.NewIdentityProvider(mock.Config())
idp.OpenIDConfig.TestProvider.WithAuthorizationResponseIssParameterSupported()
for _, tt := range []struct {
name string
iss string
assertions func(t *testing.T, lc *client.LoginCallback, err error)
}{
{
name: "happy path",
iss: idp.OpenIDConfig.TestProvider.Issuer(),
assertions: func(t *testing.T, lc *client.LoginCallback, err error) {
assert.NotNil(t, lc)
assert.NoError(t, err)
},
},
{
name: "missing issuer",
iss: "",
assertions: func(t *testing.T, lc *client.LoginCallback, err error) {
assert.Nil(t, lc)
assert.ErrorIs(t, err, client.ErrCallbackInvalidIssuer)
},
},
{
name: "wrong issuer",
iss: "https://wrong-issuer",
assertions: func(t *testing.T, lc *client.LoginCallback, err error) {
assert.Nil(t, lc)
assert.ErrorIs(t, err, client.ErrCallbackInvalidIssuer)
},
},
} {
t.Run(tt.name, func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?state=some-state"
if tt.iss != "" {
url += "&iss=" + tt.iss
}
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
tt.assertions(t, lc, err)
})
}
})
}
func TestLoginCallback_RedeemTokens(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state"
t.Run("happy path", func(t *testing.T) {
idp := mock.NewIdentityProvider(mock.Config())
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state"
tokens, err := newLoginCallback(t, url, nil)
require.NoError(t, err)
require.NotNil(t, lc)
tokens, err := lc.RedeemTokens(context.Background())
assert.NoError(t, err)
assert.NotNil(t, tokens)
require.NotNil(t, tokens)
assert.NotEmpty(t, tokens.AccessToken)
assert.NotEmpty(t, tokens.RefreshToken)
@@ -122,82 +33,111 @@ func TestLoginCallback_RedeemTokens(t *testing.T) {
assert.True(t, tokens.Expiry.Before(time.Now().Add(time.Hour)))
})
t.Run("invalid code", func(t *testing.T) {
t.Run("invalid state", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?state=some-other-state"
_, err := newLoginCallback(t, url, nil)
assert.ErrorIs(t, err, client.ErrCallbackInvalidState)
})
t.Run("missing state", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback"
_, err := newLoginCallback(t, url, nil)
assert.ErrorIs(t, err, client.ErrCallbackInvalidState)
})
t.Run("identity provider error", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?error=invalid_client&error_description=client%20authenticaion%20failed"
idp := mock.NewIdentityProvider(mock.Config())
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
require.NoError(t, err)
require.NotNil(t, lc)
idp.ProviderHandler.Codes = map[string]*mock.AuthorizeRequest{
"some-other-code": {},
"another-code": {},
}
_, err := newLoginCallback(t, url, nil)
assert.ErrorIs(t, err, client.ErrCallbackIdentityProvider)
})
tokens, err := lc.RedeemTokens(context.Background())
assert.Error(t, err)
t.Run("supports authorization response with iss parameter", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state&iss=https://some-issuer"
_, err := newLoginCallback(t, url, func(idp *mock.IdentityProvider) {
idp.OpenIDConfig.TestProvider.SetIssuer("https://some-issuer")
idp.OpenIDConfig.TestProvider.WithAuthorizationResponseIssParameterSupported()
})
assert.NoError(t, err)
})
t.Run("missing issuer", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state"
_, err := newLoginCallback(t, url, func(idp *mock.IdentityProvider) {
idp.OpenIDConfig.TestProvider.WithAuthorizationResponseIssParameterSupported()
})
assert.ErrorIs(t, err, client.ErrCallbackInvalidIssuer)
assert.ErrorContains(t, err, "missing issuer parameter")
})
t.Run("invalid issuer", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state&iss=https://wrong-issuer"
_, err := newLoginCallback(t, url, func(idp *mock.IdentityProvider) {
idp.OpenIDConfig.TestProvider.SetIssuer("https://some-issuer")
idp.OpenIDConfig.TestProvider.WithAuthorizationResponseIssParameterSupported()
})
assert.ErrorIs(t, err, client.ErrCallbackInvalidIssuer)
assert.ErrorContains(t, err, "issuer mismatch: expected \"https://some-issuer\", got \"https://wrong-issuer\"")
})
t.Run("invalid code", func(t *testing.T) {
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state"
tokens, err := newLoginCallback(t, url, func(idp *mock.IdentityProvider) {
idp.ProviderHandler.Codes = map[string]*mock.AuthorizeRequest{
"some-other-code": {},
"another-code": {},
}
})
assert.ErrorIs(t, err, client.ErrCallbackRedeemTokens)
assert.Nil(t, tokens)
})
t.Run("nonce mismatch", func(t *testing.T) {
idp := mock.NewIdentityProvider(mock.Config())
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
require.NoError(t, err)
require.NotNil(t, lc)
idp.ProviderHandler.Codes["some-code"].Nonce = "some-other-nonce"
tokens, err := lc.RedeemTokens(context.Background())
assert.Error(t, err)
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state"
tokens, err := newLoginCallback(t, url, func(idp *mock.IdentityProvider) {
idp.ProviderHandler.Codes["some-code"].Nonce = "some-other-nonce"
})
assert.ErrorIs(t, err, client.ErrCallbackRedeemTokens)
assert.Nil(t, tokens)
})
t.Run("redirect_uri mismatch", func(t *testing.T) {
idp := mock.NewIdentityProvider(mock.Config())
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
require.NoError(t, err)
require.NotNil(t, lc)
idp.ProviderHandler.Codes["some-code"].RedirectUri = "http://not-wonderwall/oauth2/callback"
tokens, err := lc.RedeemTokens(context.Background())
assert.Error(t, err)
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state"
tokens, err := newLoginCallback(t, url, func(idp *mock.IdentityProvider) {
idp.ProviderHandler.Codes["some-code"].RedirectUri = "http://not-wonderwall/oauth2/callback"
})
assert.ErrorIs(t, err, client.ErrCallbackRedeemTokens)
assert.Nil(t, tokens)
})
t.Run("unexpected audience", func(t *testing.T) {
idp := mock.NewIdentityProvider(mock.Config())
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
require.NoError(t, err)
require.NotNil(t, lc)
idp.Cfg.OpenID.ClientID = "new-client-id"
tokens, err := lc.RedeemTokens(context.Background())
assert.Error(t, err)
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state"
tokens, err := newLoginCallback(t, url, func(idp *mock.IdentityProvider) {
idp.Cfg.OpenID.ClientID = "new-client-id"
})
assert.ErrorIs(t, err, client.ErrCallbackRedeemTokens)
assert.Nil(t, tokens)
})
t.Run("invalid acr", func(t *testing.T) {
idp := mock.NewIdentityProvider(mock.Config())
defer idp.Close()
lc, err := newLoginCallback(t, idp, url)
require.NoError(t, err)
require.NotNil(t, lc)
idp.ProviderHandler.Codes["some-code"].AcrLevel = "some-invalid-acr"
tokens, err := lc.RedeemTokens(context.Background())
assert.Error(t, err)
url := mock.Ingress + "/oauth2/callback?code=some-code&state=some-state"
tokens, err := newLoginCallback(t, url, func(idp *mock.IdentityProvider) {
idp.ProviderHandler.Codes["some-code"].AcrLevel = "some-invalid-acr"
})
assert.ErrorIs(t, err, client.ErrCallbackRedeemTokens)
assert.ErrorContains(t, err, "invalid acr: got \"some-invalid-acr\", expected \"some-acr\"")
assert.Nil(t, tokens)
})
}
func newLoginCallback(t *testing.T, idp *mock.IdentityProvider, url string) (*client.LoginCallback, error) {
func newLoginCallback(t *testing.T, url string, mutateFn func(*mock.IdentityProvider)) (*openid.Tokens, error) {
cfg := mock.Config()
idp := mock.NewIdentityProvider(cfg)
defer idp.Close()
req := idp.GetRequest(url)
redirect, err := urlpkg.LoginCallback(req)
assert.NoError(t, err)
@@ -212,6 +152,10 @@ func newLoginCallback(t *testing.T, idp *mock.IdentityProvider, url string) (*cl
},
}
if mutateFn != nil {
mutateFn(idp)
}
cookie := &openid.LoginCookie{
Acr: "some-acr",
State: "some-state",