addresses multiple authn/z tickets

* fb-998 - authn/z enabled in handlers (kitchen-sink ticket)
    - authorization is enabled through the use of a bearer token (using header "Authorization")
    - authorization may occur through the use of an "Authorization" header or "molecula-chip" cookie
    - ui is updated for changes to handler
* fb-1131 - protect grpc endpoints
    - GRPC endpoints now check authorization if auth is enabled
* fb-1129 - inter-node communication
    - the following endpoints use the secretKey for authentication:
        - /internal/cluster/message: POST
        - /internal/translate/data: GET, POST

* added test to api_test.go (TestAuth_MultiNode) testing various auth/permissions stuff on a multi-node cluster

not included:
    - fb-1130 - filter response of endpoints
    - fb-1109 - improved audit logging

@jaffee [are you not entertained](https://www.youtube.com/watch?v=mutgotxrcqg)

Co-authored-by: souhailanoor <90720110+souhailanoor@users.noreply.github.com>
Co-authored-by: tgruben <tgruben@gmail.com>
Co-authored-by: 54mir <48686912+54mir@users.noreply.github.com>
Co-authored-by: kcrodgers24 <49999391+kcrodgers24@users.noreply.github.com>
This commit is contained in:
reesporte 2022-01-13 14:08:45 -06:00
parent 4a1e53421e
commit cf2410fea6
21 changed files with 1123 additions and 510 deletions

View file

@ -311,7 +311,7 @@ require-%:
install-build-deps: install-protoc-gen-gofast install-protoc install-statik install-stringer install-peg
install-statik:
go get -u github.com/rakyll/statik
go install github.com/rakyll/statik@latest
install-stringer:
GO111MODULE=off $(GO) get -u golang.org/x/tools/cmd/stringer

View file

@ -4,17 +4,23 @@ package pilosa_test
import (
"bytes"
"context"
"encoding/hex"
"errors"
"fmt"
"io"
"math"
"math/rand"
"os"
"path/filepath"
"reflect"
"sort"
"strings"
"testing"
"time"
"github.com/golang-jwt/jwt"
pilosa "github.com/molecula/featurebase/v2"
"github.com/molecula/featurebase/v2/authn"
"github.com/molecula/featurebase/v2/boltdb"
"github.com/molecula/featurebase/v2/http"
"github.com/molecula/featurebase/v2/server"
@ -1437,3 +1443,215 @@ func TestAPI_RBFDebugInfo(t *testing.T) {
t.Fatal("expected info")
}
}
// makeUser makes an authnUserInfo from groups and a name and a secret key
func makeUser(t *testing.T, groups []authn.Group, name, secret string) *authn.UserInfo {
tkn := jwt.New(jwt.SigningMethodHS256)
claims := tkn.Claims.(jwt.MapClaims)
groupString, err := authn.ToGob64(groups)
if err != nil {
t.Fatalf("gobbing groups %v", err)
}
claims["molecula-idp-groups"] = groupString
claims["oid"] = "42"
claims["name"] = name
secretKey, _ := hex.DecodeString(secret)
validToken, err := tkn.SignedString(secretKey)
if err != nil {
t.Fatalf("signing string %v", err)
}
validToken = "Bearer " + validToken
return &authn.UserInfo{
UserID: "fake" + name,
UserName: name,
Groups: groups,
Token: validToken,
Expiry: time.Time{},
}
}
func TestAuth_MultiNode(t *testing.T) {
// create permissions file
permissions := `
"user-groups":
"dca35310-ecda-4f23-86cd-876aee55906b":
"test": "read"
"dca35310-ecda-4f23-86cd-876aee55906f":
"test": "write"
admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
// authentication on
auth := server.Auth{
Enable: true,
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"},
SecretKey: "DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF",
PermissionsFile: writeTestFile(t, "permissions.yaml", permissions),
QueryLogPath: writeTestFile(t, "queryLog.log", ""),
}
config := server.NewConfig()
config.Auth = auth
// set up TLS certificates
localhostCert := `-----BEGIN CERTIFICATE-----
MIICEzCCAXygAwIBAgIQMIMChMLGrR+QvmQvpwAU6zANBgkqhkiG9w0BAQsFADAS
MRAwDgYDVQQKEwdBY21lIENvMCAXDTcwMDEwMTAwMDAwMFoYDzIwODQwMTI5MTYw
MDAwWjASMRAwDgYDVQQKEwdBY21lIENvMIGfMA0GCSqGSIb3DQEBAQUAA4GNADCB
iQKBgQDuLnQAI3mDgey3VBzWnB2L39JUU4txjeVE6myuDqkM/uGlfjb9SjY1bIw4
iA5sBBZzHi3z0h1YV8QPuxEbi4nW91IJm2gsvvZhIrCHS3l6afab4pZBl2+XsDul
rKBxKKtD1rGxlG4LjncdabFn9gvLZad2bSysqz/qTAUStTvqJQIDAQABo2gwZjAO
BgNVHQ8BAf8EBAMCAqQwEwYDVR0lBAwwCgYIKwYBBQUHAwEwDwYDVR0TAQH/BAUw
AwEB/zAuBgNVHREEJzAlggtleGFtcGxlLmNvbYcEfwAAAYcQAAAAAAAAAAAAAAAA
AAAAATANBgkqhkiG9w0BAQsFAAOBgQCEcetwO59EWk7WiJsG4x8SY+UIAA+flUI9
tyC4lNhbcF2Idq9greZwbYCqTTTr2XiRNSMLCOjKyI7ukPoPjo16ocHj+P3vZGfs
h1fIw3cSS2OolhloGw/XM6RWPWtPAlGykKLciQrBru5NAPvCMsb/I1DAceTiotQM
fblo6RBxUQ==
-----END CERTIFICATE-----`
localhostKey := `-----BEGIN RSA PRIVATE KEY-----
MIICXgIBAAKBgQDuLnQAI3mDgey3VBzWnB2L39JUU4txjeVE6myuDqkM/uGlfjb9
SjY1bIw4iA5sBBZzHi3z0h1YV8QPuxEbi4nW91IJm2gsvvZhIrCHS3l6afab4pZB
l2+XsDulrKBxKKtD1rGxlG4LjncdabFn9gvLZad2bSysqz/qTAUStTvqJQIDAQAB
AoGAGRzwwir7XvBOAy5tM/uV6e+Zf6anZzus1s1Y1ClbjbE6HXbnWWF/wbZGOpet
3Zm4vD6MXc7jpTLryzTQIvVdfQbRc6+MUVeLKwZatTXtdZrhu+Jk7hx0nTPy8Jcb
uJqFk541aEw+mMogY/xEcfbWd6IOkp+4xqjlFLBEDytgbIECQQDvH/E6nk+hgN4H
qzzVtxxr397vWrjrIgPbJpQvBsafG7b0dA4AFjwVbFLmQcj2PprIMmPcQrooz8vp
jy4SHEg1AkEA/v13/5M47K9vCxmb8QeD/asydfsgS5TeuNi8DoUBEmiSJwma7FXY
fFUtxuvL7XvjwjN5B30pNEbc6Iuyt7y4MQJBAIt21su4b3sjXNueLKH85Q+phy2U
fQtuUE9txblTu14q3N7gHRZB4ZMhFYyDy8CKrN2cPg/Fvyt0Xlp/DoCzjA0CQQDU
y2ptGsuSmgUtWj3NM9xuwYPm+Z/F84K6+ARYiZ6PYj013sovGKUFfYAqVXVlxtIX
qyUBnu3X9ps8ZfjLZO7BAkEAlT4R5Yl6cGhaJQYZHOde3JEMhNRcVFMO8dJDaFeo
f9Oeos0UUothgiDktdQHxdNEwLjQf7lJJBzV+5OtwswCWA==
-----END RSA PRIVATE KEY-----`
config.TLS.CertificateKeyPath = writeTestFile(t, "certKey.pem", localhostKey)
config.TLS.CertificatePath = writeTestFile(t, "cert.pem", localhostCert)
c := test.MustRunCluster(t, 3,
[]server.CommandOption{
server.OptCommandServerOptions(
pilosa.OptServerNodeID("node0"),
pilosa.OptServerClusterHasher(&test.ModHasher{}),
),
server.OptCommandConfig(config),
},
[]server.CommandOption{
server.OptCommandServerOptions(
pilosa.OptServerNodeID("node1"),
pilosa.OptServerClusterHasher(&test.ModHasher{}),
),
server.OptCommandConfig(config),
},
[]server.CommandOption{
server.OptCommandServerOptions(
pilosa.OptServerNodeID("node2"),
pilosa.OptServerClusterHasher(&test.ModHasher{}),
),
server.OptCommandConfig(config),
},
)
defer c.Close()
adminCtx := context.WithValue(
context.Background(),
"userinfo",
makeUser(t, []authn.Group{{GroupID: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe", GroupName: "adminGroup"}}, "admin", config.Auth.SecretKey),
)
readCtx := context.WithValue(
context.Background(),
"userinfo",
makeUser(t, []authn.Group{{GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: "readGroup"}}, "reader", config.Auth.SecretKey),
)
writeCtx := context.WithValue(
context.Background(),
"userinfo",
makeUser(t, []authn.Group{{GroupID: "dca35310-ecda-4f23-86cd-876aee55906f", GroupName: "writeGroup"}}, "writer", config.Auth.SecretKey),
)
primaryAPI := c.GetPrimary().API
// needs internal/cluster/message
indexName := "test"
_, err := primaryAPI.CreateIndex(adminCtx, indexName, pilosa.IndexOptions{})
if err != nil {
t.Fatalf("creating index: %v", err)
}
// needs internal/translate/data
fieldName := "f"
_, err = primaryAPI.CreateField(adminCtx, indexName, fieldName, pilosa.OptFieldTypeSet(pilosa.DefaultCacheType, 100))
if err != nil {
t.Fatalf("creating field: %v", err)
}
_, err = primaryAPI.Query(readCtx, &pilosa.QueryRequest{
Index: indexName,
Query: fmt.Sprintf(`Set(1, %s=1)`, fieldName),
})
if err == nil {
t.Fatalf("readCtx should not be able to set bits")
}
_, err = primaryAPI.Query(writeCtx, &pilosa.QueryRequest{
Index: indexName,
Query: fmt.Sprintf(`Set(1, %s=1)`, fieldName),
})
if err != nil {
t.Fatalf("writeCtx should be able to set bits: %v", err)
}
_, err = primaryAPI.Query(adminCtx, &pilosa.QueryRequest{
Index: indexName,
Query: fmt.Sprintf(`Set(1, %s=1)`, fieldName),
})
if err != nil {
t.Fatalf("adminCtx should be able to set bits: %v", err)
}
_, err = primaryAPI.Query(readCtx, &pilosa.QueryRequest{
Index: indexName,
Query: fmt.Sprintf(`Count(Row(%s=1))`, fieldName),
})
if err != nil {
t.Fatalf("readCtx should be able read: %v", err)
}
_, err = primaryAPI.Query(writeCtx, &pilosa.QueryRequest{
Index: indexName,
Query: fmt.Sprintf(`Count(Row(%s=1))`, fieldName),
})
if err != nil {
t.Fatalf("writeCtx should be able read: %v", err)
}
_, err = primaryAPI.Query(adminCtx, &pilosa.QueryRequest{
Index: indexName,
Query: fmt.Sprintf(`Count(Row(%s=1))`, fieldName),
})
if err != nil {
t.Fatalf("adminCtx should be able read: %v", err)
}
}
func writeTestFile(t *testing.T, filename, content string) string {
t.Helper()
fname := filepath.Join(t.TempDir(), filename)
f, err := os.Create(fname)
if err != nil {
t.Fatalf("could not create file %v with err %v", filename, err)
}
_, err = io.WriteString(f, content)
if err != nil {
t.Fatalf("could not write string %v", err)
}
defer f.Close()
return fname
}

View file

@ -4,7 +4,9 @@
package authn
import (
"context"
"bytes"
"encoding/base64"
"encoding/gob"
"encoding/hex"
"encoding/json"
"fmt"
@ -13,20 +15,20 @@ import (
"time"
"github.com/golang-jwt/jwt"
"github.com/gorilla/securecookie"
"github.com/molecula/featurebase/v2/logger"
"github.com/pkg/errors"
"golang.org/x/oauth2"
)
func init() {
gob.Register([]Group{})
}
// Auth holds state, configuration, and utilities needed for authentication.
type Auth struct {
logger logger.Logger
cookieName string
refreshWithin time.Duration
hashKey []byte
blockKey []byte
secure *securecookie.SecureCookie
secretKey []byte
groupEndpoint string
logoutEndpoint string
fbURL string // fbURL is the domain FB is hosted on, used for post logout redirection
@ -34,11 +36,10 @@ type Auth struct {
}
// NewAuth instantiates and returns a new Auth struct
func NewAuth(logger logger.Logger, url string, scopes []string, authURL, tokenURL, groupEndpoint, logout, clientID, clientSecret, hashKey, blockKey string) (*Auth, error) {
auth := &Auth{
func NewAuth(logger logger.Logger, url string, scopes []string, authURL, tokenURL, groupEndpoint, logout, clientID, clientSecret, secretKey string) (auth *Auth, err error) {
auth = &Auth{
logger: logger,
cookieName: "molecula-chip",
refreshWithin: 15 * time.Minute,
groupEndpoint: groupEndpoint,
logoutEndpoint: logout,
fbURL: url,
@ -53,72 +54,111 @@ func NewAuth(logger logger.Logger, url string, scopes []string, authURL, tokenUR
},
},
}
var err error
if auth.hashKey, err = decodeHex(hashKey); err != nil {
return nil, errors.Wrap(err, "decoding hash key")
}
if auth.blockKey, err = decodeHex(blockKey); err != nil {
return nil, errors.Wrap(err, "decoding block key")
if auth.secretKey, err = decodeHex(secretKey); err != nil {
return nil, errors.Wrap(err, "decoding secret key")
}
auth.secure = securecookie.New(auth.hashKey, auth.blockKey)
return auth, nil
}
// AuthContext holds the value of an authenticated user's cookie
type AuthContext struct {
UserID string
UserName string
GroupMembership []Group
Token *oauth2.Token
func (a Auth) SecretKey() []byte {
return a.secretKey
}
// 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"`
}
// Group holds group information for an authenticated user
type Group struct {
UserID string
GroupID string `json:"id"`
GroupName string `json:"displayName"`
}
// Groups holds a slice of Group informations for marshalling from Json
// ToGob64 encodes a []Group to a string, returning string and nil on success (todd's idea)
// it has to be a string bc we're using it in a jwt.MapClaims which needs string-y things
func ToGob64(m []Group) (string, error) {
var b bytes.Buffer
if err := gob.NewEncoder(&b).Encode(m); err != nil {
return "", err
}
return base64.StdEncoding.EncodeToString(b.Bytes()), nil
}
// FromGob64 converts a previously encoded []Group from a string to a []Group
// it has to be a string bc we're using it in a jwt.MapClaims which needs string-y things
func FromGob64(gobbed string) ([]Group, error) {
m := []Group{}
by, err := base64.StdEncoding.DecodeString(gobbed)
if err != nil {
return nil, err
}
b := bytes.Buffer{}
b.Write(by)
d := gob.NewDecoder(&b)
err = d.Decode(&m)
if err != nil {
return nil, err
}
return m, nil
}
// Groups holds a slice of Group for marshalling from Json
type Groups struct {
Groups []Group `json:"value"`
}
// UserInfo holds user information for an authenticated user
type UserInfo struct {
UserID string `json:"userid"`
UserName string `json:"username"`
}
// Authenticate reads the authentication cookie from a request, returning the
// user's group memberships on success. If the cookie is not present or has expired,
// Authenticate redirects the user to sign in. If the cookie is within the
// refresh window of expiring, the cookie is refreshed, and the updated group
// membership is returned.
func (a *Auth) Authenticate(w http.ResponseWriter, r *http.Request) ([]Group, error) {
cookie, err := a.readCookie(w, r)
if err != nil {
http.Redirect(w, r, "/signin", http.StatusTemporaryRedirect)
return nil, err
// Authenticate takes in a bearer token `bearer` and returns UserInfo from that token
func (a *Auth) Authenticate(bearer string) (*UserInfo, error) {
// parse the bearer token into a jwt.Token
token, err := jwt.Parse(bearer, func(token *jwt.Token) (interface{}, error) {
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
}
return a.secretKey, nil
})
if token == nil || token.Claims == nil || err != nil || !token.Valid {
return nil, errors.Wrap(err, fmt.Sprintf("%#v parsing jwt claims from access tokens", token))
}
if cookie.Token.Expiry.Before(time.Now().Add(a.refreshWithin)) {
err = a.refreshToken(w, cookie)
if err != nil {
a.logger.Errorf("refreshing access token: ", err)
if cookie.Token.Expiry.Before(time.Now()) {
http.Redirect(w, r, "/signin", http.StatusTemporaryRedirect)
return nil, err
userInfo := UserInfo{}
// check that token does not expire now
switch claimType := token.Claims.(type) {
case jwt.MapClaims:
if exp, ok := claimType["exp"]; ok {
var e int64
switch expType := exp.(type) {
case float64:
e = int64(expType)
case json.Number:
e, _ = expType.Int64()
}
if e <= time.Now().Unix() {
return nil, fmt.Errorf("token expired")
}
}
}
if len(cookie.GroupMembership) == 0 {
return nil, errors.New("user is not part of any groups in identity provider")
}
return cookie.GroupMembership, nil
userInfo.UserID = claimType["oid"].(string)
userInfo.UserName = claimType["name"].(string)
userInfo.Token = bearer
g := claimType["molecula-idp-groups"].(string)
groups, err := FromGob64(g)
if err != nil {
return nil, errors.Wrap(err, "decoding groups")
}
userInfo.Groups = groups
default:
return nil, fmt.Errorf("could not parse jwt claims of type %T, expected jwt.MapClaims", claimType)
}
return &userInfo, nil
}
// Login redirects a user to login to their configured oAuth authorize endpoint
@ -129,46 +169,41 @@ func (a *Auth) Login(w http.ResponseWriter, r *http.Request) {
// Logout clears out user cookie and redirects user to IdP's logout endpoint
func (a *Auth) Logout(w http.ResponseWriter, r *http.Request) {
http.SetCookie(w, a.getEmptyCookie())
http.SetCookie(w, &http.Cookie{
Name: a.cookieName,
Value: "",
Path: "/",
Secure: true,
HttpOnly: true,
SameSite: http.SameSiteStrictMode,
Expires: time.Unix(0, 0),
})
redirect := fmt.Sprintf("%s?post_logout_redirect_uri=%s/", a.logoutEndpoint, a.fbURL)
http.Redirect(w, r, redirect, http.StatusTemporaryRedirect)
}
// Redirect handles the oAuth /redirect endpoint. It gets user information from
// the identity provider and sets a secure cookie holding the user information.
// the identity provider and sets a secure cookie holding the user information
// signed by featurebase.
func (a *Auth) Redirect(w http.ResponseWriter, r *http.Request) {
code := r.FormValue("code")
token, err := a.getToken(r, code)
if err != nil {
a.logger.Warnf("getting token from IdP: %+v", err)
http.Error(w, "Bad Request: 400", http.StatusBadRequest)
http.Error(w, "Bad Request", http.StatusBadRequest)
return
}
cv, err := a.newAuthContext(token)
if err != nil || cv == nil {
a.logger.Warnf("creating cookie: %+v", err)
http.Error(w, "Bad Request: 400", http.StatusBadRequest)
return
}
a.setCookie(w, cv)
http.Redirect(w, r, "/", http.StatusTemporaryRedirect)
}
// GetUserInfo reads user's cookie and returns their username and userId
func (a *Auth) GetUserInfo(w http.ResponseWriter, r *http.Request) *UserInfo {
var resp UserInfo
cookie, err := a.readCookie(w, r)
// with vitamin A!
enrichedTkn, err := a.addGroupMembership(token.AccessToken)
if err != nil {
a.logger.Warnf("was not able to read cookie for req: %+v", r)
return &resp
a.logger.Warnf("enriching token with group membership: %+v", err)
http.Error(w, "Bad Request", http.StatusBadRequest)
return
}
return &UserInfo{
UserID: cookie.UserID,
UserName: cookie.UserName,
}
a.setCookie(w, enrichedTkn, token.Expiry)
http.Redirect(w, r, "/", http.StatusTemporaryRedirect)
}
// getToken exhanges authorization code for an oAuth2 token
@ -180,124 +215,88 @@ func (a *Auth) getToken(r *http.Request, code string) (*oauth2.Token, error) {
return token, nil
}
// newAuthContext parses a jwt `token` and returns relevant information in a cookie value struct
func (a *Auth) newAuthContext(token *oauth2.Token) (*AuthContext, error) {
if token == nil {
return nil, errors.New("baking cookie due to nil token")
}
if token.AccessToken == "" {
return nil, errors.New("no access token provided")
}
// We are using ParseUnverified here because we're using the OAuth2.0 authZ code flow
// which assumes that the IdP gives good responses. This means that if the IdP is
// insecure, then we are too. But that's the way OAuth works, unfortunately.
// Also, we assume the jwt is not tampered with bc we communicate with the IdP over HTTPS only.
accessParsed, _, err := new(jwt.Parser).ParseUnverified(token.AccessToken, jwt.MapClaims{})
if accessParsed == nil || accessParsed.Claims == nil || err != nil {
return nil, errors.Wrap(err, fmt.Sprintf("%v parsing jwt claims from access tokens", accessParsed))
}
claims := accessParsed.Claims.(jwt.MapClaims)
groups, err := a.getGroupMembership(token)
// addGroupMembership is only called in `a.Redirect`. It adds groups to a jwt's
// claims, and signs it using `a.secretKey`.
func (a *Auth) addGroupMembership(token string) (string, error) {
g, err := a.getGroups(token)
if err != nil {
return nil, errors.Wrap(err, "getting group membership")
return "", err
}
// not needed at this point in the logic and makes the encoded cookie too large
token.AccessToken = ""
return &AuthContext{
UserID: claims["oid"].(string),
UserName: claims["name"].(string),
GroupMembership: groups.Groups,
Token: token,
}, nil
// parse token into jwt
unenriched, _, err := new(jwt.Parser).ParseUnverified(token, jwt.MapClaims{})
if unenriched == nil || unenriched.Claims == nil || err != nil {
return "", errors.Wrap(err, fmt.Sprintf("%v parsing jwt claims from access tokens", token))
}
enriched := jwt.New(jwt.SigningMethodHS256)
enriched.Claims = unenriched.Claims
var tokenStr string
// parse groups into string format
switch claims := enriched.Claims.(type) {
case jwt.MapClaims:
groupString, err := ToGob64(g)
if err != nil {
return "", errors.Wrap(err, "failed to serialize groups")
}
// stick it into jwt claims
claims["molecula-idp-groups"] = groupString
// get stringified and signed jwt
tokenStr, err = enriched.SignedString(a.secretKey)
if err != nil {
return "", errors.Wrap(err, "signing jwt")
}
default:
return "", fmt.Errorf("could not parse jwt claims of type %T, expected jwt.MapClaims", claims)
}
return tokenStr, nil
}
// getGroupMembership uses a oauth2 token to retrieve group membership information from IdP
func (a *Auth) getGroupMembership(token *oauth2.Token) (Groups, error) {
// getGroups gets the group membership for a given token from configured IdP
func (a *Auth) getGroups(token string) ([]Group, error) {
var groups Groups
req, err := http.NewRequest("GET", a.groupEndpoint, nil)
if err != nil {
return groups, errors.Wrap(err, "creating new request to group endpoint")
return groups.Groups, errors.Wrap(err, "creating new request to group endpoint")
}
req.Header.Add("Authorization", fmt.Sprintf("Bearer %s", token.AccessToken))
req.Header.Add("Authorization", fmt.Sprintf("Bearer %s", token))
response, err := http.DefaultClient.Do(req)
if err != nil {
return groups, errors.Wrap(err, "getting group membership info")
return groups.Groups, errors.Wrap(err, "getting group membership info")
}
defer response.Body.Close()
rawGroups, err := io.ReadAll(response.Body)
if err != nil {
return groups, errors.Wrap(err, "failed reading group membership response")
return groups.Groups, errors.Wrap(err, "failed reading group membership response")
}
if err = json.Unmarshal(rawGroups, &groups); err != nil {
return groups, errors.Wrap(err, "failed unmarshalling group membership response")
return groups.Groups, errors.Wrap(err, "failed unmarshalling group membership response")
}
return groups, nil
return groups.Groups, nil
}
// readCookie decodes an encrypted and signed cookie and returns the contained info
func (a *Auth) readCookie(w http.ResponseWriter, r *http.Request) (*AuthContext, error) {
cookie, err := r.Cookie(a.cookieName)
if err != nil {
return nil, errors.Wrap(err, "cookie not found")
}
var value AuthContext
err = a.secure.Decode(a.cookieName, cookie.Value, &value)
if err != nil {
http.SetCookie(w, a.getEmptyCookie())
return nil, errors.Wrap(err, "decoding cookie")
}
return &value, nil
}
func (a *Auth) setCookie(w http.ResponseWriter, cookie *AuthContext) error {
encoded, err := a.secure.Encode(a.cookieName, cookie)
if err != nil {
return errors.Wrap(err, "encoding AuthContext")
}
func (a *Auth) setCookie(w http.ResponseWriter, token string, expiry time.Time) error {
http.SetCookie(w, &http.Cookie{
Name: a.cookieName,
Value: encoded,
Value: token,
Path: "/",
Secure: true,
HttpOnly: true,
SameSite: http.SameSiteStrictMode,
Expires: cookie.Token.Expiry,
Expires: expiry,
})
return nil
}
func (a *Auth) refreshToken(w http.ResponseWriter, cookie *AuthContext) error {
if cookie.Token.RefreshToken == "" {
return errors.New("no refresh token found, check auth scopes to see if refresh tokens are being provided by your IdP")
}
tokenSource := a.oAuthConfig.TokenSource(context.Background(), cookie.Token)
newToken, err := tokenSource.Token()
if err != nil {
return errors.Wrap(err, "refreshing token")
}
if newToken.Expiry != cookie.Token.Expiry {
cv, err := a.newAuthContext(newToken)
if err != nil {
return errors.Wrap(err, "creating cookie value from token")
}
a.setCookie(w, cv)
}
return nil
}
func decodeHex(hexstr string) ([]byte, error) {
data, err := hex.DecodeString(hexstr)
if err != nil {
@ -308,15 +307,3 @@ func decodeHex(hexstr string) ([]byte, error) {
}
return data, nil
}
func (a *Auth) getEmptyCookie() *http.Cookie {
return &http.Cookie{
Name: a.cookieName,
Value: "",
Path: "/",
Secure: true,
HttpOnly: true,
SameSite: http.SameSiteStrictMode,
}
}

View file

@ -8,7 +8,6 @@ import (
"time"
"github.com/molecula/featurebase/v2/logger"
"golang.org/x/oauth2"
)
func TestAuth(t *testing.T) {
@ -35,37 +34,14 @@ func TestAuth(t *testing.T) {
ClientID,
ClientSecret,
Key,
Key,
)
if err != nil {
t.Errorf("building auth object%s", err)
}
tokenNoAT := oauth2.Token{
TokenType: "Bearer",
RefreshToken: "abcdef",
Expiry: time.Now().Add(time.Hour),
}
tokenAT := oauth2.Token{
TokenType: "Bearer",
RefreshToken: "abcdef",
AccessToken: "aasdf",
Expiry: time.Now().Add(time.Hour),
}
grp := Group{
UserID: "snowstorm",
GroupID: "abcd123-A",
GroupName: "Romantic Painters",
}
validCV := AuthContext{
UserID: "snowstorm",
UserName: "J.M.W. Turner",
GroupMembership: []Group{grp},
Token: &tokenAT,
}
t.Run("SetCookie", func(t *testing.T) {
w := httptest.NewRecorder()
err := a.setCookie(w, &validCV)
err := a.setCookie(w, "a cookie value", time.Now().Add(time.Hour))
if err != nil {
t.Fatalf("expected no errors, got: %v", err)
}
@ -79,12 +55,6 @@ func TestAuth(t *testing.T) {
}
})
t.Run("GetEmptyCookie", func(t *testing.T) {
c := a.getEmptyCookie()
if c.Value != "" {
t.Fatalf("expected empty cookie, got: %+v", c.Value)
}
})
t.Run("KeyLength", func(t *testing.T) {
_, err := NewAuth(
logger.NewStandardLogger(os.Stdout),
@ -96,25 +66,10 @@ func TestAuth(t *testing.T) {
LogoutURL,
ClientID,
ClientSecret,
Key,
ShortKey,
)
if err == nil || !strings.Contains(err.Error(), "decoding block key") {
t.Fatalf("expected error decoding block key got: %v", err)
if err == nil || !strings.Contains(err.Error(), "decoding secret key") {
t.Fatalf("expected error decoding secret key got: %v", err)
}
})
t.Run("NewAuthContext-BadAccessToken", func(t *testing.T) {
_, err := a.newAuthContext(&tokenAT)
if err == nil || !strings.Contains(err.Error(), "jwt claims") {
t.Fatalf("expected failure regarding jwt claims, got: %v", err)
}
})
t.Run("AuthContext-NoAccessToken", func(t *testing.T) {
_, err := a.newAuthContext(&tokenNoAT)
if err == nil || !strings.Contains(err.Error(), "access token") {
t.Fatalf("expected failure regarding access token, got: %v", err)
}
})
}

View file

@ -68,7 +68,8 @@ func (p *GroupPermissions) ReadPermissionsFile(permsFile io.Reader) (err error)
return
}
func (p *GroupPermissions) GetPermissions(groups []authn.Group, index string) (permission Permission, errors error) {
func (p *GroupPermissions) GetPermissions(user *authn.UserInfo, index string) (permission Permission, errors error) {
groups := user.Groups
if admin := p.IsAdmin(groups); admin {
return Admin, nil
}
@ -88,7 +89,7 @@ func (p *GroupPermissions) GetPermissions(groups []authn.Group, index string) (p
if perm, ok := p.Permissions[group.GroupID][index]; ok {
allPermissions[perm] = true
} else {
return None, fmt.Errorf("user %s does not have permission to index %s", group.UserID, index)
return None, fmt.Errorf("user %s does not have permission to index %s", user.UserID, index)
}
} else {
groupsDenied = append(groupsDenied, group.GroupID)

View file

@ -107,17 +107,15 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
// initializes groups that are returned from identity provider
groupName := "name"
userId := "user-id"
groupsList1 := []authn.Group{}
groupsList2 := []authn.Group{{
UserID: userId,
GroupID: "fake-group",
GroupName: groupName}}
groupsList3 := []authn.Group{
{UserID: userId, GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: groupName},
{UserID: userId, GroupID: "dca35310-ecda-4f23-86cd-876aee559900", GroupName: groupName},
{GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: groupName},
{GroupID: "dca35310-ecda-4f23-86cd-876aee559900", GroupName: groupName},
}
groupsList4 := []authn.Group{{UserID: userId, GroupID: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe", GroupName: groupName}}
groupsList4 := []authn.Group{{GroupID: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe", GroupName: groupName}}
tests := []struct {
yamlData string
@ -187,7 +185,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
t.Errorf("Error: %s", err)
}
p1, err := p.GetPermissions(test.groups, test.index)
p1, err := p.GetPermissions(&authn.UserInfo{Groups: test.groups}, test.index)
if p1 != test.userAccess {
t.Errorf("expected permission to be %s, but got %s", test.userAccess, p1)
@ -206,11 +204,11 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
func TestAuth_IsAdmin(t *testing.T) {
group1 := []authn.Group{
{UserID: "admin-user-id", GroupID: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe", GroupName: "admin-group"},
{GroupID: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe", GroupName: "admin-group"},
}
group2 := []authn.Group{
{UserID: "user-id", GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: "group-name"},
{GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: "group-name"},
}
groupPermissions := authz.GroupPermissions{
@ -247,15 +245,15 @@ func TestAuth_IsAdmin(t *testing.T) {
func TestAuth_GetAuthorizedIndexList(t *testing.T) {
group1 := []authn.Group{
{UserID: "user-id", GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: "group-name"},
{GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: "group-name"},
}
group2 := []authn.Group{
{UserID: "admin-user-id", GroupID: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe", GroupName: "admin-group"},
{GroupID: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe", GroupName: "admin-group"},
}
group3 := []authn.Group{
{UserID: "user-id", GroupID: "dca35310-ecda-4f23-86cd-876aee559900", GroupName: "group-name"},
{GroupID: "dca35310-ecda-4f23-86cd-876aee559900", GroupName: "group-name"},
}
p := authz.GroupPermissions{

View file

@ -115,8 +115,7 @@ func BuildServerFlags(cmd *cobra.Command, srv *server.Command) {
flags.StringVar(&srv.Config.Auth.GroupEndpointURL, "auth.group-endpoint-url", srv.Config.Auth.GroupEndpointURL, "Identity Provider's Group endpoint URL.")
flags.StringVar(&srv.Config.Auth.LogoutURL, "auth.logout-url", srv.Config.Auth.LogoutURL, "Identity Provider's Logout URL.")
flags.StringSliceVar(&srv.Config.Auth.Scopes, "auth.scopes", srv.Config.Auth.Scopes, "Comma separated list of scopes obtained from IdP")
flags.StringVar(&srv.Config.Auth.HashKey, "auth.hash-key", srv.Config.Auth.HashKey, "First Secret for Auth.")
flags.StringVar(&srv.Config.Auth.BlockKey, "auth.block-key", srv.Config.Auth.BlockKey, "Second Secret for Auth.")
flags.StringVar(&srv.Config.Auth.SecretKey, "auth.secret-key", srv.Config.Auth.SecretKey, "Secret key used for auth.")
flags.StringVar(&srv.Config.Auth.PermissionsFile, "auth.permissions", srv.Config.Auth.PermissionsFile, "Permissions' file with group authorization.")
flags.StringVar(&srv.Config.Auth.QueryLogPath, "auth.query-log-path", srv.Config.Auth.QueryLogPath, "Path to log user queries")

4
go.mod
View file

@ -14,6 +14,7 @@ require (
github.com/cespare/xxhash v1.1.0
github.com/davecgh/go-spew v1.1.1
github.com/desertbit/timer v0.0.0-20180107155436-c41aec40b27f // indirect
github.com/dgrijalva/jwt-go v3.2.0+incompatible
github.com/dustin/go-humanize v1.0.0 // indirect
github.com/felixge/fgprof v0.9.1
github.com/fsnotify/fsnotify v1.4.9 // indirect
@ -57,8 +58,9 @@ require (
golang.org/x/net v0.0.0-20210805182204-aaa1db679c0d // indirect
golang.org/x/oauth2 v0.0.0-20190604053449-0f29369cfe45
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c
golang.org/x/sys v0.0.0-20220111092808-5a964db01320 // indirect
google.golang.org/grpc v1.28.0
gopkg.in/yaml.v2 v2.3.0
gopkg.in/yaml.v2 v2.4.0
modernc.org/mathutil v1.0.0
modernc.org/strutil v1.0.0
sigs.k8s.io/yaml v1.2.0 // indirect

7
go.sum
View file

@ -504,8 +504,9 @@ golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20201014080544-cc95f250f6bc/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20210816074244-15123e1e1f71 h1:ikCpsnYR+Ew0vu99XlDp55lGgDJdIMx3f4a18jfse/s=
golang.org/x/sys v0.0.0-20210816074244-15123e1e1f71/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220111092808-5a964db01320 h1:0jf+tOCoZ3LyutmCOWpVni1chK4VfFLhRsDK7MhqGRY=
golang.org/x/sys v0.0.0-20220111092808-5a964db01320/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.1-0.20180807135948-17ff2d5776d2/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
@ -596,8 +597,8 @@ gopkg.in/yaml.v2 v2.2.1/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
gopkg.in/yaml.v2 v2.2.4/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
gopkg.in/yaml.v2 v2.2.8/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
gopkg.in/yaml.v2 v2.3.0 h1:clyUAQHOM3G0M3f5vQj7LuJrETvjVot3Z5el9nffUtU=
gopkg.in/yaml.v2 v2.3.0/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI=
gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY=
gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ=
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c h1:dUUwHk2QECo/6vqA44rthZ8ie2QXMNeKRTHCNY2nXvo=
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4=

View file

@ -21,6 +21,7 @@ import (
"github.com/hashicorp/go-retryablehttp"
pilosa "github.com/molecula/featurebase/v2"
"github.com/molecula/featurebase/v2/authn"
"github.com/molecula/featurebase/v2/encoding/proto"
"github.com/molecula/featurebase/v2/ingest"
"github.com/molecula/featurebase/v2/logger"
@ -42,6 +43,9 @@ type InternalClient struct {
retryableClient *retryablehttp.Client
// the local node's API, used for operations that we can short-circuit that way
api *pilosa.API
// secret Key for auth across nodes
secretKey string
}
// NewInternalClient returns a new instance of InternalClient to connect to host.
@ -63,6 +67,14 @@ func NewInternalClient(host string, remoteClient *http.Client, opts ...InternalC
type InternalClientOption func(c *InternalClient)
// WithSecretKey adds the secretKey used for inter-node communication when auth
// is enabled
func WithSecretKey(secretKey string) InternalClientOption {
return func(c *InternalClient) {
c.secretKey = secretKey
}
}
// WithClientRetryPeriod is the max amount of total time the client will
// retry failed requests using exponential backoff.
func WithClientRetryPeriod(period time.Duration) InternalClientOption {
@ -572,6 +584,12 @@ func (c *InternalClient) QueryNode(ctx context.Context, uri *pnet.URI, index str
return nil, errors.Wrap(err, "creating request")
}
uinfo := ctx.Value("userinfo")
if uinfo != nil {
token := uinfo.(*authn.UserInfo).Token
req.Header.Set("Authorization", token)
}
req.Header.Set("Content-Length", strconv.Itoa(len(buf)))
req.Header.Set("Content-Type", "application/x-protobuf")
req.Header.Set("Accept", "application/x-protobuf")
@ -1206,10 +1224,14 @@ func (c *InternalClient) SendMessage(ctx context.Context, uri *pnet.URI, msg []b
if err != nil {
return errors.Wrap(err, "making new request")
}
req.Header.Set("Content-Type", "application/x-protobuf")
req.Header.Set("User-Agent", "pilosa/"+pilosa.Version)
req.Header.Set("Accept", "application/json")
req.Header.Set("Connection", "keep-alive")
if c.secretKey != "" {
req.Header.Set("X-Feature-Key", c.secretKey)
}
// Execute request.
resp, err := c.executeRequest(req.WithContext(ctx))

View file

@ -5,6 +5,7 @@ import (
"bytes"
"context"
"crypto/tls"
"encoding/hex"
"encoding/json"
"expvar"
"fmt"
@ -416,9 +417,9 @@ func newRouter(handler *Handler) http.Handler {
router.HandleFunc("/index/{index}/field/", handler.chkAuthZ(handler.handlePostField, authz.Write)).Methods("POST").Name("PostField")
router.HandleFunc("/index/{index}/field/{field}", handler.chkAuthZ(handler.handlePostField, authz.Write)).Methods("POST").Name("PostField")
router.HandleFunc("/index/{index}/field/{field}", handler.chkAuthZ(handler.handleDeleteField, authz.Write)).Methods("DELETE").Name("DeleteField")
router.HandleFunc("/index/{index}/field/{field}/import", handler.chkAuthZ(handler.handlePostImport, authz.Read)).Methods("POST").Name("PostImport")
router.HandleFunc("/index/{index}/field/{field}/import", handler.chkAuthZ(handler.handlePostImport, authz.Write)).Methods("POST").Name("PostImport")
router.HandleFunc("/index/{index}/field/{field}/mutex-check", handler.chkAuthZ(handler.handleGetMutexCheck, authz.Read)).Methods("GET").Name("GetMutexCheck")
router.HandleFunc("/index/{index}/field/{field}/import-roaring/{shard}", handler.chkAuthZ(handler.handlePostImportRoaring, authz.Read)).Methods("POST").Name("PostImportRoaring")
router.HandleFunc("/index/{index}/field/{field}/import-roaring/{shard}", handler.chkAuthZ(handler.handlePostImportRoaring, authz.Write)).Methods("POST").Name("PostImportRoaring")
router.HandleFunc("/index/{index}/query", handler.chkAuthZ(handler.handlePostQuery, authz.Read)).Methods("POST").Name("PostQuery")
router.HandleFunc("/info", handler.chkAuthZ(handler.handleGetInfo, authz.Admin)).Methods("GET").Name("GetInfo")
router.HandleFunc("/recalculate-caches", handler.chkAuthZ(handler.handleRecalculateCaches, authz.Admin)).Methods("POST").Name("RecalculateCaches")
@ -444,14 +445,18 @@ func newRouter(handler *Handler) http.Handler {
// /internal endpoints are for internal use only; they may change at any time.
// DO NOT rely on these for external applications!
router.HandleFunc("/internal/cluster/message", handler.chkAuthN(handler.handlePostClusterMessage)).Methods("POST").Name("PostClusterMessage")
// Truly used internally by featurebease
router.HandleFunc("/internal/cluster/message", handler.chkInternal(handler.handlePostClusterMessage)).Methods("POST").Name("PostClusterMessage")
router.HandleFunc("/internal/translate/data", handler.chkInternal(handler.handleGetTranslateData)).Methods("GET").Name("GetTranslateData")
router.HandleFunc("/internal/translate/data", handler.chkInternal(handler.handlePostTranslateData)).Methods("POST").Name("PostTranslateData")
// other ones
router.HandleFunc("/internal/fragment/block/data", handler.chkAuthN(handler.handleGetFragmentBlockData)).Methods("GET").Name("GetFragmentBlockData")
router.HandleFunc("/internal/fragment/blocks", handler.chkAuthN(handler.handleGetFragmentBlocks)).Methods("GET").Name("GetFragmentBlocks")
router.HandleFunc("/internal/fragment/data", handler.chkAuthN(handler.handleGetFragmentData)).Methods("GET").Name("GetFragmentData")
router.HandleFunc("/internal/fragment/nodes", handler.chkAuthN(handler.handleGetFragmentNodes)).Methods("GET").Name("GetFragmentNodes")
router.HandleFunc("/internal/partition/nodes", handler.chkAuthN(handler.handleGetPartitionNodes)).Methods("GET").Name("GetPartitionNodes")
router.HandleFunc("/internal/translate/data", handler.chkAuthN(handler.handleGetTranslateData)).Methods("GET").Name("GetTranslateData")
router.HandleFunc("/internal/translate/data", handler.chkAuthN(handler.handlePostTranslateData)).Methods("POST").Name("PostTranslateData")
router.HandleFunc("/internal/translate/keys", handler.chkAuthN(handler.handlePostTranslateKeys)).Methods("POST").Name("PostTranslateKeys")
router.HandleFunc("/internal/translate/ids", handler.chkAuthN(handler.handlePostTranslateIDs)).Methods("POST").Name("PostTranslateIDs")
router.HandleFunc("/internal/index/{index}/field/{field}/mutex-check", handler.chkAuthN(handler.handleInternalGetMutexCheck)).Methods("GET").Name("InternalGetMutexCheck")
@ -540,11 +545,25 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
h.Handler.ServeHTTP(w, r)
}
func (h *Handler) chkInternal(handler http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if h.auth != nil {
secret, ok := r.Header["X-Feature-Key"]
decodedString, err := hex.DecodeString(secret[0])
if err != nil || !ok || !bytes.Equal(decodedString, h.auth.SecretKey()) {
http.Error(w, errors.Wrap(err, "internal secret key validation failed").Error(), http.StatusUnauthorized)
return
}
}
handler.ServeHTTP(w, r)
}
}
func (h *Handler) chkAuthN(handler http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if h.auth != nil {
if _, err := h.auth.Authenticate(w, r); err != nil {
http.Error(w, errors.Wrap(err, "authenticating").Error(), http.StatusBadRequest)
if _, err := h.auth.Authenticate(getToken(r)); err != nil {
http.Error(w, errors.Wrap(err, "authenticating").Error(), http.StatusUnauthorized)
return
}
}
@ -556,9 +575,9 @@ func (h *Handler) chkAuthZ(handler http.HandlerFunc, perm authz.Permission) http
return func(w http.ResponseWriter, r *http.Request) {
lperm := perm
if h.auth != nil {
groups, err := h.auth.Authenticate(w, r)
uinfo, err := h.auth.Authenticate(getToken(r))
if err != nil {
http.Error(w, errors.Wrap(err, "authenticating").Error(), http.StatusBadRequest)
http.Error(w, errors.Wrap(err, "authenticating").Error(), http.StatusForbidden)
return
}
@ -568,7 +587,16 @@ func (h *Handler) chkAuthZ(handler http.HandlerFunc, perm authz.Permission) http
return
}
uinfo := h.auth.GetUserInfo(w, r)
ctx := context.WithValue(r.Context(), contextKeyGroupMembership, uinfo.Groups)
if h.permissions.IsAdmin(uinfo.Groups) {
ctx = context.WithValue(ctx, contextKeyPermission, authz.Admin)
handler.ServeHTTP(w, r.WithContext(ctx))
return
} else if lperm == authz.Admin {
http.Error(w, "Insufficient permissions: user does not have admin permission", http.StatusForbidden)
return
}
var queryString string
queryRequest := r.Context().Value(contextKeyQueryRequest)
@ -591,14 +619,18 @@ func (h *Handler) chkAuthZ(handler http.HandlerFunc, perm authz.Permission) http
h.querylogger.Infof("User ID: %s, User Name: %s, Endpoint: %s, Index: %s, Query: %s, Err: %v", uinfo.UserID, uinfo.UserName, r.URL.Path, "indexName", queryString, err)
}
ctx := context.WithValue(r.Context(), contextKeyGroupMembership, groups)
indexName, ok := mux.Vars(r)["index"]
if ok {
p, err := h.permissions.GetPermissions(groups, indexName)
ctx = context.WithValue(r.Context(), contextKeyPermission, p)
if err != nil || !p.Satisfies(lperm) {
p, err := h.permissions.GetPermissions(uinfo, indexName)
ctx = context.WithValue(ctx, contextKeyPermission, p)
if err != nil {
w.Header().Add("Content-Type", "text/plain")
w.WriteHeader(http.StatusForbidden)
http.Error(w, errors.Wrap(err, "Insufficient Permissions").Error(), http.StatusForbidden)
return
}
if !p.Satisfies(lperm) {
w.Header().Add("Content-Type", "text/plain")
http.Error(w, fmt.Sprintf("Insufficient permissions: user has %s permissions, but request requires %s permission", p, lperm), http.StatusForbidden)
return
}
}
@ -3551,10 +3583,11 @@ func (h *Handler) handleCheckAuthentication(w http.ResponseWriter, r *http.Reque
http.Error(w, "", http.StatusNoContent)
return
}
groups, err := h.auth.Authenticate(w, r)
if groups == nil || err != nil {
uinfo, err := h.auth.Authenticate(getToken(r))
if uinfo == nil || err != nil {
w.Header().Add("Content-Type", "text/plain")
w.WriteHeader(http.StatusForbidden)
http.Error(w, err.Error(), http.StatusUnauthorized)
return
}
w.Header().Add("Content-Type", "text/plain")
@ -3572,7 +3605,14 @@ func (h *Handler) handleUserInfo(w http.ResponseWriter, r *http.Request) {
http.Error(w, "", http.StatusNoContent)
return
}
if err := json.NewEncoder(w).Encode(h.auth.GetUserInfo(w, r)); err != nil {
uinfo, err := h.auth.Authenticate(getToken(r))
if err != nil {
h.logger.Errorf("error authenticating: %v", err)
http.Error(w, err.Error(), http.StatusForbidden)
return
}
if err := json.NewEncoder(w).Encode(uinfo); err != nil {
h.logger.Errorf("writing user info: %s", err)
}
}
@ -3584,3 +3624,20 @@ 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 {
if token, ok := r.Header["Authorization"]; ok && len(token) > 0 {
parts := strings.Split(token[0], "Bearer ")
if len(parts) != 2 {
return ""
}
return parts[1]
}
cookie, err := r.Cookie("molecula-chip")
if err != nil {
return ""
}
return cookie.Value
}

View file

@ -15,15 +15,14 @@ import (
"testing"
"time"
"github.com/gorilla/securecookie"
"github.com/golang-jwt/jwt"
pilosa "github.com/molecula/featurebase/v2"
"github.com/molecula/featurebase/v2/authn"
"github.com/stretchr/testify/assert"
"golang.org/x/oauth2"
"github.com/molecula/featurebase/v2/authz"
"github.com/molecula/featurebase/v2/logger"
"github.com/molecula/featurebase/v2/pql"
"golang.org/x/oauth2"
)
// Test custom UnmarshalJSON for postIndexRequest object
@ -198,12 +197,10 @@ func TestAuthentication(t *testing.T) {
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"}
HashKey = "DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF"
BlockKey = "DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF"
SecretKey = "DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF"
)
hashKey, _ := hex.DecodeString("DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF")
blockKey, _ := hex.DecodeString("DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF")
secretKey, _ := hex.DecodeString("DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF")
a, err := authn.NewAuth(
logger.NewStandardLogger(os.Stdout),
@ -215,97 +212,56 @@ func TestAuthentication(t *testing.T) {
LogoutURL,
ClientId,
ClientSecret,
HashKey,
BlockKey,
SecretKey,
)
if err != nil {
t.Errorf("building auth object%s", err)
}
h := Handler{
auth: a,
logger: logger.NewStandardLogger(os.Stdout),
querylogger: logger.NewStandardLogger(os.Stdout),
auth: a,
}
hOff := Handler{}
// make a valid token
tkn := jwt.New(jwt.SigningMethodHS256)
claims := tkn.Claims.(jwt.MapClaims)
groupString, _ := authn.ToGob64([]authn.Group{{GroupID: "thing", GroupName: "whatever"}})
claims["molecula-idp-groups"] = groupString
claims["oid"] = "42"
claims["name"] = "todd"
validToken, err := tkn.SignedString([]byte(secretKey))
if err != nil {
panic(err)
}
validToken = "Bearer " + validToken
token := oauth2.Token{
TokenType: "Bearer",
AccessToken: "asdf",
RefreshToken: "abcdef",
Expiry: time.Now().Add(time.Hour),
}
expiredToken := oauth2.Token{
TokenType: "Bearer",
RefreshToken: "abcdef",
Expiry: time.Now(),
// make an expired token
expiredTkn := jwt.New(jwt.SigningMethodHS256)
expiredClaims := expiredTkn.Claims.(jwt.MapClaims)
expiredClaims["molecula-idp-groups"] = groupString
expiredClaims["oid"] = "42"
expiredClaims["name"] = "todd"
expiredClaims["exp"] = "1"
expiredToken, err := expiredTkn.SignedString([]byte(secretKey))
if err != nil {
panic(err)
}
grp := authn.Group{
UserID: "snowstorm",
GroupID: "abcd123-A",
GroupName: "Romantic Painters",
}
validCV := authn.AuthContext{
UserID: "snowstorm",
UserName: "J.M.W. Turner",
GroupMembership: []authn.Group{grp},
Token: &token,
}
emptyCV := authn.AuthContext{
UserID: "narcissus",
UserName: "Caravaggio",
GroupMembership: []authn.Group{},
Token: &token,
}
expiredCV := authn.AuthContext{
UserID: "narcissus",
UserName: "Caravaggio",
GroupMembership: []authn.Group{grp},
Token: &expiredToken,
}
secure := securecookie.New(hashKey, blockKey)
validEncodedCV, _ := secure.Encode("molecula-chip", validCV)
noGroupEncodedCV, _ := secure.Encode("molecula-chip", emptyCV)
expiredEncodedCV, _ := secure.Encode("molecula-chip", expiredCV)
expiredToken = "Bearer " + expiredToken
validCookie := &gohttp.Cookie{
Name: "molecula-chip",
Value: validEncodedCV,
Path: "/",
Secure: true,
HttpOnly: true,
Expires: token.Expiry,
}
noGroupCookie := &gohttp.Cookie{
Name: "molecula-chip",
Value: noGroupEncodedCV,
Path: "/",
Secure: true,
HttpOnly: true,
Expires: token.Expiry,
}
expiredCookie := &gohttp.Cookie{
Name: "molecula-chip",
Value: expiredEncodedCV,
Path: "/",
Secure: true,
HttpOnly: true,
Expires: time.Now().Add(time.Minute * -1),
}
emptyCookie := &gohttp.Cookie{
Name: "molecula-chip",
Value: "",
Path: "/",
Secure: true,
HttpOnly: true,
Expires: token.Expiry,
}
unEncodedCookie := &gohttp.Cookie{
Name: "molecula-chip",
Value: "The quick brown fox",
Value: token.AccessToken,
Path: "/",
Secure: true,
HttpOnly: true,
@ -321,7 +277,9 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
name string
path string
kind string
method string
yamlData string
token string
cookie *gohttp.Cookie
handler endpoint
fn evaluate
@ -331,7 +289,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
path: "/login",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleLogin(w, r) },
handler: h.handleLogin,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if strings.Index(string(data), AuthorizeURL) != 9 {
t.Errorf("incorrect redirect url: expected: %s, got: %s", AuthorizeURL, string(data))
@ -343,7 +301,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
path: "/logout",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleLogout(w, r) },
handler: h.handleLogout,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if w.Result().Cookies()[0].Value != "" {
t.Errorf("expected cookie to be cleared, got: %+v", w.Result().Cookies()[0].Value)
@ -351,82 +309,69 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
},
},
{
name: "Authenticate-Groups",
name: "Authenticate-ValidToken",
path: "/auth",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleCheckAuthentication(w, r) },
kind: "bearer",
token: validToken,
handler: h.handleCheckAuthentication,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if w.Result().StatusCode != 200 {
t.Errorf("expected http code 200, got: %+v", w.Result().StatusCode)
body, _ := readResponse(w)
t.Errorf("expected http code 200, got: %+v with body: %+v", w.Result().StatusCode, body)
}
},
},
{
name: "Authenticate-NoGroups",
name: "Authenticate-NoToken",
path: "/auth",
kind: "type1",
cookie: noGroupCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleCheckAuthentication(w, r) },
handler: h.handleCheckAuthentication,
fn: func(w *httptest.ResponseRecorder, data []byte) {
// status forbidden
if w.Result().StatusCode != 403 {
// not token at all == status forbidden
if w.Result().StatusCode != 401 {
t.Errorf("expected http code 401, got: %+v", w.Result().StatusCode)
}
},
},
{
name: "Authenticate-InvalidToken",
path: "/auth",
kind: "type1",
token: "this isn't a real token",
handler: h.handleCheckAuthentication,
fn: func(w *httptest.ResponseRecorder, data []byte) {
// no valid token in header == Unauthorized
if w.Result().StatusCode != gohttp.StatusUnauthorized {
t.Errorf("expected http code 401, got: %+v", w.Result().StatusCode)
}
},
},
{
name: "Authenticate-ExpiredToken",
path: "/auth",
kind: "type1",
token: expiredToken,
handler: h.handleCheckAuthentication,
fn: func(w *httptest.ResponseRecorder, data []byte) {
// expired token == unauthorized
if w.Result().StatusCode != 401 {
t.Errorf("expected http code 403, got: %+v", w.Result().StatusCode)
}
},
},
{
name: "Authenticate-MalformedCookie",
path: "/auth",
kind: "type1",
cookie: unEncodedCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleCheckAuthentication(w, r) },
fn: func(w *httptest.ResponseRecorder, data []byte) {
// redirect to signin
if w.Result().StatusCode != 307 {
t.Errorf("expected http code 307, got: %+v", w.Result().StatusCode)
}
},
},
{
name: "Authenticate-Expired",
path: "/auth",
kind: "type1",
cookie: expiredCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleCheckAuthentication(w, r) },
fn: func(w *httptest.ResponseRecorder, data []byte) {
// redirect to signin
if w.Result().StatusCode != 307 {
t.Errorf("expected http code 307, got: %+v", w.Result().StatusCode)
}
},
},
{
name: "Authenticate-NoCookie",
path: "/auth",
kind: "type1",
cookie: emptyCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleCheckAuthentication(w, r) },
fn: func(w *httptest.ResponseRecorder, data []byte) {
// redirect to signin
if w.Result().StatusCode != 307 {
t.Errorf("expected http code 307, got: %+v", w.Result().StatusCode)
}
},
},
{
name: "UserInfo",
path: "/userinfo",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleUserInfo(w, r) },
kind: "bearer",
token: validToken,
handler: h.handleUserInfo,
fn: func(w *httptest.ResponseRecorder, data []byte) {
uinfo := authn.UserInfo{}
err = json.Unmarshal(data, &uinfo)
if err != nil {
t.Errorf("unmarshalling userinfo")
}
if uinfo.UserID != "snowstorm" && uinfo.UserName != "J.M.W. Turner" {
if uinfo.UserID != "42" && uinfo.UserName != "todd" {
t.Errorf("expected http code 400, got: %+v", uinfo)
}
},
@ -434,27 +379,21 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
{
name: "UserInfo-NoCookie",
path: "/userinfo",
kind: "type1",
cookie: emptyCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleUserInfo(w, r) },
kind: "bearer",
token: "",
handler: h.handleUserInfo,
fn: func(w *httptest.ResponseRecorder, data []byte) {
uinfo := authn.UserInfo{}
err = json.Unmarshal(data, &uinfo)
if err != nil {
t.Errorf("unmarshalling userinfo")
}
if uinfo.UserID != "" && uinfo.UserName != "" {
t.Errorf("expected http code 400, got: %+v", uinfo)
if got := w.Result().StatusCode; got != gohttp.StatusForbidden {
t.Errorf("expected 403, got %v", got)
}
},
},
{
name: "Redirect-NoAuthCode",
path: "/redirect",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleRedirect(w, r) },
handler: h.handleRedirect,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if strings.Index(string(data), AuthorizeURL) != 9 {
if w.Result().StatusCode != 400 {
@ -468,7 +407,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
path: "/redirect",
kind: "type2",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { h.handleRedirect(w, r) },
handler: h.handleRedirect,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if strings.Index(string(data), AuthorizeURL) != 9 {
if w.Result().StatusCode != 400 {
@ -482,7 +421,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
path: "/login",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { hOff.handleLogin(w, r) },
handler: hOff.handleLogin,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if strings.Index(string(data), AuthorizeURL) != 9 {
if w.Result().StatusCode != 204 {
@ -496,7 +435,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
path: "/logout",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { hOff.handleLogout(w, r) },
handler: hOff.handleLogout,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if strings.Index(string(data), AuthorizeURL) != 9 {
if w.Result().StatusCode != 204 {
@ -510,7 +449,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
path: "/userinfo",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { hOff.handleUserInfo(w, r) },
handler: hOff.handleUserInfo,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if strings.Index(string(data), AuthorizeURL) != 9 {
if w.Result().StatusCode != 204 {
@ -524,7 +463,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
path: "/auth",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { hOff.handleCheckAuthentication(w, r) },
handler: hOff.handleCheckAuthentication,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if strings.Index(string(data), AuthorizeURL) != 9 {
if w.Result().StatusCode != 204 {
@ -538,7 +477,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
path: "/redirect",
kind: "type1",
cookie: validCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) { hOff.handleRedirect(w, r) },
handler: hOff.handleRedirect,
fn: func(w *httptest.ResponseRecorder, data []byte) {
if strings.Index(string(data), AuthorizeURL) != 9 {
if w.Result().StatusCode != 204 {
@ -563,59 +502,57 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
},
},
{
name: "MW-ExpiredAuth",
path: "/index/{index}/query",
kind: "middleware",
cookie: expiredCookie,
// this tests that there are no permissions read in even though
// auth is turned on, so we get a 500
name: "MW-CreateIndexGood",
path: "/index/abcd",
kind: "bearer",
method: gohttp.MethodPost,
token: validToken,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) {
f := h.chkAuthN(h.handlePostQuery)
h := h
var p authz.GroupPermissions
if err := p.ReadPermissionsFile(strings.NewReader(permissions1)); err != nil {
t.Errorf("Error: %s", err)
}
h.permissions = &p
f := h.chkAuthZ(h.handlePostIndex, authz.Admin)
f(w, r)
},
fn: func(w *httptest.ResponseRecorder, data []byte) {
if w.Result().StatusCode != 307 {
t.Errorf("expected http code 307, got: %+v", w.Result().StatusCode)
if got, want := w.Result().StatusCode, gohttp.StatusForbidden; got != want {
t.Errorf("expected %v, got %v", want, got)
}
},
},
{
name: "MW-ExpiredAuth2",
path: "/index/{index}/query",
kind: "middleware",
cookie: expiredCookie,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) {
f := h.chkAuthZ(h.handlePostQuery, authz.Admin)
f(w, r)
},
fn: func(w *httptest.ResponseRecorder, data []byte) {
if w.Result().StatusCode != 307 {
t.Errorf("expected http code 307, got: %+v", w.Result().StatusCode)
}
},
},
{
name: "MW-NoPermissions",
path: "/index/{index}/query",
kind: "middleware",
cookie: validCookie,
// this tests that there are no permissions read in even though
// auth is turned on, so we get a 500
name: "MW-NoPermissions",
path: "/index/{index}/query",
kind: "bearer",
token: validToken,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) {
h := h
f := h.chkAuthZ(h.handlePostQuery, authz.Write)
assert.Panics(t, func() { f(w, r) }, "expected panic")
f(w, r)
},
fn: func(w *httptest.ResponseRecorder, data []byte) {
if got, want := w.Result().StatusCode, gohttp.StatusInternalServerError; got != want {
t.Errorf("expected %v, got %v", want, got)
}
},
fn: func(w *httptest.ResponseRecorder, data []byte) {},
},
{
name: "MW-NoIndexNoAdmin",
path: "/index/{index}/query",
kind: "middleware",
cookie: validCookie,
name: "MW-NoQuery",
path: "/index/{index}/query",
kind: "bearer",
token: validToken,
handler: func(w gohttp.ResponseWriter, r *gohttp.Request) {
h := h
permFile := strings.NewReader(permissions1)
var p authz.GroupPermissions
if err := p.ReadPermissionsFile(permFile); err != nil {
if err := p.ReadPermissionsFile(strings.NewReader(permissions1)); err != nil {
t.Errorf("Error: %s", err)
}
h.permissions = &p
@ -623,17 +560,16 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
f(w, r)
},
fn: func(w *httptest.ResponseRecorder, data []byte) {
if w.Result().StatusCode != 400 {
t.Errorf("expected http code 400, got: %+v", w.Result().StatusCode)
if got, want := w.Result().StatusCode, gohttp.StatusBadRequest; got != want {
t.Errorf("expected %v, got: %+v", want, got)
}
},
},
}
for _, test := range tests {
switch test.kind {
case "type1":
case "type1", "middleware":
t.Run(test.name, func(t *testing.T) {
r := httptest.NewRequest(gohttp.MethodGet, test.path, nil)
w := httptest.NewRecorder()
@ -664,20 +600,21 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
test.fn(w, data)
})
case "middleware":
case "bearer":
t.Run(test.name, func(t *testing.T) {
r := httptest.NewRequest(gohttp.MethodGet, test.path, nil)
w := httptest.NewRecorder()
if test.cookie != nil {
r.AddCookie(test.cookie)
if test.method == "" {
test.method = gohttp.MethodGet
}
r := httptest.NewRequest(test.method, test.path, nil)
w := httptest.NewRecorder()
if test.token != "" {
r.Header.Add("Authorization", test.token)
}
test.handler(w, r)
data, err := readResponse(w)
if err != nil {
t.Errorf("expected no errors reading response, got: %+v", err)
}
test.fn(w, data)
})
}

View file

@ -373,7 +373,7 @@ log-path = "/var/log/molecula/featurebase.log"
# ==============================================================================
# Enable/Disable AuthN/AuthZ for featurebase
# Can choose identity provider, defaults for Azure Active Directory
# Use provided keygen binary to generate hash and block keys with sufficient length and entropy
# Use provided keygen binary to generate a secret key with sufficient length and entropy
# [auth]
# enable = false
# client-id = ""
@ -383,7 +383,6 @@ log-path = "/var/log/molecula/featurebase.log"
# group-endpoint-url = ""
# logout-url = ""
# scopes = ["", ""]
# hash-key = ""
# block-key = ""
# secret-key = ""
# permissions = ""
# query-log-path = ""

View file

@ -51,7 +51,7 @@ test('test useAuth - expect authenticated', async () => {
test('test useAuth - expect not authed', async () => {
const mockResponse: AxiosResponse<any> = {
status: 200,
status: 401,
data: '',
statusText: '',
headers: {},

View file

@ -54,7 +54,7 @@ function useProvideAuth() {
// Turn on Authentication
setIsAuthOn(true);
if (res.data === 'OK') {
if (res.status === 200) {
// User is authenticated
setIsAuthenticated(true);

View file

@ -242,8 +242,7 @@ type Auth struct {
GroupEndpointURL string `toml:"group-endpoint-url"`
LogoutURL string `toml:"logout-url"`
Scopes []string `toml:"scopes"`
HashKey string `toml:"hash-key"`
BlockKey string `toml:"block-key"`
SecretKey string `toml:"secret-key"`
PermissionsFile string `toml:"permissions"`
QueryLogPath string `toml:"query-log-path"`
}
@ -614,24 +613,29 @@ func (c *Config) ValidateAuth() (errors []error) {
if !c.Auth.Enable {
return
}
authConfig := map[string]string{
"ClientId": c.Auth.ClientId,
"ClientSecret": c.Auth.ClientSecret,
"AuthorizeURL": c.Auth.AuthorizeURL,
"TokenURL": c.Auth.TokenURL,
"GroupEndpointURL": c.Auth.GroupEndpointURL,
"LogoutURL": c.Auth.LogoutURL,
"HashKey": c.Auth.HashKey,
"BlockKey": c.Auth.BlockKey,
authConfig := []struct {
name string
val string
}{
{name: "ClientId", val: c.Auth.ClientId},
{name: "ClientSecret", val: c.Auth.ClientSecret},
{name: "AuthorizeURL", val: c.Auth.AuthorizeURL},
{name: "TokenURL", val: c.Auth.TokenURL},
{name: "GroupEndpointURL", val: c.Auth.GroupEndpointURL},
{name: "LogoutURL", val: c.Auth.LogoutURL},
{name: "SecretKey", val: c.Auth.SecretKey},
{name: "QueryLogPath", val: c.Auth.QueryLogPath},
}
for name, value := range authConfig {
for _, configOpt := range authConfig {
name := configOpt.name
value := configOpt.val
if value == "" {
errors = append(errors, fmt.Errorf("empty string for auth config %s", name))
continue
}
if name == "HashKey" || name == "BlockKey" {
if name == "SecretKey" {
if len(value) != 64 {
errors = append(errors, fmt.Errorf("invalid key length for %s. exp %d, got %d", name, 64, len(value)))
}

View file

@ -319,15 +319,14 @@ func TestConfig_validateAuth(t *testing.T) {
GroupEndpointURL: emptyString,
LogoutURL: emptyString,
Scopes: validStringSlice,
HashKey: emptyString,
BlockKey: emptyString,
SecretKey: emptyString,
},
},
{
// Auth enabled, keys are invalid length
[]string{
errorMesgKey,
errorMesgKey,
errorMesgEmpty,
},
Auth{
Enable: enable,
@ -338,8 +337,7 @@ func TestConfig_validateAuth(t *testing.T) {
GroupEndpointURL: validTestURL,
LogoutURL: validTestURL,
Scopes: validStringSlice,
HashKey: validString,
BlockKey: validString,
SecretKey: validString,
},
},
{
@ -358,8 +356,8 @@ func TestConfig_validateAuth(t *testing.T) {
GroupEndpointURL: invalidURL,
LogoutURL: invalidURL,
Scopes: validStringSlice,
HashKey: validKey,
BlockKey: validKey,
SecretKey: validKey,
QueryLogPath: "thisnisfasdfPAth",
},
},
{
@ -376,8 +374,8 @@ func TestConfig_validateAuth(t *testing.T) {
GroupEndpointURL: validTestURL,
LogoutURL: validTestURL,
Scopes: emptySlice,
HashKey: validKey,
BlockKey: validKey,
SecretKey: validKey,
QueryLogPath: "thisaisdf aPath",
},
},
{
@ -392,8 +390,8 @@ func TestConfig_validateAuth(t *testing.T) {
GroupEndpointURL: validTestURL,
LogoutURL: validTestURL,
Scopes: validStringSlice,
HashKey: validKey,
BlockKey: validKey,
SecretKey: validKey,
QueryLogPath: "thisIsAPAth",
},
},
{
@ -408,8 +406,7 @@ func TestConfig_validateAuth(t *testing.T) {
GroupEndpointURL: invalidURL,
LogoutURL: validTestURL,
Scopes: validStringSlice,
HashKey: validKey,
BlockKey: emptyString,
SecretKey: emptyString,
},
},
}

View file

@ -13,10 +13,14 @@ import (
"time"
"github.com/improbable-eng/grpc-web/go/grpcweb"
"github.com/molecula/featurebase/v2"
pilosa "github.com/molecula/featurebase/v2"
"github.com/molecula/featurebase/v2/authn"
"github.com/molecula/featurebase/v2/authz"
"github.com/molecula/featurebase/v2/logger"
"github.com/molecula/featurebase/v2/pql"
pb "github.com/molecula/featurebase/v2/proto"
vdsm_pb "github.com/molecula/featurebase/v2/proto/vdsm"
"github.com/molecula/featurebase/v2/sql"
"github.com/molecula/featurebase/v2/stats"
"github.com/pkg/errors"
"google.golang.org/grpc"
@ -30,6 +34,7 @@ import (
// GRPCHandler contains methods which handle the various gRPC requests.
type GRPCHandler struct {
api *pilosa.API
perms *authz.GroupPermissions
logger logger.Logger
stats stats.StatsClient
inspectDeprecated sync.Once
@ -49,6 +54,11 @@ func (h *GRPCHandler) WithStats(stats stats.StatsClient) *GRPCHandler {
return h
}
func (h *GRPCHandler) WithPerms(perms *authz.GroupPermissions) *GRPCHandler {
h.perms = perms
return h
}
// errorToStatusError appends an appropriate grpc status code
// to the error (returning it as a status.Error).
func errToStatusError(err error) error {
@ -126,15 +136,49 @@ func (h *GRPCHandler) execSQL(ctx context.Context, queryStr string) (pb.ToRowser
return execSQL(ctx, h.api, h.logger, queryStr)
}
func isAllowed(requested []string, allowed []string) bool {
if len(allowed) == 0 {
return false
}
for _, r := range requested {
in := false
for _, a := range allowed {
if a == r {
in = true
}
}
if !in {
return false
}
}
return true
}
// QuerySQL handles the SQL request and sends RowResponses to the stream.
func (h *GRPCHandler) QuerySQL(req *pb.QuerySQLRequest, stream pb.Pilosa_QuerySQLServer) error {
ctx := stream.Context()
uinfo := ctx.Value("userinfo")
if uinfo != nil {
// authz
m := sql.NewMapper()
parsed, err := m.MapSQL(req.Sql)
if err != nil {
return errors.Wrap(err, "parsing SQL")
}
if !h.perms.IsAdmin(uinfo.(*authn.UserInfo).Groups) {
if !isAllowed(parsed.Tables, h.perms.GetAuthorizedIndexList(uinfo.(*authn.UserInfo).Groups, authz.Read)) {
return status.Error(codes.PermissionDenied, "insufficient permissions to access requested tables")
}
}
}
start := time.Now()
results, err := h.execSQL(stream.Context(), req.Sql)
duration := time.Since(start)
if err != nil {
return err
}
err = stream.SendHeader(metadata.New(map[string]string{
"duration": strconv.Itoa(int(duration)),
}))
@ -198,6 +242,23 @@ func (h *GRPCHandler) QueryPQL(req *pb.QueryPQLRequest, stream pb.Pilosa_QueryPQ
Query: req.Pql,
}
ctx := stream.Context()
uinfo := ctx.Value("userinfo")
if uinfo != nil {
lperm := authz.Read
q, err := pql.ParseString(req.Pql)
if err != nil {
return status.Error(codes.InvalidArgument, err.Error())
}
if q.WriteCallN() > 0 {
lperm = authz.Write
}
if !h.perms.IsAdmin(uinfo.(*authn.UserInfo).Groups) {
if !isAllowed([]string{req.Index}, h.perms.GetAuthorizedIndexList(uinfo.(*authn.UserInfo).Groups, lperm)) {
return status.Error(codes.PermissionDenied, "insufficient permissions to access requested indexes")
}
}
}
t := time.Now()
resp, err := h.api.Query(stream.Context(), &query)
durQuery := time.Since(t)
@ -246,6 +307,22 @@ func (h *GRPCHandler) QueryPQLUnary(ctx context.Context, req *pb.QueryPQLRequest
Index: req.Index,
Query: req.Pql,
}
uinfo := ctx.Value("userinfo")
if uinfo != nil {
lperm := authz.Read
q, err := pql.ParseString(req.Pql)
if err != nil {
return nil, status.Error(codes.InvalidArgument, err.Error())
}
if q.WriteCallN() > 0 {
lperm = authz.Write
}
if !h.perms.IsAdmin(uinfo.(*authn.UserInfo).Groups) {
if !isAllowed([]string{req.Index}, h.perms.GetAuthorizedIndexList(uinfo.(*authn.UserInfo).Groups, lperm)) {
return nil, status.Error(codes.PermissionDenied, fmt.Sprintf("insufficient permissions for %v", req.Index))
}
}
}
t := time.Now()
resp, err := h.api.Query(ctx, &query)
@ -291,6 +368,12 @@ func (h *GRPCHandler) QueryPQLUnary(ctx context.Context, req *pb.QueryPQLRequest
// CreateIndex creates a new Index
func (h *GRPCHandler) CreateIndex(ctx context.Context, req *pb.CreateIndexRequest) (*pb.CreateIndexResponse, error) {
uinfo := ctx.Value("userinfo")
if uinfo != nil {
if !h.perms.IsAdmin(uinfo.(*authn.UserInfo).Groups) {
return nil, status.Error(codes.PermissionDenied, "must be admin to create index")
}
}
// Always enable TrackExistence for gRPC-created indexes
opts := pilosa.IndexOptions{Keys: req.Keys, TrackExistence: true}
_, err := h.api.CreateIndex(ctx, req.Name, opts)
@ -302,6 +385,20 @@ func (h *GRPCHandler) CreateIndex(ctx context.Context, req *pb.CreateIndexReques
// GetIndex returns a single Index given a name
func (h *GRPCHandler) GetIndex(ctx context.Context, req *pb.GetIndexRequest) (*pb.GetIndexResponse, error) {
uinfo := ctx.Value("userinfo")
if uinfo != nil {
pp, ok := uinfo.(*authn.UserInfo)
if !ok {
return nil, status.Error(codes.InvalidArgument, "malformed auth header")
}
p, err := h.perms.GetPermissions(pp, req.Name)
if err != nil {
return nil, err
}
if !p.Satisfies(authz.Read) {
return nil, status.Error(codes.PermissionDenied, fmt.Sprintf("permission denied for index %v", req.Name))
}
}
schema, err := h.api.Schema(ctx, false)
if err != nil {
return nil, errToStatusError(err)
@ -317,20 +414,44 @@ func (h *GRPCHandler) GetIndex(ctx context.Context, req *pb.GetIndexRequest) (*p
// GetIndexes returns a list of all Indexes
func (h *GRPCHandler) GetIndexes(ctx context.Context, req *pb.GetIndexesRequest) (*pb.GetIndexesResponse, error) {
uinfo := ctx.Value("userinfo")
var pp *authn.UserInfo
if uinfo != nil {
var ok bool
pp, ok = uinfo.(*authn.UserInfo)
if !ok {
return nil, status.Error(codes.InvalidArgument, "malformed auth header")
}
}
schema, err := h.api.Schema(ctx, false)
if err != nil {
return nil, errToStatusError(err)
}
indexes := make([]*pb.Index, len(schema))
for i, index := range schema {
indexes[i] = &pb.Index{Name: index.Name}
i := 0
for _, index := range schema {
if pp != nil {
if p, err := h.perms.GetPermissions(pp, index.Name); err == nil && p.Satisfies(authz.Read) {
indexes[i] = &pb.Index{Name: index.Name}
i += 1
}
} else {
indexes[i] = &pb.Index{Name: index.Name}
i += 1
}
}
return &pb.GetIndexesResponse{Indexes: indexes}, nil
}
// DeleteIndex deletes an Index
func (h *GRPCHandler) DeleteIndex(ctx context.Context, req *pb.DeleteIndexRequest) (*pb.DeleteIndexResponse, error) {
uinfo := ctx.Value("userinfo")
if uinfo != nil {
if !h.perms.IsAdmin(uinfo.(*authn.UserInfo).Groups) {
return nil, status.Error(codes.PermissionDenied, "must be admin to delete index")
}
}
err := h.api.DeleteIndex(ctx, req.Name)
if err != nil {
return nil, errToStatusError(err)
@ -1301,6 +1422,8 @@ type grpcServer struct {
grpcServer *grpc.Server
ln net.Listener
tlsConfig *tls.Config
auth *authn.Auth
perms *authz.GroupPermissions
logger logger.Logger
stats stats.StatsClient
@ -1343,6 +1466,20 @@ func OptGRPCServerStats(stats stats.StatsClient) grpcServerOption {
}
}
func OptGRPCServerAuth(authn *authn.Auth) grpcServerOption {
return func(s *grpcServer) error {
s.auth = authn
return nil
}
}
func OptGRPCServerPerm(gp *authz.GroupPermissions) grpcServerOption {
return func(s *grpcServer) error {
s.perms = gp
return nil
}
}
func (s *grpcServer) Serve() error {
s.logger.Infof("enabled grpc listening on %s", s.ln.Addr())
@ -1398,10 +1535,37 @@ func NewGRPCServer(opts ...grpcServerOption) (*grpcServer, error) {
creds := credentials.NewTLS(server.tlsConfig)
gopts = append(gopts, grpc.Creds(creds))
}
//if auth enabled
if server.auth != nil {
gopts = append(gopts, grpc.UnaryInterceptor(
func(ctx context.Context, req interface{}, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (interface{}, error) {
ctx, err := Valid(ctx, server.auth)
if err != nil {
return nil, err
}
return handler(ctx, req)
},
))
gopts = append(gopts, grpc.StreamInterceptor(
func(srv interface{}, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
ctx, err := Valid(ss.Context(), server.auth)
if err != nil {
return err
}
return handler(srv, newWrappedStream(ss, ctx))
},
))
}
// create grpc server
server.grpcServer = grpc.NewServer(gopts...)
grpcHandler := NewGRPCHandler(server.api).WithLogger(server.logger).WithStats(server.stats)
// add server permissions if we've got 'em
if server.perms != nil {
grpcHandler.perms = server.perms
}
pb.RegisterPilosaServer(server.grpcServer, grpcHandler)
vdsm_pb.RegisterMoleculaServer(server.grpcServer, NewVDSMGRPCHandler(grpcHandler, server.api).WithLogger(server.logger).WithStats(server.stats))
@ -1410,3 +1574,58 @@ func NewGRPCServer(opts ...grpcServerOption) (*grpcServer, error) {
return server, nil
}
// wrappedStream wraps around the embedded grpc.ServerStream, and intercepts the RecvMsg and
// SendMsg method call.
type wrappedStream struct {
grpc.ServerStream
uiContext context.Context
}
func (w *wrappedStream) Context() context.Context {
return w.uiContext
}
func (w *wrappedStream) RecvMsg(m interface{}) error {
return w.ServerStream.RecvMsg(m)
}
func (w *wrappedStream) SendMsg(m interface{}) error {
return w.ServerStream.SendMsg(m)
}
func newWrappedStream(s grpc.ServerStream, ctx context.Context) grpc.ServerStream {
return &wrappedStream{s, ctx}
}
func Valid(ctx context.Context, auth *authn.Auth) (context.Context, error) {
md, ok := metadata.FromIncomingContext(ctx)
if !ok {
return ctx, status.Errorf(codes.InvalidArgument, "missing metadata")
}
authorization, ok := md["authorization"]
if !ok {
c, ok := md["cookie"]
if !ok {
return ctx, status.Errorf(codes.InvalidArgument, "missing authorization token")
}
cookies := strings.Split(c[0], "; ")
for _, cookie := range cookies {
if strings.HasPrefix(cookie, "molecula-chip") {
authorization = strings.Split(cookie, "molecula-chip=")[1:]
}
}
if len(authorization) == 0 {
return ctx, status.Errorf(codes.InvalidArgument, "missing authorization token")
}
}
token := strings.TrimPrefix(authorization[0], "Bearer ")
userinfo, err := auth.Authenticate(token)
if err != nil {
return ctx, status.Errorf(codes.Unauthenticated, err.Error())
}
return context.WithValue(ctx, "userinfo", userinfo), nil
}

View file

@ -3,13 +3,21 @@ package server_test
import (
"context"
"encoding/hex"
"fmt"
"io"
"os"
"path/filepath"
"reflect"
"strconv"
"strings"
"testing"
"time"
"github.com/molecula/featurebase/v2"
"github.com/golang-jwt/jwt"
pilosa "github.com/molecula/featurebase/v2"
"github.com/molecula/featurebase/v2/authn"
"github.com/molecula/featurebase/v2/authz"
"github.com/molecula/featurebase/v2/pql"
pb "github.com/molecula/featurebase/v2/proto"
"github.com/molecula/featurebase/v2/server"
@ -385,7 +393,7 @@ func TestQueryPQL(t *testing.T) {
m.MustCreateField(t, i.Name(), "f", pilosa.OptFieldKeys())
gh := server.NewGRPCHandler(m.API)
mock := &mockPilosa_QuerySQLServer{}
mock := &mockPilosa_QuerySQLServer{ctx: context.Background()}
err := gh.QueryPQL(&pb.QueryPQLRequest{
Index: i.Name(),
@ -941,7 +949,7 @@ func TestQuerySQL(t *testing.T) {
if strings.HasPrefix(test.sql, "drop table") {
t.Skip("drop statements can only run once")
}
mock := &mockPilosa_QuerySQLServer{}
mock := &mockPilosa_QuerySQLServer{ctx: context.Background()}
err := gh.QuerySQL(&pb.QuerySQLRequest{Sql: test.sql}, mock)
if err != nil {
t.Fatalf("sql: %s, error: %v", test.sql, err)
@ -968,7 +976,6 @@ func TestQuerySQLUnaryWithError(t *testing.T) {
ctx := grpc.NewContextWithServerTransportStream(context.Background(), stream)
gh, tearDownFunc := setUpTestQuerySQLUnary(ctx, t)
defer tearDownFunc()
tests := []struct {
sql string
err error
@ -1005,6 +1012,157 @@ func TestQuerySQLUnaryWithError(t *testing.T) {
}
})
}
permissions := `
"user-groups":
"dca35310-ecda-4f23-86cd-876aee55906b":
"grouper": "read"
"dca35310-ecda-4f23-86cd-876aee55906f":
"grouper": "write"
admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"`
permFile := writeTestFile(t, "permissions.yaml", permissions)
auth := server.Auth{
Enable: true,
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"},
SecretKey: "DEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEFDEADBEEF",
PermissionsFile: permFile,
}
var p authz.GroupPermissions
permsFile, err := os.Open(permFile)
if err != nil {
t.Fatal(err)
}
defer permsFile.Close()
if err = p.ReadPermissionsFile(permsFile); err != nil {
t.Fatal(err)
}
gh = gh.WithPerms(&p)
makeUser := func(groups []authn.Group, name string) *authn.UserInfo {
// make a valid token
tkn := jwt.New(jwt.SigningMethodHS256)
claims := tkn.Claims.(jwt.MapClaims)
groupString, _ := authn.ToGob64(groups)
claims["molecula-idp-groups"] = groupString
claims["oid"] = "42"
claims["name"] = name
secretKey, _ := hex.DecodeString(auth.SecretKey)
validToken, err := tkn.SignedString(secretKey)
if err != nil {
panic(err)
}
validToken = "Bearer " + validToken
adminUser := &authn.UserInfo{
UserID: "fake" + name,
UserName: name,
Groups: groups,
Token: validToken,
Expiry: time.Time{},
}
return adminUser
}
user := makeUser([]authn.Group{{GroupID: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe", GroupName: "adminGroup"}}, "admin")
adminCtx := context.WithValue(
ctx,
"userinfo",
user,
)
readuser := makeUser([]authn.Group{{GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: "readers"}}, "admin")
readCtx := context.WithValue(
ctx,
"userinfo",
readuser,
)
writeuser := makeUser([]authn.Group{{GroupID: "dca35310-ecda-4f23-86cd-876aee55906f", GroupName: "writers"}}, "admin")
writeCtx := context.WithValue(
ctx,
"userinfo",
writeuser,
)
sql := "select * from grouper"
t.Run("test-auth-with-admin-sqlUnary", func(t *testing.T) {
_, err := gh.QuerySQLUnary(adminCtx, &pb.QuerySQLRequest{Sql: sql})
if err != nil {
t.Fatal(err)
}
})
t.Run("test-auth-with-read-sqlUnary", func(t *testing.T) {
_, err := gh.QuerySQLUnary(readCtx, &pb.QuerySQLRequest{Sql: sql})
if err != nil {
t.Fatal(err)
}
})
t.Run("test-admin-auth-sql", func(t *testing.T) {
mock := &mockPilosa_QuerySQLServer{ctx: adminCtx}
err := gh.QuerySQL(&pb.QuerySQLRequest{Sql: sql}, mock)
if err != nil {
t.Fatal(err)
}
})
t.Run("test-admin-auth-get-index", func(t *testing.T) {
_, err := gh.GetIndex(adminCtx, &pb.GetIndexRequest{Name: "grouper"})
if err != nil {
t.Fatal(err)
}
})
t.Run("test-admin-auth-get-indexes", func(t *testing.T) {
_, err := gh.GetIndexes(adminCtx, &pb.GetIndexesRequest{})
if err != nil {
t.Fatal(err)
}
})
t.Run("test-admin-auth-pql", func(t *testing.T) {
_, err := gh.QueryPQLUnary(adminCtx, &pb.QueryPQLRequest{
Index: "grouper",
Pql: `Set(0, color="red")`,
})
if err != nil {
// Unary query should work
t.Fatal(err)
}
})
t.Run("test-write-with-read-auth-pql", func(t *testing.T) {
_, err := gh.QueryPQLUnary(readCtx, &pb.QueryPQLRequest{
Index: "grouper",
Pql: `Set(0, color="red")`,
})
if err == nil {
//should not be able to write
t.Fatal(err)
}
})
t.Run("test-write-with-write-auth-pql", func(t *testing.T) {
_, err := gh.QueryPQLUnary(writeCtx, &pb.QueryPQLRequest{
Index: "grouper",
Pql: `Set(0, color="red")`,
})
if err != nil {
//should be able to write
t.Fatal(err)
}
})
t.Run("test-write-with-admin-auth-pql", func(t *testing.T) {
_, err := gh.QueryPQLUnary(adminCtx, &pb.QueryPQLRequest{
Index: "grouper",
Pql: `Set(0, color="green")`,
})
if err != nil {
//should be able to write
t.Fatal(err)
}
})
}
func TestCRUDIndexes(t *testing.T) {
@ -1014,7 +1172,6 @@ func TestCRUDIndexes(t *testing.T) {
stream := &MockServerTransportStream{}
ctx := grpc.NewContextWithServerTransportStream(context.Background(), stream)
gh := server.NewGRPCHandler(m.API)
t.Run("CreateIndex", func(t *testing.T) {
// Try CreateIndex for testindex1
_, err := gh.CreateIndex(ctx, &pb.CreateIndexRequest{Name: "testindex1", Keys: true})
@ -1448,6 +1605,7 @@ func (stream *MockServerTransportStream) ClearMD() {
type mockPilosa_QuerySQLServer struct {
MockServerTransportStream
ctx context.Context
pb.Pilosa_QuerySQLServer
Results []*pb.RowResponse
}
@ -1470,9 +1628,19 @@ func (m *mockPilosa_QuerySQLServer) SetTrailer(md metadata.MD) {
}
func (m *mockPilosa_QuerySQLServer) Context() context.Context {
return context.Background()
return m.ctx
}
func (m *mockPilosa_QuerySQLServer) clearResults() {
m.Results = m.Results[:0]
}
func writeTestFile(t *testing.T, filename, content string) string {
fname := filepath.Join(t.TempDir(), filename)
f, err := os.Create(fname)
if err != nil {
panic(filename)
}
io.WriteString(f, content)
defer f.Close()
return fname
}

View file

@ -85,8 +85,7 @@ type Command struct {
pgserver *PostgresServer
serverOptions []pilosa.ServerOption
auth *authn.Auth
auth *authn.Auth
}
type CommandOption func(c *Command) error
@ -110,6 +109,8 @@ func OptCommandConfig(config *Config) CommandOption {
defer c.Config.MustValidate()
if c.Config != nil {
c.Config.Etcd = config.Etcd
c.Config.Auth = config.Auth
c.Config.TLS = config.TLS
return nil
}
c.Config = config
@ -481,7 +482,6 @@ func (m *Command) SetupServer() error {
pilosa.OptServerStatsClient(statsClient),
pilosa.OptServerURI(advertiseURI),
pilosa.OptServerGRPCURI(advertiseGRPCURI),
pilosa.OptServerInternalClient(http.NewInternalClientFromURI(uri, c)),
pilosa.OptServerClusterName(m.Config.Cluster.Name),
pilosa.OptServerSerializer(proto.Serializer{}),
pilosa.OptServerStorageConfig(m.Config.Storage),
@ -497,6 +497,12 @@ func (m *Command) SetupServer() error {
serverOptions = append(serverOptions, m.serverOptions...)
if m.Config.Auth.Enable {
serverOptions = append(serverOptions, pilosa.OptServerInternalClient(http.NewInternalClientFromURI(uri, c, http.WithSecretKey(m.Config.Auth.SecretKey))))
} else {
serverOptions = append(serverOptions, pilosa.OptServerInternalClient(http.NewInternalClientFromURI(uri, c)))
}
m.Server, err = pilosa.NewServer(serverOptions...)
if err != nil {
@ -514,13 +520,6 @@ func (m *Command) SetupServer() error {
// Tell server about its new API, which its client will need.
m.Server.SetAPI(m.API)
m.grpcServer, err = NewGRPCServer(
OptGRPCServerAPI(m.API),
OptGRPCServerListener(m.grpcLn),
OptGRPCServerTLSConfig(m.tlsConfig),
OptGRPCServerLogger(m.logger),
OptGRPCServerStats(statsClient),
)
if err != nil {
return errors.Wrap(err, "new grpc server")
}
@ -539,7 +538,7 @@ func (m *Command) SetupServer() error {
}
ac := m.Config.Auth
m.auth, err = authn.NewAuth(m.logger, m.listenURI.String(), ac.Scopes, ac.AuthorizeURL, ac.TokenURL, ac.GroupEndpointURL, ac.LogoutURL, ac.ClientId, ac.ClientSecret, ac.HashKey, ac.BlockKey)
m.auth, err = authn.NewAuth(m.logger, m.listenURI.String(), ac.Scopes, ac.AuthorizeURL, ac.TokenURL, ac.GroupEndpointURL, ac.LogoutURL, ac.ClientId, ac.ClientSecret, ac.SecretKey)
if err != nil {
return errors.Wrap(err, "instantiating authN object")
}
@ -562,6 +561,16 @@ func (m *Command) SetupServer() error {
}
m.grpcServer, err = NewGRPCServer(
OptGRPCServerAPI(m.API),
OptGRPCServerListener(m.grpcLn),
OptGRPCServerTLSConfig(m.tlsConfig),
OptGRPCServerLogger(m.logger),
OptGRPCServerStats(statsClient),
OptGRPCServerAuth(m.auth),
OptGRPCServerPerm(&p),
)
m.Handler, err = http.NewHandler(
http.OptHandlerAllowedOrigins(m.Config.Handler.AllowedOrigins),
http.OptHandlerAPI(m.API),

View file

@ -0,0 +1,40 @@
package server
import (
"fmt"
"testing"
)
// unit tests for internal functions
func TestIsAllowed(t *testing.T) {
cases := []struct {
requested []string
allowed []string
expected bool
}{
{
requested: []string{"a", "b", "c"},
allowed: []string{"a", "b", "c", "d", "e"},
expected: true,
},
{
requested: []string{"a", "b", "c", "f"},
allowed: []string{"a", "b", "c", "d", "e"},
expected: false,
},
{
requested: []string{"a", "b", "c"},
allowed: []string{},
expected: false,
},
}
for i, test := range cases {
t.Run(fmt.Sprint(i), func(t *testing.T) {
if res := isAllowed(test.requested, test.allowed); res != test.expected {
t.Errorf("expected %v, got %v", test.expected, res)
}
})
}
}