mirror of
https://github.com/featurebasedb/featurebase.git
synced 2026-10-07 03:17:50 +00:00
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:
parent
3986e202bf
commit
60e6900c2e
14 changed files with 624 additions and 311 deletions
|
|
@ -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]
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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()) {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
}
|
||||
|
|
|
|||
158
server/grpc_internal_test.go
Normal file
158
server/grpc_internal_test.go
Normal 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
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue