From 60e6900c2ecb88c6eccf7f7d4081c385be9059d8 Mon Sep 17 00:00:00 2001 From: reesporte <45641995+reesporte@users.noreply.github.com> Date: Fri, 20 May 2022 16:12:27 -0500 Subject: [PATCH] Add refresh token header/cookie (#2071) * Add refresh token header/cookie As part of work on automatic refreshing of access tokens in the grafana plugin (FB-1377), we will now accept a refresh token in the "X-Molecula-Refresh-Token" header or the "refresh-molecula-chip" cookie. This refresh token will be used if the access token is expired. To achieve this, there was a lot of plumbing that had to be done. Here is a list of some of it: * Added lots of constants for the new values. * Removed token cache, since we will be keeping state on the clients. * We now only refresh tokens when they are expired, which is more inline with the OAuth spec. * Refactored SetGRPCMetadata to be simpler to read. * Refactored AddAuthToken. * Update failing tests. * We now don't split GRPC cookies on ";". Not sure why we did that before tbh. I also added TODOs to add the refresh token to other subcommands. This is out of scope for my current ticket, but it would be nice to have in the future. * remove unnecessary context from Authenticate * Add comments on why we check both cases for headers It's because some GRPC clients lowercase metadata names. I've run into issues with this enough that I think it's worth the extra checks. We prefer lowercase though, because that's "standard". * Fix test that broke during rebase --- authn/authenticate.go | 239 +++++++++++++---------- authn/authenticate_internal_test.go | 176 +++++++++-------- ctl/backup.go | 9 +- ctl/import.go | 9 +- ctl/import_test.go | 10 +- ctl/restore.go | 9 +- http_handler.go | 78 +++++--- http_handler_internal_test.go | 4 +- internal/clustertests/cluster_test.go | 12 +- internal/clustertests/pause_node_test.go | 9 +- internal_client.go | 133 +++++++------ internal_client_test.go | 14 +- server/grpc.go | 75 +++++-- server/grpc_internal_test.go | 158 +++++++++++++++ 14 files changed, 624 insertions(+), 311 deletions(-) create mode 100644 server/grpc_internal_test.go diff --git a/authn/authenticate.go b/authn/authenticate.go index 66d1cfdb5..48dc5a823 100644 --- a/authn/authenticate.go +++ b/authn/authenticate.go @@ -24,8 +24,22 @@ import ( "golang.org/x/oauth2" ) -// CookieName is the name of the cookie that holds the refreshed auth token. -const CookieName = "molecula-chip" +const ( + // AccessCookieName is the name of the cookie that holds the access token. + AccessCookieName = "molecula-chip" + + // RefreshCookieName is the name of the cookie that holds the refresh token. + RefreshCookieName = "refresh-molecula-chip" + + // RefreshHeaderName is the name of the header that holds the refresh token. + RefreshHeaderName = "X-Molecula-Refresh-Token" + + // ContextValueAccessToken is the key used to set AccessTokens in a ctx. + ContextValueAccessToken = "Access" + + // ContextValueRefreshToken is the key used to set RefreshTokens in a ctx. + ContextValueRefreshToken = "Refresh" +) // cachedGroups is used to hold groups and when they were last cached type cachedGroups struct { @@ -33,19 +47,14 @@ type cachedGroups struct { groups []Group } -// cacheToken is used to hold tokens and when they were added to the cache -type cachedToken struct { - cacheTime time.Time - token *oauth2.Token -} - // UserInfo holds the information about the user from the token type UserInfo struct { - UserID string `json:"userid"` - UserName string `json:"username"` - Groups []Group `json:"groups"` - Expiry time.Time `json:"expiry"` - Token string `json:"token"` + UserID string `json:"userid"` + UserName string `json:"username"` + Groups []Group `json:"groups"` + Expiry time.Time `json:"expiry"` + Token string `json:"token"` + RefreshToken string `json:"refreshtoken"` } // Group holds group information for an authenticated user @@ -62,29 +71,29 @@ type Groups struct { // Auth holds state, configuration, and utilities needed for authentication. type Auth struct { - logger logger.Logger - cookieName string - secretKey []byte - groupEndpoint string - logoutEndpoint string - fbURL string // fbURL is the domain featurebase is hosted on, used for post logout redirection - oAuthConfig *oauth2.Config - cacheTTL time.Duration // cacheTTL is used to determine if a cached item should be refreshed or not - tokenTTR time.Duration // tokenTTR (time to refresh) is used to determine if a token should be refreshed or not - tokenCache map[string]cachedToken // tokenCache is a map of accessToken -> *oauth2.Token which we can use to refresh the tokens - groupsCache map[string]cachedGroups // groupsCache is a map of accessToken -> group memberships - lastCacheClean time.Time // last cache clean is the time that the cache was last cleaned - allowedNetworks []net.IPNet // list of allowed networks for ingest + logger logger.Logger + accessCookieName string + refreshCookieName string + secretKey []byte + groupEndpoint string + logoutEndpoint string + fbURL string // fbURL is the domain featurebase is hosted on, used for post logout redirection + oAuthConfig *oauth2.Config + cacheTTL time.Duration // cacheTTL is used to determine if a cached item should be refreshed or not + groupsCache map[string]cachedGroups // groupsCache is a map of accessToken -> group memberships + lastCacheClean time.Time // last cache clean is the time that the cache was last cleaned + allowedNetworks []net.IPNet // list of allowed networks for ingest } // NewAuth instantiates and returns a new Auth struct func NewAuth(logger logger.Logger, url string, scopes []string, authURL, tokenURL, groupEndpoint, logout, clientID, clientSecret, secretKey string, configuredIPs []string) (auth *Auth, err error) { auth = &Auth{ - logger: logger, - cookieName: CookieName, - groupEndpoint: groupEndpoint, - logoutEndpoint: logout, - fbURL: url, + logger: logger, + accessCookieName: AccessCookieName, + refreshCookieName: RefreshCookieName, + groupEndpoint: groupEndpoint, + logoutEndpoint: logout, + fbURL: url, oAuthConfig: &oauth2.Config{ RedirectURL: fmt.Sprintf("%s/redirect", url), ClientID: clientID, @@ -95,10 +104,8 @@ func NewAuth(logger logger.Logger, url string, scopes []string, authURL, tokenUR TokenURL: tokenURL, }, }, - tokenCache: map[string]cachedToken{}, groupsCache: map[string]cachedGroups{}, cacheTTL: 10 * time.Minute, - tokenTTR: 7 * time.Minute, lastCacheClean: time.Now(), } @@ -120,54 +127,57 @@ func (a Auth) SecretKey() []byte { return a.secretKey } -// Authenticate takes in a bearer token `bearer` and returns UserInfo from that token +// refreshToken refreshes a given access/refresh token pair +func (a *Auth) refreshToken(access, refresh string) (string, string, error) { + resp, err := http.PostForm(a.oAuthConfig.Endpoint.TokenURL, + url.Values{ + "grant_type": {"refresh_token"}, + "refresh_token": {refresh}, + "client_id": {a.oAuthConfig.ClientID}, + "client_secret": {a.oAuthConfig.ClientSecret}, + }, + ) + + if err != nil { + return "", "", errors.Wrap(err, "refreshing token") + } + + if resp.StatusCode != http.StatusOK { + return "", "", fmt.Errorf("refreshing token: %s", resp.Status) + } + + defer resp.Body.Close() + + var t oauth2.Token + if err := json.NewDecoder(resp.Body).Decode(&t); err != nil { + return "", "", errors.Wrap(err, "decoding refreshed token") + } + + // remove the old groups from the groups cache + delete(a.groupsCache, access) + + return t.AccessToken, t.RefreshToken, nil +} + +// Authenticate takes in a auth token `access` and returns UserInfo from that token // it is caller's responsibility to inform the user that the access token has been refreshed -func (a *Auth) Authenticate(ctx context.Context, bearer string) (*UserInfo, error) { +func (a *Auth) Authenticate(access, refresh string) (*UserInfo, error) { // clean up the cache every 30 minutes or so if time.Now().Sub(a.lastCacheClean) >= 30*time.Minute { a.cleanCache() } - if tkn, ok := a.tokenCache[bearer]; ok && (tkn.token.Expiry.Sub(time.Now()) <= a.tokenTTR || !tkn.token.Valid()) { - // refresh the token - resp, err := http.PostForm(a.oAuthConfig.Endpoint.TokenURL, - url.Values{ - "grant_type": {"refresh_token"}, - "refresh_token": {tkn.token.RefreshToken}, - "client_id": {a.oAuthConfig.ClientID}, - "client_secret": {a.oAuthConfig.ClientSecret}, - }, - ) - if err != nil { - return nil, errors.Wrap(err, "refreshing token") - } - if resp.StatusCode != http.StatusOK { - return nil, fmt.Errorf("refreshing token: %s", resp.Status) - } - defer resp.Body.Close() - var t oauth2.Token - if err := json.NewDecoder(resp.Body).Decode(&t); err != nil { - return nil, errors.Wrap(err, "decoding refreshed token") - } - - // update the cache - delete(a.tokenCache, bearer) - delete(a.groupsCache, bearer) - bearer = t.AccessToken - a.tokenCache[bearer] = cachedToken{time.Now(), &t} - } - - if len(bearer) == 0 { - return nil, fmt.Errorf("bearer token is empty") + if len(access) == 0 { + return nil, fmt.Errorf("auth token is empty") } // NOTE: we are using ParseUnverified here because the IDP validates the // token's signature when we get the user's groups, we just need to make // sure it's not expired and is well-formed - token, _, err := new(jwt.Parser).ParseUnverified(bearer, &jwt.MapClaims{}) + token, _, err := new(jwt.Parser).ParseUnverified(access, &jwt.MapClaims{}) // well-formed-ness check if token == nil || token.Claims == nil || err != nil { - return nil, fmt.Errorf("parsing bearer token: %v", err) + return nil, fmt.Errorf("parsing auth token: %v", err) } claims := *token.Claims.(*jwt.MapClaims) @@ -175,14 +185,19 @@ func (a *Auth) Authenticate(ctx context.Context, bearer string) (*UserInfo, erro // expiry check if exp, ok := claims["exp"].(string); ok { if expiry, err := strconv.ParseInt(exp, 10, 64); err != nil || expiry < time.Now().UTC().Unix() { - return nil, fmt.Errorf("token is expired") + access, refresh, err = a.refreshToken(access, refresh) + if err != nil { + return nil, fmt.Errorf("token is expired: %w", err) + } } } userInfo := UserInfo{ - Token: bearer, - Groups: []Group{}, + Token: access, + RefreshToken: refresh, + Groups: []Group{}, } + if uid, ok := claims["oid"].(string); ok { userInfo.UserID = uid } @@ -190,7 +205,7 @@ func (a *Auth) Authenticate(ctx context.Context, bearer string) (*UserInfo, erro userInfo.UserName = name } - if userInfo.Groups, err = a.getGroups(bearer); err != nil { + if userInfo.Groups, err = a.getGroups(access); err != nil { return nil, errors.Wrap(err, "getting groups") } @@ -199,18 +214,11 @@ func (a *Auth) Authenticate(ctx context.Context, bearer string) (*UserInfo, erro // cleanCache removes old items from our cache func (a *Auth) cleanCache() { - for bearer, tkn := range a.tokenCache { - // if it's been more than 24 hours since the token was cached - if time.Now().Sub(tkn.cacheTime) >= 24*time.Hour { - // remove it from our cache - delete(a.tokenCache, bearer) - } - } - for bearer, tkn := range a.groupsCache { + for access, tkn := range a.groupsCache { // if it's been more than 24 hours since the groups were cached if time.Now().Sub(tkn.cacheTime) >= 24*time.Hour { // remove it from our cache - delete(a.groupsCache, bearer) + delete(a.groupsCache, access) } } a.lastCacheClean = time.Now() @@ -225,14 +233,22 @@ func (a *Auth) Login(w http.ResponseWriter, r *http.Request) { // Logout clears out the user's cookie, removes the token from our cache, and // redirects user to IdP's logout endpoint func (a *Auth) Logout(w http.ResponseWriter, r *http.Request) { - // remove the bearer token from a.tokenCache and a.groupsCache - if bearer, err := r.Cookie(a.cookieName); err == nil { - delete(a.tokenCache, bearer.Value) - delete(a.groupsCache, bearer.Value) + // remove the access token from a.groupsCache + if access, err := r.Cookie(a.accessCookieName); err == nil { + delete(a.groupsCache, access.Value) } // clear cookie http.SetCookie(w, &http.Cookie{ - Name: a.cookieName, + Name: a.accessCookieName, + Value: "", + Path: "/", + Secure: true, + HttpOnly: true, + SameSite: http.SameSiteStrictMode, + Expires: time.Unix(0, 0), + }) + http.SetCookie(w, &http.Cookie{ + Name: a.refreshCookieName, Value: "", Path: "/", Secure: true, @@ -254,9 +270,7 @@ func (a *Auth) Redirect(w http.ResponseWriter, r *http.Request) { return } - a.tokenCache[token.AccessToken] = cachedToken{time.Now(), token} - - a.SetCookie(w, token.AccessToken, token.Expiry) + a.SetCookie(w, token.AccessToken, token.RefreshToken, token.Expiry) http.Redirect(w, r, "/", http.StatusTemporaryRedirect) } @@ -306,10 +320,20 @@ func (a *Auth) getGroups(token string) ([]Group, error) { return groups.Groups, nil } -func (a *Auth) SetCookie(w http.ResponseWriter, token string, expiry time.Time) error { +func (a *Auth) SetCookie(w http.ResponseWriter, access, refresh string, expiry time.Time) error { http.SetCookie(w, &http.Cookie{ - Name: a.cookieName, - Value: token, + Name: a.refreshCookieName, + Value: refresh, + Path: "/", + Secure: true, + HttpOnly: true, + SameSite: http.SameSiteStrictMode, + Expires: expiry, + }) + + http.SetCookie(w, &http.Cookie{ + Name: a.accessCookieName, + Value: access, Path: "/", Secure: true, HttpOnly: true, @@ -319,18 +343,23 @@ func (a *Auth) SetCookie(w http.ResponseWriter, token string, expiry time.Time) return nil } -func (a *Auth) SetGRPCMetadata(ctx context.Context, md metadata.MD, token string) error { - cookies := []string{} +func (a *Auth) SetGRPCMetadata(ctx context.Context, md metadata.MD, access, refresh string) error { + mCookies := map[string]string{} if c, ok := md["cookie"]; ok { for _, cookie := range c { - if strings.HasPrefix(cookie, a.cookieName) { - cookie = a.cookieName + "=" + token - } - cookies = append(cookies, cookie) + name, val := parseCookie(cookie) + mCookies[name] = val } - } else { - cookies = []string{a.cookieName + "=" + token} } + + mCookies[a.accessCookieName] = access + mCookies[a.refreshCookieName] = refresh + + cookies := []string{} + for name, val := range mCookies { + cookies = append(cookies, name+"="+val) + } + md["cookie"] = cookies return grpc.SetHeader(ctx, md) } @@ -381,3 +410,13 @@ func (a *Auth) CheckAllowedNetworks(clientIP string) bool { } return false } + +func parseCookie(cookie string) (name, data string) { + vals := strings.Split(cookie, "=") + if len(vals) == 0 { + vals = []string{"", ""} + } else if len(vals) < 2 { + vals = append(vals, "") + } + return vals[0], vals[1] +} diff --git a/authn/authenticate_internal_test.go b/authn/authenticate_internal_test.go index 10dd307fb..727180aa7 100644 --- a/authn/authenticate_internal_test.go +++ b/authn/authenticate_internal_test.go @@ -18,7 +18,6 @@ import ( "github.com/golang-jwt/jwt" "github.com/molecula/featurebase/v3/logger" - "golang.org/x/oauth2" "google.golang.org/grpc" "google.golang.org/grpc/metadata" ) @@ -59,9 +58,13 @@ func NewTestAuth(t *testing.T) *Auth { func TestSetGRPCMetadata(t *testing.T) { a := NewTestAuth(t) for name, md := range map[string]metadata.MD{ - "empty": {}, - "something": {"cookie": []string{a.cookieName + "=something"}}, - "otherCookies": {"cookie": []string{a.cookieName + "=something", "blah=blah"}}, + "empty": {}, + "something": {"cookie": []string{a.accessCookieName + "=something"}}, + "somethingElse": {"cookie": []string{ + a.accessCookieName + "=something", + a.refreshCookieName + "=something", + }}, + "otherCookies": {"cookie": []string{a.accessCookieName + "=something", "blah=blah"}}, } { t.Run(name, func(t *testing.T) { ogCookies, _ := md["cookie"] @@ -75,7 +78,7 @@ func TestSetGRPCMetadata(t *testing.T) { if !ok { t.Fatalf("expected ok, got: %v", ok) } - err := a.SetGRPCMetadata(ctx, md, "this is a token!") + err := a.SetGRPCMetadata(ctx, md, "accesstoken!", "refreshtoken!") if err != nil { t.Fatalf("expected no errors, got: %v", err) } @@ -90,18 +93,29 @@ func TestSetGRPCMetadata(t *testing.T) { if !ok { t.Fatalf("expected ok, got: %v", ok) } - var cookie string - for _, cookie = range c { - if strings.HasPrefix(cookie, a.cookieName) { + var accessCookie, refreshCookie string + for _, cookie := range c { + if strings.HasPrefix(cookie, a.accessCookieName) { + accessCookie = cookie + } else if strings.HasPrefix(cookie, a.refreshCookieName) { + refreshCookie = cookie + } + if refreshCookie != "" && accessCookie != "" { break } } - if exp, got := a.cookieName+"=this is a token!", cookie; got != exp { - t.Fatalf("expected '%v', got '%v'", exp, got) + + exp := a.accessCookieName + "=accesstoken!" + if accessCookie != exp { + t.Fatalf("expected '%v', got '%v'", exp, accessCookie) + } + exp = a.refreshCookieName + "=refreshtoken!" + if refreshCookie != exp { + t.Fatalf("expected '%v', got '%v'", exp, refreshCookie) } - for _, cookie = range c { - if strings.HasPrefix(cookie, a.cookieName) { + for _, cookie := range c { + if strings.HasPrefix(cookie, a.accessCookieName) || strings.HasPrefix(cookie, a.refreshCookieName) { continue } found := false @@ -124,7 +138,7 @@ func TestAuth(t *testing.T) { a := NewTestAuth(t) t.Run("SetCookie", func(t *testing.T) { w := httptest.NewRecorder() - err := a.SetCookie(w, "a cookie value", time.Now().Add(time.Hour)) + err := a.SetCookie(w, "access", "refresh", time.Now().Add(time.Hour)) if err != nil { t.Fatalf("expected no errors, got: %v", err) } @@ -171,7 +185,7 @@ func TestAuthenticate(t *testing.T) { uname string exp int64 refresh bool - errOnRefresh bool + refreshToken string malformed bool empty bool groups []Group @@ -191,12 +205,12 @@ func TestAuthenticate(t *testing.T) { { name: "Malformed", malformed: true, - err: fmt.Errorf("parsing bearer token: token contains an invalid number of segments"), + err: fmt.Errorf("parsing auth token: token contains an invalid number of segments"), }, { name: "Empty", empty: true, - err: fmt.Errorf("bearer token is empty"), + err: fmt.Errorf("auth token is empty"), }, { name: "ExpiredTokenNoRefresh", @@ -209,7 +223,7 @@ func TestAuthenticate(t *testing.T) { }, }, exp: -17764800, - err: fmt.Errorf("token is expired"), + err: fmt.Errorf("token is expired: refreshing token: 400 Bad Request"), }, { name: "ExpiredTokenYesRefresh", @@ -221,8 +235,9 @@ func TestAuthenticate(t *testing.T) { GroupName: "adminGroup", }, }, - refresh: true, - exp: -17764800, + refresh: true, + refreshToken: "refreshToken", + exp: -17764800, }, { name: "ExpiredTokenYesRefreshButError", @@ -235,9 +250,9 @@ func TestAuthenticate(t *testing.T) { }, }, refresh: true, - errOnRefresh: true, + refreshToken: "blah!!", exp: -17764800, - err: fmt.Errorf("refreshing token: 500 Internal Server Error"), + err: fmt.Errorf("token is expired: refreshing token: 403 Forbidden"), }, } for _, test := range cases { @@ -266,41 +281,39 @@ func TestAuthenticate(t *testing.T) { } if test.refresh { var srv *httptest.Server - if !test.errOnRefresh { - srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - tkn := jwt.New(jwt.SigningMethodHS256) - claims := tkn.Claims.(jwt.MapClaims) - claims["oid"] = test.uid - claims["name"] = test.uname - expiry := strconv.Itoa(int(time.Now().Add(2 * time.Hour).Unix())) - claims["exp"] = expiry - fresh, err := tkn.SignedString(a.SecretKey()) - if err != nil { - t.Fatalf("unexpected error when signing token %v", err) - } + srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + t.Fatalf("unexpected error: %v", err) + } + refresh := r.Form.Get("refresh_token") + if refresh != test.refreshToken { + t.Fatalf("refresh token not passed properly, expected %v, got %v", test.refreshToken, refresh) + return + } + if refresh != "refreshToken" { + http.Error(w, "bad token", http.StatusForbidden) + } - a.groupsCache[fresh] = cachedGroups{time.Now(), test.groups} - fmt.Fprintf(w, `{"access_token": "`+fresh+`", "refresh_token": "blah", "token_type": "bearer", "expires": `+expiry+` }`) - })) - } else { - srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - http.Error(w, "bad", http.StatusInternalServerError) - })) - } + tkn := jwt.New(jwt.SigningMethodHS256) + claims := tkn.Claims.(jwt.MapClaims) + claims["oid"] = test.uid + claims["name"] = test.uname + expiry := strconv.Itoa(int(time.Now().Add(2 * time.Hour).Unix())) + claims["exp"] = expiry + fresh, err := tkn.SignedString(a.SecretKey()) + if err != nil { + t.Fatalf("unexpected error when signing token %v", err) + } + + a.groupsCache[fresh] = cachedGroups{time.Now(), test.groups} + fmt.Fprintf(w, `{"access_token": "`+fresh+`", "refresh_token": "blah", "token_type": "bearer", "expires": `+expiry+` }`) + })) defer srv.Close() a.oAuthConfig.Endpoint.TokenURL = srv.URL - a.tokenCache[token] = cachedToken{ - time.Now(), - &oauth2.Token{ - AccessToken: token, - RefreshToken: "blah", - Expiry: time.Unix(test.exp, 0), - }, - } } // do the actual testing - uinfo, err := a.Authenticate(context.TODO(), token) + uinfo, err := a.Authenticate(token, test.refreshToken) // okay this part kind of sucks bc we need to check errors and i // dont want to write a whole new test for things that should have // errors just to avoid this mess. errors.Is doesn't work either @@ -334,11 +347,9 @@ func TestAuthenticate_CleanCache(t *testing.T) { now := time.Now() a.groupsCache["oldy"] = cachedGroups{now.Add(-24 * time.Hour), []Group{}} a.groupsCache["goldy"] = cachedGroups{now.Add(-4 * time.Hour), []Group{}} - a.tokenCache["oldy"] = cachedToken{now.Add(-24 * time.Hour), &oauth2.Token{}} - a.tokenCache["goldy"] = cachedToken{now.Add(-4 * time.Hour), &oauth2.Token{}} a.lastCacheClean = now.Add(-45 * time.Minute) - _, _ = a.Authenticate(context.TODO(), "this doesn't matter") + _, _ = a.Authenticate("this doesn't matter", "this doesn't matter?") if a.lastCacheClean.Sub(now) <= time.Nanosecond { t.Fatalf("cache should have been cleaned") } @@ -348,23 +359,15 @@ func TestAuthenticate_CleanCache(t *testing.T) { if _, ok := a.groupsCache["goldy"]; !ok { t.Errorf("goldy should not have been deleted") } - if _, ok := a.tokenCache["oldy"]; ok { - t.Errorf("oldy should have been deleted") - } - if _, ok := a.tokenCache["goldy"]; !ok { - t.Errorf("goldy should not have been deleted") - } }) t.Run("shouldn't clean", func(t *testing.T) { a := NewTestAuth(t) now := time.Now() a.groupsCache["oldy"] = cachedGroups{now.Add(-24 * time.Hour), []Group{}} a.groupsCache["goldy"] = cachedGroups{now.Add(-4 * time.Hour), []Group{}} - a.tokenCache["oldy"] = cachedToken{now.Add(-24 * time.Hour), &oauth2.Token{}} - a.tokenCache["goldy"] = cachedToken{now.Add(-4 * time.Hour), &oauth2.Token{}} a.lastCacheClean = now - _, _ = a.Authenticate(context.TODO(), "this doesn't matter") + _, _ = a.Authenticate("this doesn't matter", "this doesn't matter?") if a.lastCacheClean.Sub(now) >= time.Nanosecond { t.Fatalf("cache should not have been cleaned") } @@ -374,12 +377,6 @@ func TestAuthenticate_CleanCache(t *testing.T) { if _, ok := a.groupsCache["goldy"]; !ok { t.Errorf("goldy should not have been deleted") } - if _, ok := a.tokenCache["oldy"]; !ok { - t.Errorf("oldy should not have been deleted") - } - if _, ok := a.tokenCache["goldy"]; !ok { - t.Errorf("goldy should not have been deleted") - } }) } @@ -517,7 +514,7 @@ func TestHandlers(t *testing.T) { w := httptest.NewRecorder() req.AddCookie( &http.Cookie{ - Name: a.cookieName, + Name: a.accessCookieName, Value: "test", Path: "/", Secure: true, @@ -526,8 +523,20 @@ func TestHandlers(t *testing.T) { Expires: time.Unix(3000000, 0), }, ) + + req.AddCookie( + &http.Cookie{ + Name: a.refreshCookieName, + Value: "test", + Path: "/", + Secure: true, + HttpOnly: true, + SameSite: http.SameSiteStrictMode, + Expires: time.Unix(3000000, 0), + }, + ) + a.groupsCache["test"] = cachedGroups{} - a.tokenCache["test"] = cachedToken{time.Now(), &oauth2.Token{}} a.Logout(w, req) resp := w.Result() if resp.StatusCode != http.StatusTemporaryRedirect { @@ -538,7 +547,7 @@ func TestHandlers(t *testing.T) { t.Fatalf("expected %v, got %v", redirect, got.Path) } for _, c := range resp.Cookies() { - if c.Name == a.cookieName { + if c.Name == a.accessCookieName || c.Name == a.refreshCookieName { if c.Value != "" { t.Fatalf("cookie not set to empty value!") } @@ -547,15 +556,11 @@ func TestHandlers(t *testing.T) { if want != got { t.Fatalf("expected %v, got %v", want, got) } - break } } if _, ok := a.groupsCache["test"]; ok { t.Fatalf("groups not deleted!") } - if _, ok := a.tokenCache["test"]; ok { - t.Fatalf("token not deleted!") - } }) t.Run("redirectGood", func(t *testing.T) { req := httptest.NewRequest("GET", "/redirect", nil) @@ -572,11 +577,6 @@ func TestHandlers(t *testing.T) { if err != nil { t.Fatalf("unexpected error when signing token %v", err) } - freshToken := oauth2.Token{ - AccessToken: fresh, - RefreshToken: "blah", - Expiry: exp, - } srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { body := `{"access_token": "` + fresh + `", "refresh_token": "blah", "expires_in": "` + strconv.Itoa(int(expiresIn.Seconds())) + `"}` @@ -593,15 +593,13 @@ func TestHandlers(t *testing.T) { if got, err := resp.Location(); err != nil || got.String() != "/" { t.Fatalf("expected %v, got %v", "/", got.Path) } - cachedToken := a.tokenCache[fresh].token - if cachedToken.AccessToken != freshToken.AccessToken { - t.Fatalf("expected %v, got %v", freshToken.AccessToken, cachedToken.AccessToken) - } - if cachedToken.RefreshToken != freshToken.RefreshToken { - t.Fatalf("expected %v, got %v", freshToken.RefreshToken, cachedToken.RefreshToken) - } - if cachedToken.Expiry.Sub(freshToken.Expiry) > time.Second { - t.Fatalf("expected %v, got %v", freshToken.Expiry, cachedToken.Expiry) + cookies := resp.Cookies() + for _, c := range cookies { + if c.Name == a.accessCookieName && c.Value != fresh { + t.Fatalf("expected %v, got %v", exp, c.Value) + } else if c.Name == a.refreshCookieName && c.Value != "blah" { + t.Fatalf("expected %v, got %v", "blah", c.Value) + } } }) diff --git a/ctl/backup.go b/ctl/backup.go index 6327ea012..6bdd7a83a 100644 --- a/ctl/backup.go +++ b/ctl/backup.go @@ -13,6 +13,7 @@ import ( "time" pilosa "github.com/molecula/featurebase/v3" + "github.com/molecula/featurebase/v3/authn" "github.com/molecula/featurebase/v3/encoding/proto" "github.com/molecula/featurebase/v3/server" "github.com/molecula/featurebase/v3/topology" @@ -20,6 +21,8 @@ import ( "golang.org/x/sync/errgroup" ) +// TODO(rdp): add refresh token to this as well + // BackupCommand represents a command for backing up a FeatureBase node. type BackupCommand struct { // nolint: maligned tlsConfig *tls.Config @@ -108,7 +111,11 @@ func (cmd *BackupCommand) Run(ctx context.Context) (err error) { cmd.client = client if cmd.AuthToken != "" { - ctx = context.WithValue(ctx, "token", "Bearer "+cmd.AuthToken) + ctx = context.WithValue( + ctx, + authn.ContextValueAccessToken, + "Bearer "+cmd.AuthToken, + ) } // Determine the field type in order to correctly handle the input data. diff --git a/ctl/import.go b/ctl/import.go index 09bc4e0c4..8d44fd98e 100644 --- a/ctl/import.go +++ b/ctl/import.go @@ -12,11 +12,14 @@ import ( "time" pilosa "github.com/molecula/featurebase/v3" + "github.com/molecula/featurebase/v3/authn" "github.com/molecula/featurebase/v3/pql" "github.com/molecula/featurebase/v3/server" "github.com/pkg/errors" ) +// TODO(rdp): add refresh token to this as well + // ImportCommand represents a command for bulk importing data. type ImportCommand struct { // nolint: maligned // Destination host and port. @@ -87,7 +90,11 @@ func (cmd *ImportCommand) Run(ctx context.Context) error { cmd.client = client if cmd.AuthToken != "" { - ctx = context.WithValue(ctx, "token", "Bearer "+cmd.AuthToken) + ctx = context.WithValue( + ctx, + authn.ContextValueAccessToken, + "Bearer "+cmd.AuthToken, + ) } if cmd.CreateSchema { diff --git a/ctl/import_test.go b/ctl/import_test.go index 1aca06ad3..a5ac14b08 100644 --- a/ctl/import_test.go +++ b/ctl/import_test.go @@ -685,14 +685,14 @@ func TestImport_AuthOn(t *testing.T) { Field: "field1", CreateSchema: false, Token: invalidToken, - Err: fmt.Errorf("bearer token is empty"), + Err: fmt.Errorf("auth token is empty"), }, { Index: "test", Field: "field1", CreateSchema: true, Token: invalidToken, - Err: fmt.Errorf("bearer token is empty"), + Err: fmt.Errorf("auth token is empty"), }, } @@ -723,7 +723,11 @@ func TestImport_AuthOn(t *testing.T) { cm.Field = test.Field cm.CreateSchema = test.CreateSchema cm.Paths = []string{file.Name()} - ctx := context.WithValue(context.Background(), "token", test.Token) + ctx := context.WithValue( + context.Background(), + authn.ContextValueAccessToken, + test.Token, + ) err = cm.Run(ctx) if test.Err != nil { if !strings.Contains(err.Error(), test.Err.Error()) { diff --git a/ctl/restore.go b/ctl/restore.go index 0c8cb51b0..59b7ca73f 100644 --- a/ctl/restore.go +++ b/ctl/restore.go @@ -18,6 +18,7 @@ import ( "github.com/hashicorp/go-retryablehttp" pilosa "github.com/molecula/featurebase/v3" + "github.com/molecula/featurebase/v3/authn" "github.com/molecula/featurebase/v3/logger" "github.com/molecula/featurebase/v3/server" "github.com/molecula/featurebase/v3/topology" @@ -25,6 +26,8 @@ import ( "golang.org/x/sync/errgroup" ) +// TODO(rdp): add refresh token to this as well + // RestoreCommand represents a command for restoring a backup to type RestoreCommand struct { tlsConfig *tls.Config @@ -92,7 +95,7 @@ func (cmd *RestoreCommand) Run(ctx context.Context) (err error) { cmd.client = client if cmd.AuthToken != "" { - ctx = context.WithValue(ctx, "token", "Bearer "+cmd.AuthToken) + ctx = context.WithValue(ctx, authn.ContextValueAccessToken, "Bearer "+cmd.AuthToken) } nodes, err := cmd.client.Nodes(ctx) @@ -153,7 +156,7 @@ func (cmd *RestoreCommand) restoreSchema(ctx context.Context, primary *topology. req = req.WithContext(ctx) req.Header.Add("Accept", "application/json") - token, ok := ctx.Value("token").(string) + token, ok := ctx.Value(authn.ContextValueAccessToken).(string) if ok && token != "" { req.Header.Set("Authorization", token) } @@ -323,7 +326,7 @@ func (cmd *RestoreCommand) restoreShard(ctx context.Context, filename string) er req = req.WithContext(ctx) req.Header.Set("Content-Type", "application/octet-stream") - token, ok := ctx.Value("token").(string) + token, ok := ctx.Value(authn.ContextValueAccessToken).(string) if ok && token != "" { req.Header.Set("Authorization", token) } diff --git a/http_handler.go b/http_handler.go index 7ba2f5b4e..7fcba2816 100644 --- a/http_handler.go +++ b/http_handler.go @@ -625,14 +625,16 @@ func (h *Handler) chkAuthN(handler http.HandlerFunc) http.HandlerFunc { return } - uinfo, err := h.auth.Authenticate(ctx, getToken(r)) + access, refresh := getTokens(r) + uinfo, err := h.auth.Authenticate(access, refresh) if err != nil { http.Error(w, errors.Wrap(err, "authenticating").Error(), http.StatusUnauthorized) return } // just in case it got refreshed - h.auth.SetCookie(w, uinfo.Token, uinfo.Expiry) - ctx = context.WithValue(ctx, "token", r.Header["Authorization"]) + ctx = context.WithValue(ctx, authn.ContextValueAccessToken, "Bearer "+access) + ctx = context.WithValue(ctx, authn.ContextValueRefreshToken, refresh) + h.auth.SetCookie(w, uinfo.Token, uinfo.RefreshToken, uinfo.Expiry) } handler.ServeHTTP(w, r.WithContext(ctx)) } @@ -640,14 +642,14 @@ func (h *Handler) chkAuthN(handler http.HandlerFunc) http.HandlerFunc { func (h *Handler) chkAuthZ(handler http.HandlerFunc, perm authz.Permission) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - ctx := r.Context() - // if auth isn't turned on, just serve the request if h.auth == nil { handler.ServeHTTP(w, r) return } + ctx := r.Context() + // check if IP is in allowed networks, if yes give it admin permissions allowedNetwork, ctx := h.chkAllowedNetworks(r) if allowedNetwork { @@ -660,18 +662,24 @@ func (h *Handler) chkAuthZ(handler http.HandlerFunc, perm authz.Permission) http lperm := perm // check if the user is authenticated - uinfo, err := h.auth.Authenticate(ctx, getToken(r)) + access, refresh := getTokens(r) + + uinfo, err := h.auth.Authenticate(access, refresh) + + ctx = context.WithValue(ctx, authn.ContextValueAccessToken, "Bearer "+access) + ctx = context.WithValue(ctx, authn.ContextValueRefreshToken, refresh) + if err != nil { http.Error(w, errors.Wrap(err, "authenticating").Error(), http.StatusForbidden) return } // just in case it got refreshed - h.auth.SetCookie(w, uinfo.Token, uinfo.Expiry) - - // put the user's groups in the context - ctx = context.WithValue(ctx, contextKeyGroupMembership, uinfo.Groups) - ctx = context.WithValue(ctx, "token", "Bearer "+uinfo.Token) + h.auth.SetCookie(w, uinfo.Token, uinfo.RefreshToken, uinfo.Expiry) + // put the user's authN/Z info in the context + ctx = context.WithValue(r.Context(), contextKeyGroupMembership, uinfo.Groups) + ctx = context.WithValue(ctx, authn.ContextValueAccessToken, "Bearer "+uinfo.Token) + ctx = context.WithValue(ctx, authn.ContextValueRefreshToken, uinfo.RefreshToken) // unlikely h.permissions will be nil, but we'll check to be safe if h.permissions == nil { h.logger.Errorf("authentication is turned on without authorization permissions set") @@ -3795,14 +3803,15 @@ func (h *Handler) handleCheckAuthentication(w http.ResponseWriter, r *http.Reque return } - uinfo, err := h.auth.Authenticate(r.Context(), getToken(r)) + access, refresh := getTokens(r) + uinfo, err := h.auth.Authenticate(access, refresh) if uinfo == nil || err != nil { w.Header().Add("Content-Type", "text/plain") http.Error(w, err.Error(), http.StatusUnauthorized) return } // just in case it got refreshed - h.auth.SetCookie(w, uinfo.Token, uinfo.Expiry) + h.auth.SetCookie(w, uinfo.Token, uinfo.RefreshToken, uinfo.Expiry) w.Header().Add("Content-Type", "text/plain") w.WriteHeader(http.StatusOK) @@ -3818,14 +3827,16 @@ func (h *Handler) handleUserInfo(w http.ResponseWriter, r *http.Request) { http.Error(w, "", http.StatusNoContent) return } - uinfo, err := h.auth.Authenticate(r.Context(), getToken(r)) + + access, refresh := getTokens(r) + uinfo, err := h.auth.Authenticate(access, refresh) if err != nil { h.logger.Errorf("error authenticating: %v", err) http.Error(w, err.Error(), http.StatusForbidden) return } // just in case it got refreshed - h.auth.SetCookie(w, uinfo.Token, uinfo.Expiry) + h.auth.SetCookie(w, uinfo.Token, uinfo.RefreshToken, uinfo.Expiry) if err := json.NewEncoder(w).Encode(uinfo); err != nil { h.logger.Errorf("writing user info: %s", err) @@ -3840,19 +3851,36 @@ func (h *Handler) handleLogout(w http.ResponseWriter, r *http.Request) { h.auth.Logout(w, r) } -// getToken gets the access token from the request, returning empty string on -// error -func getToken(r *http.Request) string { +// getTokens gets the access and refresh tokens from the request, +// returning empty strings if they aren't in the request. +func getTokens(r *http.Request) (string, string) { + var access, refresh string if token, ok := r.Header["Authorization"]; ok && len(token) > 0 { parts := strings.Split(token[0], "Bearer ") - if len(parts) != 2 { - return "" + if len(parts) >= 2 { + access = parts[1] } - return parts[1] } - cookie, err := r.Cookie(authn.CookieName) - if err != nil { - return "" + + if token, ok := r.Header[authn.RefreshHeaderName]; ok && len(token) > 0 { + refresh = token[0] } - return cookie.Value + + if access == "" { + accessCookie, err := r.Cookie(authn.AccessCookieName) + if err != nil { + return access, refresh + } + access = accessCookie.Value + } + + if refresh == "" { + refreshCookie, err := r.Cookie(authn.RefreshCookieName) + if err != nil { + return access, refresh + } + refresh = refreshCookie.Value + } + + return access, refresh } diff --git a/http_handler_internal_test.go b/http_handler_internal_test.go index 14129bb4b..1ce0d3584 100644 --- a/http_handler_internal_test.go +++ b/http_handler_internal_test.go @@ -275,7 +275,7 @@ func TestAuthentication(t *testing.T) { expiredToken = "Bearer " + expiredToken validCookie := &http.Cookie{ - Name: authn.CookieName, + Name: authn.AccessCookieName, Value: token.AccessToken, Path: "/", Secure: true, @@ -687,7 +687,7 @@ func TestChkAuthN(t *testing.T) { name: "Invalid", token: invalidToken, handler: h.chkAuthN(testingHandler), - err: "authenticating: parsing bearer token", + err: "authenticating: parsing auth token", }, { name: "Expired", diff --git a/internal/clustertests/cluster_test.go b/internal/clustertests/cluster_test.go index 604161cc4..696f6e107 100644 --- a/internal/clustertests/cluster_test.go +++ b/internal/clustertests/cluster_test.go @@ -102,7 +102,11 @@ func TestClusterStuff(t *testing.T) { // generate auth token and add to context if auth { token = GetAuthToken(t) - ctx = context.WithValue(ctx, "token", "Bearer "+token) + ctx = context.WithValue( + ctx, + authn.ContextValueAccessToken, + "Bearer "+token, + ) } if err := cli[0].CreateIndex(ctx, "testidx", pilosa.IndexOptions{}); err != nil { @@ -325,7 +329,11 @@ func TestRetryLogic(t *testing.T) { } if auth { token := GetAuthToken(t) - ctx = context.WithValue(ctx, "token", "Bearer "+token) + ctx = context.WithValue( + ctx, + authn.ContextValueAccessToken, + "Bearer "+token, + ) } var addrs = []string{"pilosa1:10101", "pilosa2:10101", "pilosa3:10101"} diff --git a/internal/clustertests/pause_node_test.go b/internal/clustertests/pause_node_test.go index 164765f26..f4b49078f 100644 --- a/internal/clustertests/pause_node_test.go +++ b/internal/clustertests/pause_node_test.go @@ -15,6 +15,7 @@ import ( "time" pilosa "github.com/molecula/featurebase/v3" + "github.com/molecula/featurebase/v3/authn" boltdb "github.com/molecula/featurebase/v3/boltdb" "github.com/molecula/featurebase/v3/disco" "github.com/molecula/featurebase/v3/encoding/proto" @@ -23,6 +24,8 @@ import ( "github.com/pkg/errors" ) +// TODO(rdp): add refresh token to this test + func startCmd(cmd string, args ...string) (*exec.Cmd, error) { pcmd := exec.Command(cmd, args...) pcmd.Stdout = os.Stdout @@ -297,7 +300,11 @@ func TestPauseReplica(t *testing.T) { ctx := context.Background() if auth { token := GetAuthToken(t) - ctx = context.WithValue(ctx, "token", "Bearer "+token) + ctx = context.WithValue( + ctx, + authn.ContextValueAccessToken, + "Bearer "+token, + ) } ctx, cancel := context.WithCancel(ctx) diff --git a/internal_client.go b/internal_client.go index c4431de64..99774f015 100644 --- a/internal_client.go +++ b/internal_client.go @@ -144,23 +144,39 @@ func NewInternalClientFromURI(defaultURI *pnet.URI, remoteClient *http.Client, o return ic } -// AddAuthToken checks in a couple spots for our authorization token and adds it to -// the Authorization Header in the request if it finds it. -func AddAuthToken(ctx context.Context, req *http.Request) *http.Request { - if token, ok := ctx.Value("token").(string); ok && token != "" { - // the "token" value should be prefixed with "Bearer" - req.Header.Set("Authorization", token) - } else if uinfo := ctx.Value("userinfo"); uinfo != nil { - // UserInfo.Token is not prefixed with "Bearer" - req.Header.Set("Authorization", "Bearer "+uinfo.(*authn.UserInfo).Token) +// AddAuthToken checks in a couple spots for our authorization token and +// adds it to the Authorization Header in the request if it finds it. It does the +// same for refresh tokens as well. +func AddAuthToken(ctx context.Context, header *http.Header) { + var access, refresh string + if token, ok := ctx.Value(authn.ContextValueAccessToken).(string); ok { + // the AccessToken value should be prefixed with "Bearer" + access = token + } + if token, ok := ctx.Value(authn.ContextValueRefreshToken).(string); ok { + refresh = token + } + + // not combining these ifs so we don't call ctx.Value unless we have to + if access == "" || refresh == "" { + if uinfo := ctx.Value("userinfo"); uinfo != nil { + if access == "" { + // UserInfo.Token is not prefixed with "Bearer" + access = "Bearer " + uinfo.(*authn.UserInfo).Token + } + if refresh == "" { + refresh = uinfo.(*authn.UserInfo).RefreshToken + } + } } // set ogIP to request for remote calls if ogIP, ok := ctx.Value(OriginalIPHeader).(string); ok && ogIP != "" { - req.Header.Set(OriginalIPHeader, ogIP) + header.Set(OriginalIPHeader, ogIP) } - return req + header.Set("Authorization", access) + header.Set(authn.RefreshHeaderName, refresh) } // MaxShardByIndex returns the number of shards on a server by index. @@ -183,7 +199,7 @@ func (c *InternalClient) maxShardByIndex(ctx context.Context) (map[string]uint64 req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/json") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -216,7 +232,7 @@ func (c *InternalClient) AvailableShards(ctx context.Context, indexName string) req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/json") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -250,7 +266,7 @@ func (c *InternalClient) SchemaNode(ctx context.Context, uri *pnet.URI, views bo req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/json") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -282,7 +298,7 @@ func (c *InternalClient) Schema(ctx context.Context) ([]*IndexInfo, error) { req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/json") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -319,7 +335,7 @@ func (c *InternalClient) IngestSchema(ctx context.Context, uri *pnet.URI, buf [] req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx), giveRawResponse(true)) if err != nil { @@ -370,7 +386,7 @@ func (c *InternalClient) IngestOperations(ctx context.Context, uri *pnet.URI, in req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx)) if err != nil { @@ -403,7 +419,7 @@ func (c *InternalClient) IngestNodeOperations(ctx context.Context, uri *pnet.URI req.Header.Set("Content-Type", "application/x-protobuf") req.Header.Set("Accept", "application/x-protobuf") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx)) if err != nil { @@ -432,7 +448,7 @@ func (c *InternalClient) MutexCheck(ctx context.Context, uri *pnet.URI, indexNam } req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx)) if err != nil { @@ -463,7 +479,7 @@ func (c *InternalClient) PostSchema(ctx context.Context, uri *pnet.URI, s *Schem req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx)) if err != nil { @@ -510,7 +526,7 @@ func (c *InternalClient) CreateIndex(ctx context.Context, index string, opt Inde req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request against the host. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -540,7 +556,7 @@ func (c *InternalClient) FragmentNodes(ctx context.Context, index string, shard req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/json") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -572,7 +588,7 @@ func (c *InternalClient) Nodes(ctx context.Context) ([]*topology.Node, error) { req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/json") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -617,7 +633,7 @@ func (c *InternalClient) QueryNode(ctx context.Context, uri *pnet.URI, index str return nil, errors.Wrap(err, "creating request") } - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) req.Header.Set("Content-Length", strconv.Itoa(len(buf))) req.Header.Set("Content-Type", "application/x-protobuf") @@ -711,7 +727,7 @@ func (c *InternalClient) importNode(ctx context.Context, node *topology.Node, in req.Header.Set("Accept", "application/x-protobuf") req.Header.Set("X-Pilosa-Row", "roaring") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request against the host. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -932,7 +948,7 @@ func (c *InternalClient) ImportRoaring(ctx context.Context, uri *pnet.URI, index httpReq.Header.Set("Accept", "application/x-protobuf") httpReq.Header.Set("X-Pilosa-Row", "roaring") httpReq.Header.Set("User-Agent", "pilosa/"+Version) - httpReq = AddAuthToken(ctx, httpReq) + AddAuthToken(ctx, &httpReq.Header) // Execute request against the host. resp, err := c.executeRequest(httpReq.WithContext(ctx)) @@ -1007,7 +1023,7 @@ func (c *InternalClient) exportNodeCSV(ctx context.Context, node *topology.Node, } req.Header.Set("Accept", "text/csv") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request against the host. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1050,7 +1066,7 @@ func (c *InternalClient) RetrieveShardFromURI(ctx context.Context, index, field, } req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1142,7 +1158,7 @@ func (c *InternalClient) CreateFieldWithOptions(ctx context.Context, index, fiel req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request against the host. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1181,7 +1197,7 @@ func (c *InternalClient) FragmentBlocks(ctx context.Context, uri *pnet.URI, inde req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/json") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1231,7 +1247,7 @@ func (c *InternalClient) BlockData(ctx context.Context, uri *pnet.URI, index, fi req.Header.Set("Accept", "application/protobuf") req.Header.Set("X-Pilosa-Row", "roaring") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx)) if err != nil { @@ -1313,7 +1329,7 @@ func (c *InternalClient) TranslateKeysNode(ctx context.Context, uri *pnet.URI, i req.Header.Set("Accept", "application/x-protobuf") req.Header.Set("X-Pilosa-Row", "roaring") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request against the host. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1368,7 +1384,7 @@ func (c *InternalClient) TranslateIDsNode(ctx context.Context, uri *pnet.URI, in req.Header.Set("Accept", "application/x-protobuf") req.Header.Set("X-Pilosa-Row", "roaring") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request against the host. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1400,7 +1416,7 @@ func (c *InternalClient) GetPastQueries(ctx context.Context, uri *pnet.URI) ([]P req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request against the host. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1442,7 +1458,7 @@ func (c *InternalClient) FindIndexKeysNode(ctx context.Context, uri *pnet.URI, i req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Send the request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1491,7 +1507,7 @@ func (c *InternalClient) FindFieldKeysNode(ctx context.Context, uri *pnet.URI, i req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Send the request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1541,7 +1557,7 @@ func (c *InternalClient) CreateIndexKeysNode(ctx context.Context, uri *pnet.URI, req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Send the request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1594,7 +1610,7 @@ func (c *InternalClient) CreateFieldKeysNode(ctx context.Context, uri *pnet.URI, req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Send the request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1639,7 +1655,7 @@ func (c *InternalClient) MatchFieldKeysNode(ctx context.Context, uri *pnet.URI, req.Header.Set("Content-Length", strconv.Itoa(len(like))) req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Send the request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -1679,7 +1695,7 @@ func (c *InternalClient) Transactions(ctx context.Context) (map[string]*Transact } req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx)) if err != nil { @@ -1718,7 +1734,7 @@ func (c *InternalClient) StartTransaction(ctx context.Context, id string, timeou req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx), giveRawResponse(true)) if err != nil { @@ -1753,7 +1769,7 @@ func (c *InternalClient) FinishTransaction(ctx context.Context, id string) (*Tra req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx), giveRawResponse(true)) if err != nil { @@ -1790,7 +1806,7 @@ func (c *InternalClient) GetTransaction(ctx context.Context, id string) (*Transa } req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) resp, err := c.executeRequest(req.WithContext(ctx), giveRawResponse(true)) if err != nil { @@ -1856,7 +1872,9 @@ func (c *InternalClient) executeRetryableRequest(req *retryablehttp.Request, opt rc.HTTPClient = &http.Client{ CheckRedirect: func(req *http.Request, via []*http.Request) error { if len(via) > 0 { - req.Header.Set("Authorization", "Bearer "+getToken(via[0])) + access, refresh := getTokens(via[0]) + req.Header.Set("Authorization", "Bearer "+access) + req.Header.Set(authn.RefreshHeaderName, refresh) } return nil }, @@ -2154,7 +2172,7 @@ func (c *InternalClient) RetrieveTranslatePartitionFromURI(ctx context.Context, } req.Header.Set("User-Agent", "pilosa/"+Version) - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -2189,10 +2207,8 @@ func (c *InternalClient) ImportIndexKeys(ctx context.Context, uri *pnet.URI, ind return errors.Wrap(err, "creating request") } httpReq.Header.Set("User-Agent", "pilosa/"+Version) - token, ok := ctx.Value("token").(string) - if ok && token != "" { - httpReq.Header.Set("Authorization", token) - } + + AddAuthToken(ctx, &httpReq.Header) // Execute request against the host. resp, err := c.executeRetryableRequest(httpReq.WithContext(ctx)) @@ -2226,10 +2242,7 @@ func (c *InternalClient) ImportFieldKeys(ctx context.Context, uri *pnet.URI, ind } httpReq.Header.Set("User-Agent", "pilosa/"+Version) - token, ok := ctx.Value("token").(string) - if ok && token != "" { - httpReq.Header.Set("Authorization", token) - } + AddAuthToken(ctx, &httpReq.Header) // Execute request against the host. resp, err := c.executeRetryableRequest(httpReq.WithContext(ctx)) @@ -2256,7 +2269,7 @@ func (c *InternalClient) ShardReader(ctx context.Context, index string, shard ui req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/octet-stream") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -2279,7 +2292,7 @@ func (c *InternalClient) IDAllocDataReader(ctx context.Context) (io.ReadCloser, req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/octet-stream") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -2303,7 +2316,7 @@ func (c *InternalClient) IDAllocDataWriter(ctx context.Context, f io.Reader, pri req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/octet-stream") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. _, err = c.executeRequest(req.WithContext(ctx)) @@ -2330,7 +2343,7 @@ func (c *InternalClient) IndexTranslateDataReader(ctx context.Context, index str req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/octet-stream") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx), forwardAuthHeader(true)) @@ -2360,7 +2373,7 @@ func (c *InternalClient) FieldTranslateDataReader(ctx context.Context, index, fi req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/octet-stream") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -2389,7 +2402,7 @@ func (c *InternalClient) Status(ctx context.Context) (string, error) { req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/json") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) @@ -2421,7 +2434,7 @@ func (c *InternalClient) PartitionNodes(ctx context.Context, partitionID int) ([ req.Header.Set("User-Agent", "pilosa/"+Version) req.Header.Set("Accept", "application/json") - req = AddAuthToken(ctx, req) + AddAuthToken(ctx, &req.Header) // Execute request. resp, err := c.executeRequest(req.WithContext(ctx)) diff --git a/internal_client_test.go b/internal_client_test.go index 0eaae9ce4..c0f4b55c8 100644 --- a/internal_client_test.go +++ b/internal_client_test.go @@ -1576,7 +1576,7 @@ func TestAddAuthToken(t *testing.T) { if err != nil { t.Fatalf("unexpected error: %v", err) } - pilosa.AddAuthToken(context.Background(), req) + pilosa.AddAuthToken(context.Background(), &req.Header) if req.Header.Get("Authorization") != "" { t.Fatalf("Authorization header set when it should be empty") } @@ -1587,7 +1587,7 @@ func TestAddAuthToken(t *testing.T) { t.Fatalf("unexpected error: %v", err) } uinfo := &authn.UserInfo{Token: "ayo"} - pilosa.AddAuthToken(context.WithValue(context.Background(), "userinfo", uinfo), req) + pilosa.AddAuthToken(context.WithValue(context.Background(), "userinfo", uinfo), &req.Header) if got := req.Header.Get("Authorization"); got != "Bearer "+uinfo.Token { t.Fatalf("got '%v', expected 'Bearer %v'", got, uinfo.Token) } @@ -1598,7 +1598,13 @@ func TestAddAuthToken(t *testing.T) { t.Fatalf("unexpected error: %v", err) } tok := "Bearer thisisatoken" - pilosa.AddAuthToken(context.WithValue(context.Background(), "token", tok), req) + pilosa.AddAuthToken( + context.WithValue(context.Background(), + authn.ContextValueAccessToken, + tok, + ), + &req.Header, + ) if got := req.Header.Get("Authorization"); got != tok { t.Fatalf("got '%v', expected '%v'", got, tok) } @@ -1609,7 +1615,7 @@ func TestAddAuthToken(t *testing.T) { t.Fatalf("unexpected error: %v", err) } ogIP := "10.0.0.1" - pilosa.AddAuthToken(context.WithValue(context.Background(), pilosa.OriginalIPHeader, ogIP), req) + pilosa.AddAuthToken(context.WithValue(context.Background(), pilosa.OriginalIPHeader, ogIP), &req.Header) if got := req.Header.Get(pilosa.OriginalIPHeader); got != ogIP { t.Fatalf("got '%v', expected '%v'", got, ogIP) } diff --git a/server/grpc.go b/server/grpc.go index 99099f761..cad41bc4f 100644 --- a/server/grpc.go +++ b/server/grpc.go @@ -423,6 +423,9 @@ func (h *GRPCHandler) CreateIndex(ctx context.Context, req *pb.CreateIndexReques if err != nil { return nil, errToStatusError(err) } + if err := grpc.SendHeader(ctx, metadata.MD{}); err != nil { + return nil, errToStatusError(err) + } return &pb.CreateIndexResponse{}, nil } @@ -452,6 +455,9 @@ func (h *GRPCHandler) GetIndex(ctx context.Context, req *pb.GetIndexRequest) (*p return &pb.GetIndexResponse{Index: &pb.Index{Name: index.Name}}, nil } } + if err := grpc.SendHeader(ctx, metadata.MD{}); err != nil { + return nil, errToStatusError(err) + } return nil, status.Error(codes.NotFound, fmt.Sprintf("Index with name %s not found", req.Name)) } @@ -481,6 +487,11 @@ func (h *GRPCHandler) GetIndexes(ctx context.Context, req *pb.GetIndexesRequest) indexes = append(indexes, &pb.Index{Name: index.Name}) } } + + if err := grpc.SendHeader(ctx, metadata.MD{}); err != nil { + return nil, errToStatusError(err) + } + return &pb.GetIndexesResponse{Indexes: indexes}, nil } @@ -496,6 +507,9 @@ func (h *GRPCHandler) DeleteIndex(ctx context.Context, req *pb.DeleteIndexReques if err != nil { return nil, errToStatusError(err) } + if err := grpc.SendHeader(ctx, metadata.MD{}); err != nil { + return nil, errToStatusError(err) + } return &pb.DeleteIndexResponse{}, nil } @@ -1602,7 +1616,7 @@ func NewGRPCServer(opts ...grpcServerOption) (*grpcServer, error) { // reset the molecula-chip cookie just in case the token was refreshed md, ok := metadata.FromIncomingContext(ctx) if uinfo, yeah := ctx.Value("userinfo").(*authn.UserInfo); ok && yeah { - server.auth.SetGRPCMetadata(ctx, md, uinfo.Token) + server.auth.SetGRPCMetadata(ctx, md, uinfo.Token, uinfo.RefreshToken) } return handler(ctx, req) }, @@ -1616,7 +1630,7 @@ func NewGRPCServer(opts ...grpcServerOption) (*grpcServer, error) { // reset the molecula-chip cookie just in case the token was refreshed md, ok := metadata.FromIncomingContext(ctx) if uinfo, yeah := ctx.Value("userinfo").(*authn.UserInfo); ok && yeah { - server.auth.SetGRPCMetadata(ctx, md, uinfo.Token) + server.auth.SetGRPCMetadata(ctx, md, uinfo.Token, uinfo.RefreshToken) } return handler(srv, &wrappedStream{ss, ctx}) }, @@ -1689,29 +1703,50 @@ func Valid(ctx context.Context, auth *authn.Auth) (context.Context, error) { return ctx, status.Errorf(codes.InvalidArgument, "missing metadata") } - authorization, ok := md["authorization"] - if !ok { - c, there := md["cookie"] - if !there { - return ctx, status.Errorf(codes.InvalidArgument, "missing authorization token") - } - cookies := strings.Split(c[0], "; ") - for _, cookie := range cookies { - if strings.HasPrefix(cookie, authn.CookieName) { - authorization = strings.Split(cookie, authn.CookieName+"=")[1:] - break - } - } - } - if len(authorization) == 0 { + access, refresh := getTokensFromMetadata(md) + if access == "" { return ctx, status.Errorf(codes.InvalidArgument, "missing authorization token") } - - token := strings.TrimPrefix(authorization[0], "Bearer ") - uinfo, err := auth.Authenticate(ctx, token) + uinfo, err := auth.Authenticate(access, refresh) if err != nil { return ctx, status.Errorf(codes.Unauthenticated, err.Error()) } return context.WithValue(ctx, "userinfo", uinfo), nil } + +func getTokensFromMetadata(md metadata.MD) (string, string) { + // We check lowercase and uppercase because some GRPC clients lowercase metadata + // names. This is the only place we get tokens from metadata in GRPC calls. + access, ok := md["authorization"] + if !ok { + access, ok = md["Authorization"] + } + + refresh, ok2 := md[strings.ToLower(authn.RefreshHeaderName)] + if !ok2 { + refresh, ok2 = md[authn.RefreshHeaderName] + } + + if !ok || !ok2 { + if cookies, there := md["cookie"]; there { + for _, cookie := range cookies { + if strings.HasPrefix(cookie, authn.AccessCookieName+"=") && len(access) == 0 { + access = strings.Split(cookie, authn.AccessCookieName+"=")[1:] + } else if strings.HasPrefix(cookie, authn.RefreshCookieName+"=") && len(refresh) == 0 { + refresh = strings.Split(cookie, authn.RefreshCookieName+"=")[1:] + } + if len(access) > 0 && len(refresh) > 0 { + break + } + } + } + } + if len(access) == 0 { + access = []string{""} + } + if len(refresh) == 0 { + refresh = []string{""} + } + return strings.TrimPrefix(access[0], "Bearer "), refresh[0] +} diff --git a/server/grpc_internal_test.go b/server/grpc_internal_test.go new file mode 100644 index 000000000..024664a13 --- /dev/null +++ b/server/grpc_internal_test.go @@ -0,0 +1,158 @@ +package server + +import ( + "context" + "strings" + "testing" + + "github.com/molecula/featurebase/v3/authn" + "github.com/molecula/featurebase/v3/logger" + "google.golang.org/grpc" + "google.golang.org/grpc/metadata" +) + +func TestGetTokensFromMetadata(t *testing.T) { + for name, test := range map[string]struct { + access string + refresh string + setCookie bool + md metadata.MD + }{ + "empty": { + access: "", + refresh: "", + md: metadata.MD{}, + }, + "inTheCookieNoRefresh": { + access: "something", + refresh: "", + md: metadata.MD{}, + setCookie: true, + }, + "inTheCookieYesRefresh": { + access: "something", + refresh: "somethingElse", + md: metadata.MD{}, + setCookie: true, + }, + "otherCookies": { + access: "something", + refresh: "somethingElse", + setCookie: true, + md: metadata.MD{ + "cookie": []string{ + "okay=okay", + "blah=blah", + }, + }, + }, + "inTheHeaderNoRefresh": { + access: "something", + refresh: "", + md: metadata.MD{"authorization": []string{"something"}}, + }, + "inTheHeaderYesRefresh": { + access: "something", + refresh: "somethingElse", + md: metadata.MD{ + "authorization": []string{"something"}, + strings.ToLower(authn.RefreshHeaderName): []string{"somethingElse"}, + }, + }, + "inTheHeaderYesRefreshCaps": { + access: "something", + refresh: "somethingElse", + md: metadata.MD{ + "authorization": []string{"something"}, + authn.RefreshHeaderName: []string{"somethingElse"}, + }, + }, + } { + t.Run(name, func(t *testing.T) { + if test.setCookie { + a := NewTestAuth(t) + ctx := grpc.NewContextWithServerTransportStream( + metadata.NewIncomingContext(context.TODO(), + test.md, + ), + NewServerTransportStream(), + ) + err := a.SetGRPCMetadata(ctx, test.md, test.access, test.refresh) + if err != nil { + t.Errorf("unexpected error setting GRPC metadata: %v", err) + } + } + accessGot, refreshGot := getTokensFromMetadata(test.md) + if accessGot != test.access { + t.Errorf("access: expected %v, got %v", test.access, accessGot) + } + if refreshGot != test.refresh { + t.Errorf("refresh: expected %v, got %v", test.refresh, refreshGot) + } + }) + } +} + +// This type is used for mocking ServerTransportStreams in tests +type ServerTransportStream struct { + md metadata.MD + method string +} + +func NewServerTransportStream() *ServerTransportStream { + return &ServerTransportStream{ + md: metadata.MD{}, + method: "test", + } +} + +func (s *ServerTransportStream) Method() string { + return s.method +} + +func (s *ServerTransportStream) SetHeader(md metadata.MD) error { + s.md = md + return nil +} + +func (s *ServerTransportStream) SendHeader(md metadata.MD) error { + _ = md + return nil +} + +func (s *ServerTransportStream) SetTrailer(md metadata.MD) error { + _ = md + return nil +} + +func NewTestAuth(t *testing.T) *authn.Auth { + t.Helper() + var ( + ClientID = "e9088663-eb08-41d7-8f65-efb5f54bbb71" + ClientSecret = "DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF" + AuthorizeURL = "https://login.microsoftonline.com/4a137d66-d161-4ae4-b1e6-07e9920874b8/oauth2/v2.0/authorize" + TokenURL = "https://login.microsoftonline.com/4a137d66-d161-4ae4-b1e6-07e9920874b8/oauth2/v2.0/token" + GroupEndpointURL = "https://graph.microsoft.com/v1.0/me/transitiveMemberOf/microsoft.graph.group?$count=true" + LogoutURL = "https://login.microsoftonline.com/common/oauth2/v2.0/logout" + Scopes = []string{"https://graph.microsoft.com/.default", "offline_access"} + Key = "DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF" + ) + + a, err := authn.NewAuth( + logger.NopLogger, + "http://localhost:10101/", + Scopes, + AuthorizeURL, + TokenURL, + GroupEndpointURL, + LogoutURL, + ClientID, + ClientSecret, + Key, + []string{}, + ) + if err != nil { + t.Fatalf("building auth object%s", err) + } + return a +}