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
This commit is contained in:
reesporte 2022-05-20 16:12:27 -05:00 • committed by GitHub
parent 3986e202bf
commit 60e6900c2e
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
14 changed files with 624 additions and 311 deletions

View file

@ -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]
}

View file

@ -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)
}
}
})

View file

@ -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.

View file

@ -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 {

View file

@ -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()) {

View file

@ -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)
}

View file

@ -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
}

View file

@ -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",

View file

@ -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"}

View file

@ -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)

View file

@ -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))

View file

@ -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)
}

View file

@ -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]
}

View file

@ -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
}