From a528a70d74d3535e795c9e9b21cfe35adbd5ba21 Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Fri, 17 Jun 2022 17:05:47 +0100 Subject: [PATCH 01/13] Add support of OpenID providers With OpenID flow, instead of using /userinfo endpoint, an ID token issued by the authorisation server is used. Information in this token ususally includes extra params and options, not available in userinfo response. --- auth.go | 9 +++ go.mod | 2 + go.sum | 4 ++ provider/dev_provider.go | 113 ++++++++++++++++++++++++++++++++++--- provider/oauth2.go | 119 +++++++++++++++++++++++++++++++-------- provider/oauth2_test.go | 3 +- provider/openid_test.go | 90 +++++++++++++++++++++++++++++ 7 files changed, 307 insertions(+), 33 deletions(-) create mode 100644 provider/openid_test.go diff --git a/auth.go b/auth.go index d41c6eeb..b32fa25c 100644 --- a/auth.go +++ b/auth.go @@ -218,9 +218,16 @@ func (s *Service) Middleware() middleware.Authenticator { return s.authMiddleware } +func (s *Service) AddOpenIDProvider(name, cid, csecret string) { + s.addProvider(name, cid, csecret, true) +} + // AddProvider adds provider for given name func (s *Service) AddProvider(name, cid, csecret string) { + s.addProvider(name, cid, csecret, false) +} +func (s *Service) addProvider(name, cid, csecret string, useOpenID bool) { p := provider.Params{ URL: s.opts.URL, JwtService: s.jwtService, @@ -229,6 +236,7 @@ func (s *Service) AddProvider(name, cid, csecret string) { Cid: cid, Csecret: csecret, L: s.logger, + UseOpenID: useOpenID, } switch strings.ToLower(name) { @@ -266,6 +274,7 @@ func (s *Service) AddDevProvider(port int) { AvatarSaver: s.avatarProxy, L: s.logger, Port: port, + UseOpenID: true, } s.providers = append(s.providers, provider.NewService(provider.NewDev(p))) } diff --git a/go.mod b/go.mod index 869e836d..d9998f2d 100644 --- a/go.mod +++ b/go.mod @@ -19,9 +19,11 @@ require ( require ( cloud.google.com/go/compute v1.6.1 // indirect + github.com/MicahParks/keyfunc v1.1.0 // indirect github.com/aymerick/douceur v0.2.0 // indirect github.com/davecgh/go-spew v1.1.1 // indirect github.com/go-stack/stack v1.8.1 // indirect + github.com/golang-jwt/jwt/v4 v4.4.1 // indirect github.com/golang/protobuf v1.5.2 // indirect github.com/golang/snappy v0.0.4 // indirect github.com/google/uuid v1.1.2 // indirect diff --git a/go.sum b/go.sum index 8c21340d..04074ce2 100644 --- a/go.sum +++ b/go.sum @@ -54,6 +54,8 @@ cloud.google.com/go/storage v1.10.0/go.mod h1:FLPqc6j+Ki4BU591ie1oL6qBQGu2Bl/tZ9 dmitri.shuralyov.com/gpu/mtl v0.0.0-20190408044501-666a987793e9/go.mod h1:H6x//7gZCb22OMCxBHrMx7a5I7Hp++hsVxbQ4BYO7hU= github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= github.com/BurntSushi/xgb v0.0.0-20160522181843-27f122750802/go.mod h1:IVnqGOEym/WlBOVXweHU+Q+/VP0lqqI8lqeDx9IjBqo= +github.com/MicahParks/keyfunc v1.1.0 h1:9NcnRwS0ciuVeVNi+vTdYVMTmk62OID7VlG6y9BgLK0= +github.com/MicahParks/keyfunc v1.1.0/go.mod h1:a4yfunv77gZ0RgTNw7tOYS+bjtHk5565e+1dPz+YJI8= github.com/OneOfOne/xxhash v1.2.2/go.mod h1:HSdplMjZKSmBqAxg5vPj2TmRDmfkzw+cTzAElWljhcU= github.com/ajg/form v1.5.1 h1:t9c7v8JUKu/XxOGBU0yjNpaMloxGEJhUkqFRq0ibGeU= github.com/ajg/form v1.5.1/go.mod h1:uL1WgH+h2mgNtvBq0339dVnzXdBETtL2LeUXaIv25UY= @@ -115,6 +117,8 @@ github.com/go-stack/stack v1.8.1/go.mod h1:dcoOX6HbPZSZptuspn9bctJ+N/CnF5gGygcUP github.com/golang-jwt/jwt v3.2.1+incompatible/go.mod h1:8pz2t5EyA70fFQQSrl6XZXzqecmYZeUEB8OUGHkxJ+I= github.com/golang-jwt/jwt v3.2.2+incompatible h1:IfV12K8xAKAnZqdXVzCZ+TOjboZ2keLg81eXfW3O+oY= github.com/golang-jwt/jwt v3.2.2+incompatible/go.mod h1:8pz2t5EyA70fFQQSrl6XZXzqecmYZeUEB8OUGHkxJ+I= +github.com/golang-jwt/jwt/v4 v4.4.1 h1:pC5DB52sCeK48Wlb9oPcdhnjkz1TKt1D/P7WKJ0kUcQ= +github.com/golang-jwt/jwt/v4 v4.4.1/go.mod h1:m21LjoU+eqJr34lmDMbreY2eSTRJ1cv77w39/MY0Ch0= github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q= github.com/golang/groupcache v0.0.0-20190702054246-869f871628b6/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc= github.com/golang/groupcache v0.0.0-20191227052852-215e87163ea7/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc= diff --git a/provider/dev_provider.go b/provider/dev_provider.go index b3d2b769..94cb01f9 100644 --- a/provider/dev_provider.go +++ b/provider/dev_provider.go @@ -2,8 +2,14 @@ package provider import ( "context" + "crypto/rand" + "crypto/rsa" + "encoding/base64" + "encoding/json" "fmt" + "github.com/golang-jwt/jwt" "html/template" + "math/big" "net/http" "strings" "sync" @@ -25,9 +31,10 @@ const defDevAuthPort = 8084 // desired user name, this is the mode used for development. Non-interactive mode for tests only. type DevAuthServer struct { logger.L - Provider Oauth2Handler - Automatic bool - GetEmailFn func(string) string + Provider Oauth2Handler + Automatic bool + GetEmailFn func(string) string + CustomizeIdTokenFn func(map[string]interface{}) map[string]interface{} username string // unsafe, but fine for dev httpServer *http.Server @@ -50,6 +57,12 @@ func (d *DevAuthServer) Run(ctx context.Context) { //nolint (gocyclo) return } + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + d.Logf("[ERROR] failed to generate keys") + return + } + d.httpServer = &http.Server{ Addr: fmt.Sprintf(":%d", d.Provider.Port), Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -74,26 +87,97 @@ func (d *DevAuthServer) Run(ctx context.Context) { //nolint (gocyclo) } state := r.URL.Query().Get("state") - callbackURL := fmt.Sprintf("%s?code=g0ZGZmNjVmOWI&state=%s", d.Provider.conf.RedirectURL, state) + redirectURI := r.URL.Query().Get("redirect_uri") + callbackURL := fmt.Sprintf("%s?code=g0ZGZmNjVmOWI&state=%s", redirectURI, state) d.Logf("[DEBUG] callback url=%s", callbackURL) w.Header().Add("Location", callbackURL) w.WriteHeader(http.StatusFound) case strings.HasPrefix(r.URL.Path, "/login/oauth/access_token"): - res := `{ + email := d.username + if d.GetEmailFn != nil { + email = d.GetEmailFn(d.username) + } + + idClaims := map[string]interface{}{ + // required OpenID claims + "iss": "dev-auth", + "sub": "%s", + "aud": "client-id", + "iat": time.Now().Unix(), + "exp": time.Now().Add(1 * time.Hour).Unix(), + + // optional OpenID claims + "picture": fmt.Sprintf("http://127.0.0.1:%d/avatar?user=%s", d.Provider.Port, d.username), + "given_name": d.username, + "email": email, + } + + if d.CustomizeIdTokenFn != nil { + idClaims = d.CustomizeIdTokenFn(idClaims) + } + + tk := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims(idClaims)) + tk.Header["kid"] = "dev-auth-key-1" + signedTk, err := tk.SignedString(privateKey) + if err != nil { + d.Logf("[ERROR] failed to sign ID token") + w.WriteHeader(http.StatusInternalServerError) + return + } + + res := fmt.Sprintf(`{ "access_token":"MTQ0NjJkZmQ5OTM2NDE1ZTZjNGZmZjI3", + "id_token": "%s", "token_type":"bearer", "expires_in":3600, "refresh_token":"IwOGYzYTlmM2YxOTQ5MGE3YmNmMDFkNTVk", "scope":"create", "state":"12345678" - }` + }`, signedTk) + w.Header().Set("Content-Type", "application/json; charset=utf-8") if _, err = w.Write([]byte(res)); err != nil { w.WriteHeader(http.StatusInternalServerError) return } + case strings.HasPrefix(r.URL.Path, "/jwks"): + type jwkKey struct { + Kty string `json:"kty"` + N string `json:"n"` + E string `json:"e"` + Alg string `json:"alg"` + Kid string `json:"kid"` + } + + e := big.NewInt(int64(privateKey.E)) + key := jwkKey{ + Kty: "RSA", + Alg: "RS256", + Kid: "dev-auth-key-1", + N: base64.RawURLEncoding.EncodeToString(privateKey.N.Bytes()), + E: base64.RawURLEncoding.EncodeToString(e.Bytes()), + } + + jwks, err := json.Marshal(struct { + Keys []jwkKey `json:"keys"` + }{ + Keys: []jwkKey{key}, + }) + if err != nil { + d.Logf("[ERROR] failed to marshal jwks") + w.WriteHeader(http.StatusInternalServerError) + return + } + + w.WriteHeader(http.StatusOK) + wr, err := w.Write(jwks) + if err != nil || wr == 0 { + d.Logf("[ERROR] failed to write jwks") + w.WriteHeader(http.StatusInternalServerError) + } + case strings.HasPrefix(r.URL.Path, "/user"): ava := fmt.Sprintf("http://127.0.0.1:%d/avatar?user=%s", d.Provider.Port, d.username) res := fmt.Sprintf(`{ @@ -162,6 +246,10 @@ func (d *DevAuthServer) Shutdown() { d.lock.Unlock() } +func (d *DevAuthServer) ParseToken(value string) (claims token.Claims, err error) { + return d.Provider.JwtService.Parse(value) +} + // NewDev makes dev oauth2 provider for admin user func NewDev(p Params) Oauth2Handler { if p.Port == 0 { @@ -175,14 +263,23 @@ func NewDev(p Params) Oauth2Handler { }, scopes: []string{"user:email"}, infoURL: fmt.Sprintf("http://127.0.0.1:%d/user", p.Port), + jwksURL: fmt.Sprintf("http://127.0.0.1:%d/jwks", p.Port), mapUser: func(data UserData, _ []byte) token.User { - userInfo := token.User{ + if p.UseOpenID { + return token.User{ + ID: data.Value("sub"), + Name: data.Value("given_name"), + Picture: data.Value("picture"), + Email: data.Value("email"), + } + } + + return token.User{ ID: data.Value("id"), Name: data.Value("name"), Picture: data.Value("picture"), Email: data.Value("email"), } - return userInfo }, }) diff --git a/provider/oauth2.go b/provider/oauth2.go index d2c8dd7c..4542bdd1 100644 --- a/provider/oauth2.go +++ b/provider/oauth2.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "fmt" + "github.com/MicahParks/keyfunc" "io" "net/http" "strings" @@ -11,6 +12,7 @@ import ( "github.com/go-pkgz/rest" "github.com/golang-jwt/jwt" + jwtv4 "github.com/golang-jwt/jwt/v4" "golang.org/x/oauth2" "github.com/go-pkgz/auth/logger" @@ -24,10 +26,12 @@ type Oauth2Handler struct { // all of these fields specific to particular oauth2 provider name string infoURL string + jwksURL string endpoint oauth2.Endpoint scopes []string mapUser func(UserData, []byte) token.User // map info from InfoURL to User conf oauth2.Config + keyfunc jwt.Keyfunc } // Params to make initialized and ready to use provider @@ -39,6 +43,7 @@ type Params struct { Csecret string Issuer string AvatarSaver AvatarSaver + UseOpenID bool // switch to OpenID flow instead of pure OAuth2, i.e. load userinfo from an ID token and do not use self-signed JWT tokens Port int // relevant for providers supporting port customization, for example dev oauth2 } @@ -69,8 +74,37 @@ func initOauth2Handler(p Params, service Oauth2Handler) Oauth2Handler { Endpoint: service.endpoint, } + if p.UseOpenID { + kf, err := keyfunc.Get(service.jwksURL, keyfunc.Options{ + Client: http.DefaultClient, + Ctx: context.Background(), + RefreshErrorHandler: func(err error) { + p.Logf("[WARN] failed to refresh jwks: %s", err) + }, + RefreshInterval: 1 * time.Hour, + RefreshRateLimit: 1 * time.Minute, + RefreshTimeout: 30 * time.Second, + RefreshUnknownKID: true, + }) + + if err != nil { + p.Logf("[ERROR] failed to init jwks: %s", err) + } + + service.keyfunc = func(t *jwt.Token) (interface{}, error) { + // only to pass kid across, to manage jwt v3 vs v4 compatibility + v4token := jwtv4.Token{ + Header: map[string]interface{}{ + "kid": t.Header["kid"], + }, + } + + return kf.Keyfunc(&v4token) + } + } + p.Logf("[DEBUG] created %s oauth2, id=%s, redir=%s, endpoint=%s", - service.name, service.Cid, service.makeRedirURL("/{route}/"+service.name+"/"), service.endpoint) + service.name, service.Cid, service.makeRedirURLFromPath("/{route}/"+service.name+"/"), service.endpoint) return service } @@ -121,7 +155,7 @@ func (p Oauth2Handler) LoginHandler(w http.ResponseWriter, r *http.Request) { // setting RedirectURL to rootURL/routingPath/provider/callback // e.g. http://localhost:8080/auth/github/callback - p.conf.RedirectURL = p.makeRedirURL(r.URL.Path) + p.conf.RedirectURL = p.makeRedirURL(r) // return login url loginURL := p.conf.AuthCodeURL(state) @@ -150,7 +184,7 @@ func (p Oauth2Handler) AuthHandler(w http.ResponseWriter, r *http.Request) { return } - p.conf.RedirectURL = p.makeRedirURL(r.URL.Path) + p.conf.RedirectURL = p.makeRedirURL(r) p.Logf("[DEBUG] token with state %s", retrievedState) tok, err := p.conf.Exchange(context.Background(), r.URL.Query().Get("code")) @@ -160,32 +194,59 @@ func (p Oauth2Handler) AuthHandler(w http.ResponseWriter, r *http.Request) { } client := p.conf.Client(context.Background(), tok) - uinfo, err := client.Get(p.infoURL) - if err != nil { - rest.SendErrorJSON(w, r, p.L, http.StatusServiceUnavailable, err, "failed to get client info") - return - } - defer func() { - if e := uinfo.Body.Close(); e != nil { - p.Logf("[WARN] failed to close response body, %s", e) + var u token.User + if p.UseOpenID { + idToken, ok := tok.Extra("id_token").(string) + if !ok || idToken == "" { + rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, nil, "id_token is empty") + return } - }() - data, err := io.ReadAll(uinfo.Body) - if err != nil { - rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to read user info") - return + claims := jwt.MapClaims{} + parsedIDToken, err := jwt.ParseWithClaims(idToken, &claims, p.keyfunc) + if err != nil { + rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to parse id token") + return + } + + if !parsedIDToken.Valid { + rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "invalid id token") + return + } + + u = p.mapUser(UserData(claims), []byte(idToken)) } - jData := map[string]interface{}{} - if e := json.Unmarshal(data, &jData); e != nil { - rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to unmarshal user info") - return + if !p.UseOpenID { + uinfo, err := client.Get(p.infoURL) + if err != nil { + rest.SendErrorJSON(w, r, p.L, http.StatusServiceUnavailable, err, "failed to get client info") + return + } + + defer func() { + if e := uinfo.Body.Close(); e != nil { + p.Logf("[WARN] failed to close response body, %s", e) + } + }() + + data, err := io.ReadAll(uinfo.Body) + if err != nil { + rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to read user info") + return + } + + jData := map[string]interface{}{} + if e := json.Unmarshal(data, &jData); e != nil { + rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to unmarshal user info") + return + } + p.Logf("[DEBUG] got raw user info %+v", jData) + + u = p.mapUser(jData, data) } - p.Logf("[DEBUG] got raw user info %+v", jData) - u := p.mapUser(jData, data) if oauthClaims.NoAva { u.Picture = "" // reset picture on no avatar request } @@ -235,7 +296,19 @@ func (p Oauth2Handler) LogoutHandler(w http.ResponseWriter, r *http.Request) { p.JwtService.Reset(w) } -func (p Oauth2Handler) makeRedirURL(path string) string { +func (p Oauth2Handler) makeRedirURL(r *http.Request) string { + host := p.URL + if host == "" { // Base URL is not configured, use one from the request + host = "http://" + r.Host + } + + elems := strings.Split(r.URL.Path, "/") + newPath := strings.Join(elems[:len(elems)-1], "/") + + return strings.TrimSuffix(host, "/") + strings.TrimSuffix(newPath, "/") + urlCallbackSuffix +} + +func (p Oauth2Handler) makeRedirURLFromPath(path string) string { elems := strings.Split(path, "/") newPath := strings.Join(elems[:len(elems)-1], "/") diff --git a/provider/oauth2_test.go b/provider/oauth2_test.go index 42fe5c72..6ecb6afb 100644 --- a/provider/oauth2_test.go +++ b/provider/oauth2_test.go @@ -213,12 +213,11 @@ func TestMakeRedirURL(t *testing.T) { for i := range cases { c := cases[i] oh := initOauth2Handler(Params{URL: c.rootURL}, Oauth2Handler{}) - assert.Equal(t, c.out, oh.makeRedirURL(c.route)) + assert.Equal(t, c.out, oh.makeRedirURLFromPath(c.route)) } } func prepOauth2Test(t *testing.T, loginPort, authPort int) func() { - provider := Oauth2Handler{ name: "mock", endpoint: oauth2.Endpoint{ diff --git a/provider/openid_test.go b/provider/openid_test.go new file mode 100644 index 00000000..2a66a762 --- /dev/null +++ b/provider/openid_test.go @@ -0,0 +1,90 @@ +package provider_test + +import ( + "context" + "fmt" + "github.com/go-pkgz/auth" + "github.com/go-pkgz/auth/avatar" + "github.com/go-pkgz/auth/logger" + "github.com/go-pkgz/auth/provider" + "github.com/go-pkgz/auth/token" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "math/rand" + "net/http" + "net/http/cookiejar" + "net/http/httptest" + "testing" + "time" +) + +func TestNewOpenID(t *testing.T) { + rand.Seed(time.Now().UnixNano()) + devPort := rand.Intn(10_000) + 50_000 + expectedTestUserSub := fmt.Sprintf("test-user-%d", devPort) + svc := auth.NewService(auth.Opts{ + SecretReader: token.SecretFunc(func(aud string) (string, error) { + return "some-signing-key", nil + }), + Logger: logger.Std, + AvatarStore: avatar.NewNoOp(), + }) + + devParams := provider.Params{ + L: logger.Std, + URL: fmt.Sprintf("http://localhost:%d", devPort), + JwtService: svc.TokenService(), + Cid: "client-id", + Csecret: "client-secret", + Issuer: "test-issuer", + AvatarSaver: svc.AvatarProxy(), + UseOpenID: true, + Port: devPort, + } + + dev := provider.NewDev(devParams) + devAuth := &provider.DevAuthServer{Provider: dev, L: logger.Std} + devAuth.Automatic = true + devAuth.CustomizeIdTokenFn = func(m map[string]interface{}) map[string]interface{} { + m["sub"] = expectedTestUserSub + return m + } + + go devAuth.Run(context.Background()) + defer devAuth.Shutdown() + + time.Sleep(300 * time.Millisecond) + + svc.AddDevProvider(devPort) + + authHandler, _ := svc.Handlers() + server := httptest.NewServer(authHandler) + defer server.Close() + + jar, err := cookiejar.New(nil) + require.NoError(t, err) + + client := http.Client{ + Jar: jar, + } + + resp, err := client.Get(server.URL + "/auth/dev/login") + require.NoError(t, err) + assert.Equal(t, http.StatusOK, resp.StatusCode) + + cookies := resp.Cookies() + var cookie *http.Cookie + for _, c := range cookies { + if c.Name == "JWT" { + cookie = c + break + } + } + + require.NotNil(t, cookie) + claims, err := devAuth.ParseToken(cookie.Value) + require.NoError(t, err) + + // check user details are from the ID token + assert.Equal(t, expectedTestUserSub, claims.User.ID) +} From 242fcfc190ac9739b9882b0ae1873c37f823ce2f Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Tue, 21 Jun 2022 11:11:31 +0100 Subject: [PATCH 02/13] Allow to add custom OpenID providers --- auth.go | 17 +++++++++++++++++ provider/custom_server.go | 2 ++ 2 files changed, 19 insertions(+) diff --git a/auth.go b/auth.go index b32fa25c..154a3b54 100644 --- a/auth.go +++ b/auth.go @@ -315,6 +315,23 @@ func (s *Service) AddCustomProvider(name string, client Client, copts provider.C s.authMiddleware.Providers = s.providers } +// AddCustomOpenIDProvider adds custom provider (e.g. https://gopkg.in/oauth2.v3) that uses OpenID instead of pure OAuth2 +func (s *Service) AddCustomOpenIDProvider(name string, client Client, copts provider.CustomHandlerOpt) { + p := provider.Params{ + URL: s.opts.URL, + JwtService: s.jwtService, + Issuer: s.issuer, + AvatarSaver: s.avatarProxy, + Cid: client.Cid, + Csecret: client.Csecret, + L: s.logger, + UseOpenID: true, + } + + s.providers = append(s.providers, provider.NewService(provider.NewCustom(name, p, copts))) + s.authMiddleware.Providers = s.providers +} + // AddDirectProvider adds provider with direct check against data store // it doesn't do any handshake and uses provided credChecker to verify user and password from the request func (s *Service) AddDirectProvider(name string, credChecker provider.CredChecker) { diff --git a/provider/custom_server.go b/provider/custom_server.go index f5bde31c..9a07f2bd 100644 --- a/provider/custom_server.go +++ b/provider/custom_server.go @@ -24,6 +24,7 @@ import ( type CustomHandlerOpt struct { Endpoint oauth2.Endpoint InfoURL string + JwksURL string MapUserFn func(UserData, []byte) token.User Scopes []string } @@ -208,6 +209,7 @@ func NewCustom(name string, p Params, copts CustomHandlerOpt) Oauth2Handler { endpoint: copts.Endpoint, scopes: copts.Scopes, infoURL: copts.InfoURL, + jwksURL: copts.JwksURL, mapUser: copts.MapUserFn, }) } From e3fd05411c97c3b4145bd3d898d82a95f6d40777 Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Tue, 21 Jun 2022 16:10:10 +0100 Subject: [PATCH 03/13] Refactor keys loading to retry on token request --- auth.go | 13 +++++++ provider/dev_provider.go | 77 +++++++++++++++++++++++----------------- provider/oauth2.go | 69 ++++++++++++++++++++++------------- provider/openid_test.go | 25 ++++--------- 4 files changed, 109 insertions(+), 75 deletions(-) diff --git a/auth.go b/auth.go index 154a3b54..97ccb56e 100644 --- a/auth.go +++ b/auth.go @@ -267,6 +267,19 @@ func (s *Service) addProvider(name, cid, csecret string, useOpenID bool) { // AddDevProvider with a custom port func (s *Service) AddDevProvider(port int) { + p := provider.Params{ + URL: s.opts.URL, + JwtService: s.jwtService, + Issuer: s.issuer, + AvatarSaver: s.avatarProxy, + L: s.logger, + Port: port, + } + s.providers = append(s.providers, provider.NewService(provider.NewDev(p))) +} + +// AddDevOpenIDProvider with a custom port that is service and using OpenID tokens +func (s *Service) AddDevOpenIDProvider(port int) { p := provider.Params{ URL: s.opts.URL, JwtService: s.jwtService, diff --git a/provider/dev_provider.go b/provider/dev_provider.go index 94cb01f9..b4e1ea1f 100644 --- a/provider/dev_provider.go +++ b/provider/dev_provider.go @@ -99,42 +99,53 @@ func (d *DevAuthServer) Run(ctx context.Context) { //nolint (gocyclo) email = d.GetEmailFn(d.username) } - idClaims := map[string]interface{}{ - // required OpenID claims - "iss": "dev-auth", - "sub": "%s", - "aud": "client-id", - "iat": time.Now().Unix(), - "exp": time.Now().Add(1 * time.Hour).Unix(), - - // optional OpenID claims - "picture": fmt.Sprintf("http://127.0.0.1:%d/avatar?user=%s", d.Provider.Port, d.username), - "given_name": d.username, - "email": email, - } - - if d.CustomizeIdTokenFn != nil { - idClaims = d.CustomizeIdTokenFn(idClaims) - } - - tk := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims(idClaims)) - tk.Header["kid"] = "dev-auth-key-1" - signedTk, err := tk.SignedString(privateKey) - if err != nil { - d.Logf("[ERROR] failed to sign ID token") - w.WriteHeader(http.StatusInternalServerError) - return - } - - res := fmt.Sprintf(`{ + res := `{ "access_token":"MTQ0NjJkZmQ5OTM2NDE1ZTZjNGZmZjI3", - "id_token": "%s", "token_type":"bearer", "expires_in":3600, "refresh_token":"IwOGYzYTlmM2YxOTQ5MGE3YmNmMDFkNTVk", "scope":"create", "state":"12345678" - }`, signedTk) + }` + + if d.Provider.UseOpenID { + idClaims := map[string]interface{}{ + // required OpenID claims + "iss": "dev-auth", + "sub": "%s", + "aud": "client-id", + "iat": time.Now().Unix(), + "exp": time.Now().Add(1 * time.Hour).Unix(), + + // optional OpenID claims + "picture": fmt.Sprintf("http://127.0.0.1:%d/avatar?user=%s", d.Provider.Port, d.username), + "given_name": d.username, + "email": email, + } + + if d.CustomizeIdTokenFn != nil { + idClaims = d.CustomizeIdTokenFn(idClaims) + } + + tk := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims(idClaims)) + tk.Header["kid"] = "dev-auth-key-1" + signedTk, err := tk.SignedString(privateKey) + if err != nil { + d.Logf("[ERROR] failed to sign ID token") + w.WriteHeader(http.StatusInternalServerError) + return + } + + res = fmt.Sprintf(`{ + "access_token":"MTQ0NjJkZmQ5OTM2NDE1ZTZjNGZmZjI3", + "id_token": "%s", + "token_type":"bearer", + "expires_in":3600, + "refresh_token":"IwOGYzYTlmM2YxOTQ5MGE3YmNmMDFkNTVk", + "scope":"create", + "state":"12345678" + }`, signedTk) + } w.Header().Set("Content-Type", "application/json; charset=utf-8") if _, err = w.Write([]byte(res)); err != nil { @@ -142,7 +153,7 @@ func (d *DevAuthServer) Run(ctx context.Context) { //nolint (gocyclo) return } - case strings.HasPrefix(r.URL.Path, "/jwks"): + case strings.HasPrefix(r.URL.Path, "/jwks") && d.Provider.UseOpenID: type jwkKey struct { Kty string `json:"kty"` N string `json:"n"` @@ -274,12 +285,14 @@ func NewDev(p Params) Oauth2Handler { } } - return token.User{ + userInfo := token.User{ ID: data.Value("id"), Name: data.Value("name"), Picture: data.Value("picture"), Email: data.Value("email"), } + + return userInfo }, }) diff --git a/provider/oauth2.go b/provider/oauth2.go index 4542bdd1..fb2bde40 100644 --- a/provider/oauth2.go +++ b/provider/oauth2.go @@ -8,6 +8,7 @@ import ( "io" "net/http" "strings" + "sync" "time" "github.com/go-pkgz/rest" @@ -32,6 +33,7 @@ type Oauth2Handler struct { mapUser func(UserData, []byte) token.User // map info from InfoURL to User conf oauth2.Config keyfunc jwt.Keyfunc + kfLock *sync.Mutex } // Params to make initialized and ready to use provider @@ -75,31 +77,10 @@ func initOauth2Handler(p Params, service Oauth2Handler) Oauth2Handler { } if p.UseOpenID { - kf, err := keyfunc.Get(service.jwksURL, keyfunc.Options{ - Client: http.DefaultClient, - Ctx: context.Background(), - RefreshErrorHandler: func(err error) { - p.Logf("[WARN] failed to refresh jwks: %s", err) - }, - RefreshInterval: 1 * time.Hour, - RefreshRateLimit: 1 * time.Minute, - RefreshTimeout: 30 * time.Second, - RefreshUnknownKID: true, - }) - + service.kfLock = &sync.Mutex{} + err := service.tryInitJWKSKeyfunc() if err != nil { - p.Logf("[ERROR] failed to init jwks: %s", err) - } - - service.keyfunc = func(t *jwt.Token) (interface{}, error) { - // only to pass kid across, to manage jwt v3 vs v4 compatibility - v4token := jwtv4.Token{ - Header: map[string]interface{}{ - "kid": t.Header["kid"], - }, - } - - return kf.Keyfunc(&v4token) + p.Logf("[ERROR] failed to load JWT keys to enable OpenID, will retry on token request: %s", err) } } @@ -203,6 +184,14 @@ func (p Oauth2Handler) AuthHandler(w http.ResponseWriter, r *http.Request) { return } + if p.keyfunc == nil { + err = p.tryInitJWKSKeyfunc() + if err != nil { + rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, nil, "can't load JWKS keys") + return + } + } + claims := jwt.MapClaims{} parsedIDToken, err := jwt.ParseWithClaims(idToken, &claims, p.keyfunc) if err != nil { @@ -314,3 +303,35 @@ func (p Oauth2Handler) makeRedirURLFromPath(path string) string { return strings.TrimSuffix(p.URL, "/") + strings.TrimSuffix(newPath, "/") + urlCallbackSuffix } + +func (p *Oauth2Handler) tryInitJWKSKeyfunc() error { + p.kfLock.Lock() + defer p.kfLock.Unlock() + if p.keyfunc != nil { + return nil + } + + kf, err := keyfunc.Get(p.jwksURL, keyfunc.Options{ + Client: http.DefaultClient, + Ctx: context.Background(), + RefreshUnknownKID: true, // to support key rotation, re-load keys if KID is unknown + RefreshRateLimit: 1 * time.Minute, // but no often than once per minute + }) + + if err != nil { + return err + } + + p.keyfunc = func(t *jwt.Token) (interface{}, error) { + // only to pass kid across, to manage jwt v3 vs v4 compatibility + v4token := jwtv4.Token{ + Header: map[string]interface{}{ + "kid": t.Header["kid"], + }, + } + + return kf.Keyfunc(&v4token) + } + + return nil +} diff --git a/provider/openid_test.go b/provider/openid_test.go index 2a66a762..c5a672a3 100644 --- a/provider/openid_test.go +++ b/provider/openid_test.go @@ -6,7 +6,6 @@ import ( "github.com/go-pkgz/auth" "github.com/go-pkgz/auth/avatar" "github.com/go-pkgz/auth/logger" - "github.com/go-pkgz/auth/provider" "github.com/go-pkgz/auth/token" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -30,20 +29,10 @@ func TestNewOpenID(t *testing.T) { AvatarStore: avatar.NewNoOp(), }) - devParams := provider.Params{ - L: logger.Std, - URL: fmt.Sprintf("http://localhost:%d", devPort), - JwtService: svc.TokenService(), - Cid: "client-id", - Csecret: "client-secret", - Issuer: "test-issuer", - AvatarSaver: svc.AvatarProxy(), - UseOpenID: true, - Port: devPort, - } + svc.AddDevOpenIDProvider(devPort) + devAuth, err := svc.DevAuth() + require.NoError(t, err) - dev := provider.NewDev(devParams) - devAuth := &provider.DevAuthServer{Provider: dev, L: logger.Std} devAuth.Automatic = true devAuth.CustomizeIdTokenFn = func(m map[string]interface{}) map[string]interface{} { m["sub"] = expectedTestUserSub @@ -53,10 +42,6 @@ func TestNewOpenID(t *testing.T) { go devAuth.Run(context.Background()) defer devAuth.Shutdown() - time.Sleep(300 * time.Millisecond) - - svc.AddDevProvider(devPort) - authHandler, _ := svc.Handlers() server := httptest.NewServer(authHandler) defer server.Close() @@ -68,9 +53,11 @@ func TestNewOpenID(t *testing.T) { Jar: jar, } + time.Sleep(200 * time.Millisecond) + resp, err := client.Get(server.URL + "/auth/dev/login") require.NoError(t, err) - assert.Equal(t, http.StatusOK, resp.StatusCode) + require.Equal(t, http.StatusOK, resp.StatusCode) cookies := resp.Cookies() var cookie *http.Cookie From 89969d4b853d46ee21f4bd70e31cec35201be3d1 Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Tue, 21 Jun 2022 16:26:04 +0100 Subject: [PATCH 04/13] Only support OpenID for custom providers --- auth.go | 17 ++++++----------- 1 file changed, 6 insertions(+), 11 deletions(-) diff --git a/auth.go b/auth.go index 97ccb56e..75df4c6e 100644 --- a/auth.go +++ b/auth.go @@ -218,16 +218,8 @@ func (s *Service) Middleware() middleware.Authenticator { return s.authMiddleware } -func (s *Service) AddOpenIDProvider(name, cid, csecret string) { - s.addProvider(name, cid, csecret, true) -} - // AddProvider adds provider for given name func (s *Service) AddProvider(name, cid, csecret string) { - s.addProvider(name, cid, csecret, false) -} - -func (s *Service) addProvider(name, cid, csecret string, useOpenID bool) { p := provider.Params{ URL: s.opts.URL, JwtService: s.jwtService, @@ -236,7 +228,6 @@ func (s *Service) addProvider(name, cid, csecret string, useOpenID bool) { Cid: cid, Csecret: csecret, L: s.logger, - UseOpenID: useOpenID, } switch strings.ToLower(name) { @@ -328,8 +319,8 @@ func (s *Service) AddCustomProvider(name string, client Client, copts provider.C s.authMiddleware.Providers = s.providers } -// AddCustomOpenIDProvider adds custom provider (e.g. https://gopkg.in/oauth2.v3) that uses OpenID instead of pure OAuth2 -func (s *Service) AddCustomOpenIDProvider(name string, client Client, copts provider.CustomHandlerOpt) { +// AddOpenIDProvider adds custom provider (e.g. https://gopkg.in/oauth2.v3) that uses OpenID instead of pure OAuth2 +func (s *Service) AddOpenIDProvider(name string, client Client, copts provider.CustomHandlerOpt) { p := provider.Params{ URL: s.opts.URL, JwtService: s.jwtService, @@ -341,6 +332,10 @@ func (s *Service) AddCustomOpenIDProvider(name string, client Client, copts prov UseOpenID: true, } + if copts.Scopes == nil { + copts.Scopes = []string{"openid"} + } + s.providers = append(s.providers, provider.NewService(provider.NewCustom(name, p, copts))) s.authMiddleware.Providers = s.providers } From 7411bff4fa67c9dae1658f01f5c0f762dfd0804f Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Tue, 21 Jun 2022 16:27:47 +0100 Subject: [PATCH 05/13] Fix typos --- auth.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/auth.go b/auth.go index 75df4c6e..149bd4bb 100644 --- a/auth.go +++ b/auth.go @@ -269,7 +269,7 @@ func (s *Service) AddDevProvider(port int) { s.providers = append(s.providers, provider.NewService(provider.NewDev(p))) } -// AddDevOpenIDProvider with a custom port that is service and using OpenID tokens +// AddDevOpenIDProvider with a custom port that is using OpenID tokens func (s *Service) AddDevOpenIDProvider(port int) { p := provider.Params{ URL: s.opts.URL, From 8026918ec7d78ea6d4b1f90d7b1d0f5c007a1e05 Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Tue, 21 Jun 2022 17:01:46 +0100 Subject: [PATCH 06/13] Cleanup makeRedirURL should work from a request, but it's not part of this PR --- provider/dev_provider.go | 14 +++++--------- provider/oauth2.go | 22 +++++----------------- provider/oauth2_test.go | 2 +- provider/openid_test.go | 28 +++++++++++++++++++++++++--- 4 files changed, 36 insertions(+), 30 deletions(-) diff --git a/provider/dev_provider.go b/provider/dev_provider.go index b4e1ea1f..ce47d234 100644 --- a/provider/dev_provider.go +++ b/provider/dev_provider.go @@ -94,11 +94,6 @@ func (d *DevAuthServer) Run(ctx context.Context) { //nolint (gocyclo) w.WriteHeader(http.StatusFound) case strings.HasPrefix(r.URL.Path, "/login/oauth/access_token"): - email := d.username - if d.GetEmailFn != nil { - email = d.GetEmailFn(d.username) - } - res := `{ "access_token":"MTQ0NjJkZmQ5OTM2NDE1ZTZjNGZmZjI3", "token_type":"bearer", @@ -109,6 +104,11 @@ func (d *DevAuthServer) Run(ctx context.Context) { //nolint (gocyclo) }` if d.Provider.UseOpenID { + email := d.username + if d.GetEmailFn != nil { + email = d.GetEmailFn(d.username) + } + idClaims := map[string]interface{}{ // required OpenID claims "iss": "dev-auth", @@ -257,10 +257,6 @@ func (d *DevAuthServer) Shutdown() { d.lock.Unlock() } -func (d *DevAuthServer) ParseToken(value string) (claims token.Claims, err error) { - return d.Provider.JwtService.Parse(value) -} - // NewDev makes dev oauth2 provider for admin user func NewDev(p Params) Oauth2Handler { if p.Port == 0 { diff --git a/provider/oauth2.go b/provider/oauth2.go index fb2bde40..b87394f5 100644 --- a/provider/oauth2.go +++ b/provider/oauth2.go @@ -45,7 +45,7 @@ type Params struct { Csecret string Issuer string AvatarSaver AvatarSaver - UseOpenID bool // switch to OpenID flow instead of pure OAuth2, i.e. load userinfo from an ID token and do not use self-signed JWT tokens + UseOpenID bool // switch to OpenID flow, load user from an ID token instead of userinfo Port int // relevant for providers supporting port customization, for example dev oauth2 } @@ -85,7 +85,7 @@ func initOauth2Handler(p Params, service Oauth2Handler) Oauth2Handler { } p.Logf("[DEBUG] created %s oauth2, id=%s, redir=%s, endpoint=%s", - service.name, service.Cid, service.makeRedirURLFromPath("/{route}/"+service.name+"/"), service.endpoint) + service.name, service.Cid, service.makeRedirURL("/{route}/"+service.name+"/"), service.endpoint) return service } @@ -136,7 +136,7 @@ func (p Oauth2Handler) LoginHandler(w http.ResponseWriter, r *http.Request) { // setting RedirectURL to rootURL/routingPath/provider/callback // e.g. http://localhost:8080/auth/github/callback - p.conf.RedirectURL = p.makeRedirURL(r) + p.conf.RedirectURL = p.makeRedirURL(r.URL.Path) // return login url loginURL := p.conf.AuthCodeURL(state) @@ -165,7 +165,7 @@ func (p Oauth2Handler) AuthHandler(w http.ResponseWriter, r *http.Request) { return } - p.conf.RedirectURL = p.makeRedirURL(r) + p.conf.RedirectURL = p.makeRedirURL(r.URL.Path) p.Logf("[DEBUG] token with state %s", retrievedState) tok, err := p.conf.Exchange(context.Background(), r.URL.Query().Get("code")) @@ -285,19 +285,7 @@ func (p Oauth2Handler) LogoutHandler(w http.ResponseWriter, r *http.Request) { p.JwtService.Reset(w) } -func (p Oauth2Handler) makeRedirURL(r *http.Request) string { - host := p.URL - if host == "" { // Base URL is not configured, use one from the request - host = "http://" + r.Host - } - - elems := strings.Split(r.URL.Path, "/") - newPath := strings.Join(elems[:len(elems)-1], "/") - - return strings.TrimSuffix(host, "/") + strings.TrimSuffix(newPath, "/") + urlCallbackSuffix -} - -func (p Oauth2Handler) makeRedirURLFromPath(path string) string { +func (p Oauth2Handler) makeRedirURL(path string) string { elems := strings.Split(path, "/") newPath := strings.Join(elems[:len(elems)-1], "/") diff --git a/provider/oauth2_test.go b/provider/oauth2_test.go index 6ecb6afb..c804f99a 100644 --- a/provider/oauth2_test.go +++ b/provider/oauth2_test.go @@ -213,7 +213,7 @@ func TestMakeRedirURL(t *testing.T) { for i := range cases { c := cases[i] oh := initOauth2Handler(Params{URL: c.rootURL}, Oauth2Handler{}) - assert.Equal(t, c.out, oh.makeRedirURLFromPath(c.route)) + assert.Equal(t, c.out, oh.makeRedirURL(c.route)) } } diff --git a/provider/openid_test.go b/provider/openid_test.go index c5a672a3..9860e762 100644 --- a/provider/openid_test.go +++ b/provider/openid_test.go @@ -10,6 +10,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "math/rand" + "net" "net/http" "net/http/cookiejar" "net/http/httptest" @@ -20,6 +21,8 @@ import ( func TestNewOpenID(t *testing.T) { rand.Seed(time.Now().UnixNano()) devPort := rand.Intn(10_000) + 50_000 + testSrvHost := fmt.Sprintf("127.0.0.1:%d", rand.Intn(10_000)+50_000) + expectedTestUserSub := fmt.Sprintf("test-user-%d", devPort) svc := auth.NewService(auth.Opts{ SecretReader: token.SecretFunc(func(aud string) (string, error) { @@ -27,6 +30,7 @@ func TestNewOpenID(t *testing.T) { }), Logger: logger.Std, AvatarStore: avatar.NewNoOp(), + URL: fmt.Sprintf("http://%s", testSrvHost), }) svc.AddDevOpenIDProvider(devPort) @@ -43,7 +47,11 @@ func TestNewOpenID(t *testing.T) { defer devAuth.Shutdown() authHandler, _ := svc.Handlers() - server := httptest.NewServer(authHandler) + server := httptest.NewUnstartedServer(authHandler) + server.Listener, err = net.Listen("tcp", testSrvHost) + require.NoError(t, err) + server.Start() + defer server.Close() jar, err := cookiejar.New(nil) @@ -53,7 +61,8 @@ func TestNewOpenID(t *testing.T) { Jar: jar, } - time.Sleep(200 * time.Millisecond) + require.NoError(t, waitFor(testSrvHost)) + require.NoError(t, waitFor(fmt.Sprintf("localhost:%d", devPort))) resp, err := client.Get(server.URL + "/auth/dev/login") require.NoError(t, err) @@ -69,9 +78,22 @@ func TestNewOpenID(t *testing.T) { } require.NotNil(t, cookie) - claims, err := devAuth.ParseToken(cookie.Value) + claims, err := devAuth.Provider.JwtService.Parse(cookie.Value) require.NoError(t, err) // check user details are from the ID token assert.Equal(t, expectedTestUserSub, claims.User.ID) } + +func waitFor(host string) error { + for i := 1; i < 20; i++ { + time.Sleep(time.Duration(i*10) * time.Millisecond) + + dial, err := net.Dial("tcp", host) + if err == nil { + return dial.Close() + } + } + + return fmt.Errorf("timeout waiting for %s", host) +} From e5a20c7a40a8c7bf3eacbbbc6504ac629bbcfccd Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Wed, 22 Jun 2022 09:09:45 +0100 Subject: [PATCH 07/13] Only generate private key if OpenID is enabled Key generation is slow(-ish) so usual sleeps of 50ms sometimes not enough, that makes tests flaky. --- provider/dev_provider.go | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/provider/dev_provider.go b/provider/dev_provider.go index ce47d234..edb2bba8 100644 --- a/provider/dev_provider.go +++ b/provider/dev_provider.go @@ -57,10 +57,13 @@ func (d *DevAuthServer) Run(ctx context.Context) { //nolint (gocyclo) return } - privateKey, err := rsa.GenerateKey(rand.Reader, 2048) - if err != nil { - d.Logf("[ERROR] failed to generate keys") - return + var privateKey *rsa.PrivateKey + if d.Provider.UseOpenID { + privateKey, err = rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + d.Logf("[ERROR] failed to generate keys") + return + } } d.httpServer = &http.Server{ From ab93f7d7c3d000e699647ca9a35231a405a4cdf2 Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Wed, 22 Jun 2022 09:38:28 +0100 Subject: [PATCH 08/13] Fix linter issues and refactor to reduce cyclomatic complexity --- provider/dev_provider.go | 19 ++++--- provider/oauth2.go | 120 +++++++++++++++++++++------------------ provider/openid_test.go | 16 +++--- 3 files changed, 82 insertions(+), 73 deletions(-) diff --git a/provider/dev_provider.go b/provider/dev_provider.go index edb2bba8..2249d32b 100644 --- a/provider/dev_provider.go +++ b/provider/dev_provider.go @@ -34,7 +34,7 @@ type DevAuthServer struct { Provider Oauth2Handler Automatic bool GetEmailFn func(string) string - CustomizeIdTokenFn func(map[string]interface{}) map[string]interface{} + CustomizeIDTokenFn func(map[string]interface{}) map[string]interface{} username string // unsafe, but fine for dev httpServer *http.Server @@ -126,14 +126,15 @@ func (d *DevAuthServer) Run(ctx context.Context) { //nolint (gocyclo) "email": email, } - if d.CustomizeIdTokenFn != nil { - idClaims = d.CustomizeIdTokenFn(idClaims) + if d.CustomizeIDTokenFn != nil { + idClaims = d.CustomizeIDTokenFn(idClaims) } tk := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims(idClaims)) tk.Header["kid"] = "dev-auth-key-1" - signedTk, err := tk.SignedString(privateKey) - if err != nil { + + signedTk, e := tk.SignedString(privateKey) + if e != nil { d.Logf("[ERROR] failed to sign ID token") w.WriteHeader(http.StatusInternalServerError) return @@ -174,20 +175,20 @@ func (d *DevAuthServer) Run(ctx context.Context) { //nolint (gocyclo) E: base64.RawURLEncoding.EncodeToString(e.Bytes()), } - jwks, err := json.Marshal(struct { + jwks, er := json.Marshal(struct { Keys []jwkKey `json:"keys"` }{ Keys: []jwkKey{key}, }) - if err != nil { + if er != nil { d.Logf("[ERROR] failed to marshal jwks") w.WriteHeader(http.StatusInternalServerError) return } w.WriteHeader(http.StatusOK) - wr, err := w.Write(jwks) - if err != nil || wr == 0 { + wr, er := w.Write(jwks) + if er != nil || wr == 0 { d.Logf("[ERROR] failed to write jwks") w.WriteHeader(http.StatusInternalServerError) } diff --git a/provider/oauth2.go b/provider/oauth2.go index b87394f5..b4eb1626 100644 --- a/provider/oauth2.go +++ b/provider/oauth2.go @@ -177,65 +177,23 @@ func (p Oauth2Handler) AuthHandler(w http.ResponseWriter, r *http.Request) { client := p.conf.Client(context.Background(), tok) var u token.User - if p.UseOpenID { - idToken, ok := tok.Extra("id_token").(string) - if !ok || idToken == "" { - rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, nil, "id_token is empty") - return - } - - if p.keyfunc == nil { - err = p.tryInitJWKSKeyfunc() - if err != nil { - rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, nil, "can't load JWKS keys") - return - } - } - - claims := jwt.MapClaims{} - parsedIDToken, err := jwt.ParseWithClaims(idToken, &claims, p.keyfunc) - if err != nil { - rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to parse id token") - return - } - - if !parsedIDToken.Valid { - rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "invalid id token") - return - } - - u = p.mapUser(UserData(claims), []byte(idToken)) + var userData UserData + var rawUserData []byte + + switch p.UseOpenID { + case true: + userData, rawUserData, err = p.loadUserFromIDToken(tok) + case false: + userData, rawUserData, err = p.loadUserFromEndpoint(client) } - if !p.UseOpenID { - uinfo, err := client.Get(p.infoURL) - if err != nil { - rest.SendErrorJSON(w, r, p.L, http.StatusServiceUnavailable, err, "failed to get client info") - return - } - - defer func() { - if e := uinfo.Body.Close(); e != nil { - p.Logf("[WARN] failed to close response body, %s", e) - } - }() - - data, err := io.ReadAll(uinfo.Body) - if err != nil { - rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to read user info") - return - } - - jData := map[string]interface{}{} - if e := json.Unmarshal(data, &jData); e != nil { - rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to unmarshal user info") - return - } - p.Logf("[DEBUG] got raw user info %+v", jData) - - u = p.mapUser(jData, data) + if err != nil { + rest.SendErrorJSON(w, r, p.L, http.StatusInternalServerError, err, "failed to load user data") + return } + u = p.mapUser(userData, rawUserData) + if oauthClaims.NoAva { u.Picture = "" // reset picture on no avatar request } @@ -285,6 +243,58 @@ func (p Oauth2Handler) LogoutHandler(w http.ResponseWriter, r *http.Request) { p.JwtService.Reset(w) } +func (p Oauth2Handler) loadUserFromIDToken(tok *oauth2.Token) (UserData, []byte, error) { + idToken, ok := tok.Extra("id_token").(string) + if !ok || idToken == "" { + return nil, nil, fmt.Errorf("id_token not found") + } + + if p.keyfunc == nil { + err := p.tryInitJWKSKeyfunc() + if err != nil { + return nil, nil, fmt.Errorf("can't load JWKS keys") + } + } + + claims := jwt.MapClaims{} + parsedIDToken, err := jwt.ParseWithClaims(idToken, &claims, p.keyfunc) + if err != nil { + return nil, nil, fmt.Errorf("failed to parse id token") + } + + if !parsedIDToken.Valid { + return nil, nil, fmt.Errorf("invalid id token") + } + + return UserData(claims), []byte(idToken), nil +} + +func (p Oauth2Handler) loadUserFromEndpoint(client *http.Client) (UserData, []byte, error) { + uinfo, err := client.Get(p.infoURL) + if err != nil { + return nil, nil, fmt.Errorf("failed to get client info") + } + + defer func() { + if e := uinfo.Body.Close(); e != nil { + p.Logf("[WARN] failed to close response body, %s", e) + } + }() + + data, err := io.ReadAll(uinfo.Body) + if err != nil { + return nil, nil, fmt.Errorf("failed to read user info") + } + + jData := map[string]interface{}{} + if e := json.Unmarshal(data, &jData); e != nil { + return nil, nil, fmt.Errorf("failed to unmarshal user info") + } + p.Logf("[DEBUG] got raw user info %+v", jData) + + return jData, data, nil +} + func (p Oauth2Handler) makeRedirURL(path string) string { elems := strings.Split(path, "/") newPath := strings.Join(elems[:len(elems)-1], "/") diff --git a/provider/openid_test.go b/provider/openid_test.go index 9860e762..762c1977 100644 --- a/provider/openid_test.go +++ b/provider/openid_test.go @@ -9,7 +9,6 @@ import ( "github.com/go-pkgz/auth/token" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "math/rand" "net" "net/http" "net/http/cookiejar" @@ -19,9 +18,8 @@ import ( ) func TestNewOpenID(t *testing.T) { - rand.Seed(time.Now().UnixNano()) - devPort := rand.Intn(10_000) + 50_000 - testSrvHost := fmt.Sprintf("127.0.0.1:%d", rand.Intn(10_000)+50_000) + testSrvPort := 9091 + devPort := 9092 expectedTestUserSub := fmt.Sprintf("test-user-%d", devPort) svc := auth.NewService(auth.Opts{ @@ -30,7 +28,7 @@ func TestNewOpenID(t *testing.T) { }), Logger: logger.Std, AvatarStore: avatar.NewNoOp(), - URL: fmt.Sprintf("http://%s", testSrvHost), + URL: fmt.Sprintf("http://127.0.0.1:%d", testSrvPort), }) svc.AddDevOpenIDProvider(devPort) @@ -38,7 +36,7 @@ func TestNewOpenID(t *testing.T) { require.NoError(t, err) devAuth.Automatic = true - devAuth.CustomizeIdTokenFn = func(m map[string]interface{}) map[string]interface{} { + devAuth.CustomizeIDTokenFn = func(m map[string]interface{}) map[string]interface{} { m["sub"] = expectedTestUserSub return m } @@ -48,7 +46,7 @@ func TestNewOpenID(t *testing.T) { authHandler, _ := svc.Handlers() server := httptest.NewUnstartedServer(authHandler) - server.Listener, err = net.Listen("tcp", testSrvHost) + server.Listener, err = net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", testSrvPort)) require.NoError(t, err) server.Start() @@ -61,8 +59,8 @@ func TestNewOpenID(t *testing.T) { Jar: jar, } - require.NoError(t, waitFor(testSrvHost)) - require.NoError(t, waitFor(fmt.Sprintf("localhost:%d", devPort))) + require.NoError(t, waitFor(fmt.Sprintf("127.0.0.1:%d", testSrvPort))) + require.NoError(t, waitFor(fmt.Sprintf("127.0.0.1:%d", devPort))) resp, err := client.Get(server.URL + "/auth/dev/login") require.NoError(t, err) From 41a0ffed19081a91c99b33278482d5200c363659 Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Wed, 22 Jun 2022 10:11:26 +0100 Subject: [PATCH 09/13] Make sure AddDevOpenIDProvider is called in auth_test.go Weirdly coveralls thinks this method is not covered, because it is tested in another package. However there isn't much to test really, so at best I can check jwks URL is correctly served. --- auth_test.go | 29 +++++++++++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/auth_test.go b/auth_test.go index 601cd92d..823748cb 100644 --- a/auth_test.go +++ b/auth_test.go @@ -374,6 +374,35 @@ func TestDirectProvider(t *testing.T) { assert.NoError(t, resp.Body.Close()) } +func TestDevOpenIDProvider(t *testing.T) { + service := NewService(Opts{Logger: logger.Std}) + service.AddDevOpenIDProvider(18089) + + devAuth, err := service.DevAuth() + require.NoError(t, err) + + go devAuth.Run(context.Background()) + defer devAuth.Shutdown() + + for i := 1; i < 20; i++ { + time.Sleep(time.Duration(i*10) * time.Millisecond) + + dial, e := net.Dial("tcp", "localhost:18089") + if e == nil { + e = dial.Close() + require.NoError(t, e) + + break + } + } + + jwksResp, err := http.Get("http://localhost:18089/jwks") + require.NoError(t, err) + + require.Equal(t, 200, jwksResp.StatusCode) + // actual OpenID flow is tested in the openid_test.go, but coverage tool isn't picking it up +} + func TestDirectProvider_WithCustomUserIDFunc(t *testing.T) { _, teardown := prepService(t) defer teardown() From 9de085f52d0bae1b38600f7ef7023a473fcb58e1 Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Wed, 22 Jun 2022 11:36:03 +0100 Subject: [PATCH 10/13] Add auth tests Actual login flow is tested already, and these two new methods are called in provider/openid_test.go. However the coverage tool is not detecting these calls, and instead seems to be requiring the methods to be called in the matching test file. So this test is a weird artifact to make coverage tool happy. --- auth_test.go | 25 ++++++++++++++++++++++--- 1 file changed, 22 insertions(+), 3 deletions(-) diff --git a/auth_test.go b/auth_test.go index 823748cb..747513cb 100644 --- a/auth_test.go +++ b/auth_test.go @@ -3,6 +3,7 @@ package auth import ( "context" "encoding/json" + "golang.org/x/oauth2" "io" "io/ioutil" "net" @@ -375,7 +376,9 @@ func TestDirectProvider(t *testing.T) { } func TestDevOpenIDProvider(t *testing.T) { - service := NewService(Opts{Logger: logger.Std}) + service := NewService(Opts{Logger: logger.Std, SecretReader: token.SecretFunc(func(aud string) (string, error) { + return "secret", nil + })}) service.AddDevOpenIDProvider(18089) devAuth, err := service.DevAuth() @@ -398,9 +401,25 @@ func TestDevOpenIDProvider(t *testing.T) { jwksResp, err := http.Get("http://localhost:18089/jwks") require.NoError(t, err) + assert.Equal(t, 200, jwksResp.StatusCode) - require.Equal(t, 200, jwksResp.StatusCode) - // actual OpenID flow is tested in the openid_test.go, but coverage tool isn't picking it up + service.AddOpenIDProvider("openid", Client{Cid: "cid", Csecret: "csecret"}, provider.CustomHandlerOpt{ + Endpoint: oauth2.Endpoint{ + AuthURL: "http://localhost:18089/login/oauth/authorize", + TokenURL: "http://localhost:18089/login/oauth/access_token", + AuthStyle: oauth2.AuthStyleAutoDetect, + }, + InfoURL: "http://localhost:18089/user", + JwksURL: "http://localhost:18089/jwks", + MapUserFn: func(data provider.UserData, bytes []byte) token.User { + return token.User{ + Name: data.Value("sub"), + } + }, + }) + + assert.Len(t, service.Providers(), 2) + // OpenID flow is tested in the openid_test.go, but coverage tool isn't picking it up } func TestDirectProvider_WithCustomUserIDFunc(t *testing.T) { From 28a499c36dcff3e3ac5d7cd3942adda0de4f3a06 Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Fri, 24 Jun 2022 11:22:46 +0100 Subject: [PATCH 11/13] Fix token validation and update README golang-jwt library is trying to validate iat claim of the ID token and due to not accounting for clock skew, validation pretty randomly fails. There is an open issue https://github.com/golang-jwt/jwt/issues/98 and seems like that is fixed in v4. However it is still unclear why iat is validation in the first place, that's not required by RFC and doesn't seem like the right thing to do. Only nbf and exp claims should be used for token lifetime validity check. Also, update README to show how to configure OpenID providers. --- README.md | 31 +++++++++++++++++++++++++++++++ provider/oauth2.go | 19 ++++++++++++++++++- provider/openid_test.go | 5 +++++ 3 files changed, 54 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index cbd5146c..8ccab4a0 100644 --- a/README.md +++ b/README.md @@ -4,6 +4,7 @@ This library provides "social login" with Github, Google, Facebook, Microsoft, Twitter, Yandex, Battle.net, Apple, Patreon and Telegram as well as custom auth providers and email verification. - Multiple oauth2 providers can be used at the same time +- Support of ID Tokens (OpenID) for loading user details - Special `dev` provider allows local testing and development - JWT stored in a secure cookie with XSRF protection. Cookies can be session-only - Minimal scopes with user name, id and picture (avatar) only @@ -320,6 +321,36 @@ In order to add a new oauth2 provider following input is required: service.AddCustomProvider("custom123", auth.Client{Cid: "cid", Csecret: "csecret"}, prov.HandlerOpt) ``` +### Using ID Tokens (OpenID Connect) + +Example of configuring OAuth2 with OpenID Connect: + +```go +c := auth.Client{ + Cid: os.Getenv("AEXMPL_CIDE"), + Csecret: os.Getenv("AEXMPL_CSED"), +} + +service.AddOpenIDProvider("my-openid", c, provider.CustomHandlerOpt{ + Endpoint: oauth2.Endpoint{ + AuthURL: "https://my-open-id-provider.com/oauth2/authorize", + TokenURL: "https://my-open-id-provider.com/oauth2/token", + }, + JwksURL: "https://my-open-id-provider.com/.well-known/jwks", + InfoURL: "https://my-open-id-provider.com/user/", + MapUserFn: func (data provider.UserData, _ []byte) token.User { + userInfo := token.User{ + ID: data.Value("sub"), // standard OpenID Connect claims are available + Name: data.Value("given_name"), + } + return userInfo + }, + Scopes: []string{"openid", "email", "profile"}, // defaulted to "openid" if not specified +}) +``` + +JWKS are loaded on application start, and then cached. There is no background refresh, but requesting unknown key (kid) will trigger keys reload. + ### Self-implemented auth handler Additionally it is possible to implement own auth handler. It may be useful if auth provider does not conform to oauth standard. Self-implemented handler has to implement `provider.Provider` interface. ```go diff --git a/provider/oauth2.go b/provider/oauth2.go index b4eb1626..f0c3f045 100644 --- a/provider/oauth2.go +++ b/provider/oauth2.go @@ -20,6 +20,8 @@ import ( "github.com/go-pkgz/auth/token" ) +const clockSkew = 10 * time.Second + // Oauth2Handler implements /login, /callback and /logout handlers from aouth2 flow type Oauth2Handler struct { Params @@ -257,7 +259,13 @@ func (p Oauth2Handler) loadUserFromIDToken(tok *oauth2.Token) (UserData, []byte, } claims := jwt.MapClaims{} - parsedIDToken, err := jwt.ParseWithClaims(idToken, &claims, p.keyfunc) + parser := jwt.Parser{ + // claims validation is not considering clock skew and randomly failing with iat validation + // nbf and exp are validated below + SkipClaimsValidation: true, + } + + parsedIDToken, err := parser.ParseWithClaims(idToken, &claims, p.keyfunc) if err != nil { return nil, nil, fmt.Errorf("failed to parse id token") } @@ -266,6 +274,15 @@ func (p Oauth2Handler) loadUserFromIDToken(tok *oauth2.Token) (UserData, []byte, return nil, nil, fmt.Errorf("invalid id token") } + now := time.Now().Add(clockSkew).Unix() + if !claims.VerifyExpiresAt(now, false) { + return nil, nil, fmt.Errorf("id token expired") + } + + if !claims.VerifyNotBefore(now, false) { + return nil, nil, fmt.Errorf("id token is not yet valid") + } + return UserData(claims), []byte(idToken), nil } diff --git a/provider/openid_test.go b/provider/openid_test.go index 762c1977..d6dc4e56 100644 --- a/provider/openid_test.go +++ b/provider/openid_test.go @@ -38,6 +38,11 @@ func TestNewOpenID(t *testing.T) { devAuth.Automatic = true devAuth.CustomizeIDTokenFn = func(m map[string]interface{}) map[string]interface{} { m["sub"] = expectedTestUserSub + + now := time.Now().Add(1 * time.Second) // simulate clock difference + m["iat"] = now.Unix() + m["nbf"] = now.Unix() + m["exp"] = now.Add(1 * time.Minute).Unix() return m } From d6b8f26b07e6e83876ef33add58b517d9cc952b6 Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Fri, 24 Jun 2022 13:48:07 +0100 Subject: [PATCH 12/13] Wrap errors when loading user info --- provider/oauth2.go | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/provider/oauth2.go b/provider/oauth2.go index f0c3f045..c8fb354e 100644 --- a/provider/oauth2.go +++ b/provider/oauth2.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "github.com/MicahParks/keyfunc" + "github.com/pkg/errors" "io" "net/http" "strings" @@ -254,7 +255,7 @@ func (p Oauth2Handler) loadUserFromIDToken(tok *oauth2.Token) (UserData, []byte, if p.keyfunc == nil { err := p.tryInitJWKSKeyfunc() if err != nil { - return nil, nil, fmt.Errorf("can't load JWKS keys") + return nil, nil, errors.Wrap(err, "can't load JWKS keys") } } @@ -267,7 +268,7 @@ func (p Oauth2Handler) loadUserFromIDToken(tok *oauth2.Token) (UserData, []byte, parsedIDToken, err := parser.ParseWithClaims(idToken, &claims, p.keyfunc) if err != nil { - return nil, nil, fmt.Errorf("failed to parse id token") + return nil, nil, errors.Wrap(err, "failed to parse id token") } if !parsedIDToken.Valid { @@ -289,7 +290,7 @@ func (p Oauth2Handler) loadUserFromIDToken(tok *oauth2.Token) (UserData, []byte, func (p Oauth2Handler) loadUserFromEndpoint(client *http.Client) (UserData, []byte, error) { uinfo, err := client.Get(p.infoURL) if err != nil { - return nil, nil, fmt.Errorf("failed to get client info") + return nil, nil, errors.Wrap(err, "failed to get client info") } defer func() { @@ -300,12 +301,12 @@ func (p Oauth2Handler) loadUserFromEndpoint(client *http.Client) (UserData, []by data, err := io.ReadAll(uinfo.Body) if err != nil { - return nil, nil, fmt.Errorf("failed to read user info") + return nil, nil, errors.Wrap(err, "failed to read user info") } jData := map[string]interface{}{} if e := json.Unmarshal(data, &jData); e != nil { - return nil, nil, fmt.Errorf("failed to unmarshal user info") + return nil, nil, errors.Wrap(e, "failed to unmarshal user info") } p.Logf("[DEBUG] got raw user info %+v", jData) From 3bf3b6252353d1f9147f208aaffc9fca780291bd Mon Sep 17 00:00:00 2001 From: Alexey Nesterov Date: Mon, 27 Jun 2022 10:42:19 +0100 Subject: [PATCH 13/13] Tidy up --- provider/oauth2.go | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/provider/oauth2.go b/provider/oauth2.go index c8fb354e..56f77996 100644 --- a/provider/oauth2.go +++ b/provider/oauth2.go @@ -183,10 +183,11 @@ func (p Oauth2Handler) AuthHandler(w http.ResponseWriter, r *http.Request) { var userData UserData var rawUserData []byte - switch p.UseOpenID { - case true: + if p.UseOpenID { userData, rawUserData, err = p.loadUserFromIDToken(tok) - case false: + } + + if !p.UseOpenID { userData, rawUserData, err = p.loadUserFromEndpoint(client) }