diff --git a/http/handler.go b/http/handler.go index efd20d7e2..b3d6339b8 100644 --- a/http/handler.go +++ b/http/handler.go @@ -433,20 +433,20 @@ func newRouter(handler *Handler) http.Handler { router.HandleFunc("/transaction/{id}", handler.chkAuthZ(handler.handlePostTransaction, authz.Read)).Methods("POST").Name("PostTransaction") router.HandleFunc("/transaction/{id}/finish", handler.chkAuthZ(handler.handlePostFinishTransaction, authz.Read)).Methods("POST").Name("PostFinishTransaction") router.HandleFunc("/transactions", handler.chkAuthZ(handler.handleGetTransactions, authz.Read)).Methods("GET").Name("GetTransactions") - router.HandleFunc("/queries", handler.chkAuthZ(handler.handleGetActiveQueries, authz.Read)).Methods("GET").Name("GetActiveQueries") - router.HandleFunc("/query-history", handler.chkAuthZ(handler.handleGetPastQueries, authz.Read)).Methods("GET").Name("GetPastQueries") - router.HandleFunc("/version", handler.chkAuthZ(handler.handleGetVersion, authz.Read)).Methods("GET").Name("GetVersion") + router.HandleFunc("/queries", handler.chkAuthZ(handler.handleGetActiveQueries, authz.Admin)).Methods("GET").Name("GetActiveQueries") + router.HandleFunc("/query-history", handler.chkAuthZ(handler.handleGetPastQueries, authz.Admin)).Methods("GET").Name("GetPastQueries") + router.HandleFunc("/version", handler.handleGetVersion).Methods("GET").Name("GetVersion") // /ui endpoints are for UI use; they may change at any time. router.HandleFunc("/ui/usage", handler.chkAuthZ(handler.handleGetUsage, authz.Read)).Methods("GET").Name("GetUsage") router.HandleFunc("/ui/transaction", handler.chkAuthZ(handler.handleGetTransactionList, authz.Read)).Methods("GET").Name("GetTransactionList") router.HandleFunc("/ui/transaction/", handler.chkAuthZ(handler.handleGetTransactionList, authz.Read)).Methods("GET").Name("GetTransactionList") - router.HandleFunc("/ui/shard-distribution", handler.chkAuthZ(handler.handleGetShardDistribution, authz.Read)).Methods("GET").Name("GetShardDistribution") + router.HandleFunc("/ui/shard-distribution", handler.chkAuthZ(handler.handleGetShardDistribution, authz.Admin)).Methods("GET").Name("GetShardDistribution") // /internal endpoints are for internal use only; they may change at any time. // DO NOT rely on these for external applications! - // Truly used internally by featurebease + // Truly used internally by featurebase 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") @@ -459,23 +459,23 @@ func newRouter(handler *Handler) http.Handler { router.HandleFunc("/internal/partition/nodes", handler.chkAuthN(handler.handleGetPartitionNodes)).Methods("GET").Name("GetPartitionNodes") 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") - router.HandleFunc("/internal/index/{index}/field/{field}/remote-available-shards/{shardID}", handler.chkAuthN(handler.handleDeleteRemoteAvailableShard)).Methods("DELETE") - router.HandleFunc("/internal/index/{index}/shard/{shard}/snapshot", handler.chkAuthN(handler.handleGetIndexShardSnapshot)).Methods("GET").Name("GetIndexShardSnapshot") - router.HandleFunc("/internal/index/{index}/shards", handler.chkAuthN(handler.handleGetIndexAvailableShards)).Methods("GET").Name("GetIndexAvailableShards") + router.HandleFunc("/internal/index/{index}/field/{field}/mutex-check", handler.chkAuthZ(handler.handleInternalGetMutexCheck, authz.Read)).Methods("GET").Name("InternalGetMutexCheck") + router.HandleFunc("/internal/index/{index}/field/{field}/remote-available-shards/{shardID}", handler.chkAuthZ(handler.handleDeleteRemoteAvailableShard, authz.Admin)).Methods("DELETE") + router.HandleFunc("/internal/index/{index}/shard/{shard}/snapshot", handler.chkAuthZ(handler.handleGetIndexShardSnapshot, authz.Read)).Methods("GET").Name("GetIndexShardSnapshot") + router.HandleFunc("/internal/index/{index}/shards", handler.chkAuthZ(handler.handleGetIndexAvailableShards, authz.Read)).Methods("GET").Name("GetIndexAvailableShards") router.HandleFunc("/internal/nodes", handler.chkAuthN(handler.handleGetNodes)).Methods("GET").Name("GetNodes") router.HandleFunc("/internal/shards/max", handler.chkAuthN(handler.handleGetShardsMax)).Methods("GET").Name("GetShardsMax") // TODO: deprecate, but it's being used by the client - router.HandleFunc("/internal/ingest/{index}", handler.chkAuthN(handler.handlePostIngestData)).Methods("POST").Name("PostIngestData") - router.HandleFunc("/internal/ingest/{index}/node", handler.chkAuthN(handler.handlePostIngestNode)).Methods("POST").Name("PostIngestNode") + router.HandleFunc("/internal/ingest/{index}", handler.chkAuthZ(handler.handlePostIngestData, authz.Write)).Methods("POST").Name("PostIngestData") + router.HandleFunc("/internal/ingest/{index}/node", handler.chkAuthZ(handler.handlePostIngestNode, authz.Write)).Methods("POST").Name("PostIngestNode") - router.HandleFunc("/internal/schema", handler.chkAuthN(handler.handleIngestSchema)).Methods("POST").Name("PostIngestSchema") - router.HandleFunc("/internal/translate/index/{index}/keys/find", handler.chkAuthN(handler.handleFindIndexKeys)).Methods("POST").Name("FindIndexKeys") - router.HandleFunc("/internal/translate/index/{index}/keys/create", handler.chkAuthN(handler.handleCreateIndexKeys)).Methods("POST").Name("CreateIndexKeys") - router.HandleFunc("/internal/translate/index/{index}/{partition}", handler.chkAuthN(handler.handlePostTranslateIndexDB)).Methods("POST").Name("PostTranslateIndexDB") - router.HandleFunc("/internal/translate/field/{index}/{field}", handler.chkAuthN(handler.handlePostTranslateFieldDB)).Methods("POST").Name("PostTranslateFieldDB") - router.HandleFunc("/internal/translate/field/{index}/{field}/keys/find", handler.chkAuthN(handler.handleFindFieldKeys)).Methods("POST").Name("FindFieldKeys") - router.HandleFunc("/internal/translate/field/{index}/{field}/keys/create", handler.chkAuthN(handler.handleCreateFieldKeys)).Methods("POST").Name("CreateFieldKeys") - router.HandleFunc("/internal/translate/field/{index}/{field}/keys/like", handler.chkAuthN(handler.handleMatchField)).Methods("POST").Name("MatchFieldKeys") + router.HandleFunc("/internal/schema", handler.chkAuthZ(handler.handleIngestSchema, authz.Admin)).Methods("POST").Name("PostIngestSchema") + router.HandleFunc("/internal/translate/index/{index}/keys/find", handler.chkAuthZ(handler.handleFindIndexKeys, authz.Admin)).Methods("POST").Name("FindIndexKeys") + router.HandleFunc("/internal/translate/index/{index}/keys/create", handler.chkAuthZ(handler.handleCreateIndexKeys, authz.Admin)).Methods("POST").Name("CreateIndexKeys") + router.HandleFunc("/internal/translate/index/{index}/{partition}", handler.chkAuthZ(handler.handlePostTranslateIndexDB, authz.Admin)).Methods("POST").Name("PostTranslateIndexDB") + router.HandleFunc("/internal/translate/field/{index}/{field}", handler.chkAuthZ(handler.handlePostTranslateFieldDB, authz.Admin)).Methods("POST").Name("PostTranslateFieldDB") + router.HandleFunc("/internal/translate/field/{index}/{field}/keys/find", handler.chkAuthZ(handler.handleFindFieldKeys, authz.Admin)).Methods("POST").Name("FindFieldKeys") + router.HandleFunc("/internal/translate/field/{index}/{field}/keys/create", handler.chkAuthZ(handler.handleCreateFieldKeys, authz.Admin)).Methods("POST").Name("CreateFieldKeys") + router.HandleFunc("/internal/translate/field/{index}/{field}/keys/like", handler.chkAuthZ(handler.handleMatchField, authz.Read)).Methods("POST").Name("MatchFieldKeys") router.HandleFunc("/internal/idalloc/reserve", handler.chkAuthN(handler.handleReserveIDs)).Methods("POST").Name("ReserveIDs") router.HandleFunc("/internal/idalloc/commit", handler.chkAuthN(handler.handleCommitIDs)).Methods("POST").Name("CommitIDs") @@ -483,9 +483,9 @@ func newRouter(handler *Handler) http.Handler { router.HandleFunc("/internal/idalloc/reset/{index}", handler.chkAuthN(handler.handleResetIDAlloc)).Methods("POST").Name("ResetIDAlloc") router.HandleFunc("/internal/idalloc/data", handler.chkAuthN(handler.handleIDAllocData)).Methods("GET").Name("IDAllocData") - router.HandleFunc("/internal/restore/{index}/{shardID}", handler.chkAuthN(handler.handlePostRestore)).Methods("POST").Name("Restore") + router.HandleFunc("/internal/restore/{index}/{shardID}", handler.chkAuthZ(handler.handlePostRestore, authz.Admin)).Methods("POST").Name("Restore") - router.HandleFunc("/internal/debug/rbf", handler.chkAuthN(handler.handleGetInternalDebugRBFJSON)).Methods("GET").Name("GetInternalDebugRBFJSON") + router.HandleFunc("/internal/debug/rbf", handler.chkAuthZ(handler.handleGetInternalDebugRBFJSON, authz.Admin)).Methods("GET").Name("GetInternalDebugRBFJSON") // endpoints for collecting cpu profiles from a chosen begin point to // when the client wants to stop. Used for profiling imports that @@ -619,10 +619,16 @@ 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, uinfo.Groups) indexName, ok := mux.Vars(r)["index"] - if ok { + + if !ok { + indexName = r.URL.Query().Get("index") + } + + if indexName != "" { p, err := h.permissions.GetPermissions(uinfo, indexName) - ctx = context.WithValue(ctx, contextKeyPermission, p) + ctx = context.WithValue(r.Context(), contextKeyPermission, p) if err != nil { w.Header().Add("Content-Type", "text/plain") http.Error(w, errors.Wrap(err, "Insufficient Permissions").Error(), http.StatusForbidden) @@ -811,30 +817,6 @@ func headerAcceptRoaringRow(header http.Header) bool { return false } -func (h *Handler) filterResponse(w http.ResponseWriter, r *http.Request, schema []*pilosa.IndexInfo) []*pilosa.IndexInfo { - if h.auth != nil { - g := r.Context().Value(contextKeyGroupMembership) - if g == nil { - http.Error(w, "Forbidden", http.StatusForbidden) - return nil - } - indexes := h.permissions.GetAuthorizedIndexList(g.([]authn.Group), authz.Read) - var new []*pilosa.IndexInfo - for _, s := range schema { - for _, index := range indexes { - if s.Name == index { - new = append(new, s) - } - } - - } - return new - - } - return schema - -} - // handleGetSchema handles GET /schema requests. func (h *Handler) handleGetSchema(w http.ResponseWriter, r *http.Request) { if !validHeaderAcceptJSON(r.Header) { @@ -851,7 +833,27 @@ func (h *Handler) handleGetSchema(w http.ResponseWriter, r *http.Request) { h.logger.Printf("getting schema error: %s", err) } - schema = h.filterResponse(w, r, schema) + // if auth is turned on, filter response to only include authorized indexes + if h.auth != nil { + g := r.Context().Value(contextKeyGroupMembership) + if g == nil { + http.Error(w, "Forbidden", http.StatusForbidden) + return + } + if !h.permissions.IsAdmin(g.([]authn.Group)) { + var filtered []*pilosa.IndexInfo + allowed := h.permissions.GetAuthorizedIndexList(g.([]authn.Group), authz.Read) + for _, s := range schema { + for _, index := range allowed { + if s.Name == index { + filtered = append(filtered, s) + break + } + } + } + schema = filtered + } + } if err := json.NewEncoder(w).Encode(pilosa.Schema{Indexes: schema}); err != nil { h.logger.Errorf("write schema response error: %s", err) @@ -871,7 +873,28 @@ func (h *Handler) handleGetSchemaDetails(w http.ResponseWriter, r *http.Request) h.logger.Printf("error getting detailed schema: %s", err) return } - schema = h.filterResponse(w, r, schema) + + // if auth is turned on, filter response to only include authorized indexes + if h.auth != nil { + g := r.Context().Value(contextKeyGroupMembership) + if g == nil { + http.Error(w, "Forbidden", http.StatusForbidden) + return + } + if !h.permissions.IsAdmin(g.([]authn.Group)) { + var filtered []*pilosa.IndexInfo + allowed := h.permissions.GetAuthorizedIndexList(g.([]authn.Group), authz.Read) + for _, s := range schema { + for _, index := range allowed { + if s.Name == index { + filtered = append(filtered, s) + break + } + } + } + schema = filtered + } + } if err := json.NewEncoder(w).Encode(pilosa.Schema{Indexes: schema}); err != nil { h.logger.Printf("write schema response error: %s", err) } @@ -917,6 +940,38 @@ func (h *Handler) handleGetUsage(w http.ResponseWriter, r *http.Request) { http.Error(w, err.Error(), http.StatusInternalServerError) } + // if auth is turned on, filter results + if h.auth != nil { + g := r.Context().Value(contextKeyGroupMembership) + if g == nil { + http.Error(w, "Forbidden", http.StatusForbidden) + return + } + if !h.permissions.IsAdmin(g.([]authn.Group)) { + allowed := h.permissions.GetAuthorizedIndexList(g.([]authn.Group), authz.Read) + filteredNodeUsages := map[string]pilosa.NodeUsage{} + + for nodeId, nodeUsage := range nodeUsages { + filteredIndexUsage := pilosa.NodeUsage{ + Disk: pilosa.DiskUsage{ + IndexUsage: map[string]pilosa.IndexUsage{}, + }, + } + for index, idxUsage := range nodeUsage.Disk.IndexUsage { + // is it in auth list + for _, authd := range allowed { + if index == authd { + filteredIndexUsage.Disk.IndexUsage[index] = idxUsage + break + } + } + } + filteredNodeUsages[nodeId] = filteredIndexUsage + } + nodeUsages = filteredNodeUsages + } + } + w.Header().Set("Content-Type", "application/json") if err := json.NewEncoder(w).Encode(nodeUsages); err != nil { h.logger.Errorf("write status response error: %s", err) diff --git a/server/grpc.go b/server/grpc.go index 3683de99a..804dd91cd 100644 --- a/server/grpc.go +++ b/server/grpc.go @@ -166,15 +166,17 @@ func (h *GRPCHandler) QuerySQL(req *pb.QuerySQLRequest, stream pb.Pilosa_QuerySQ if err != nil { return errors.Wrap(err, "parsing SQL") } + allowed := h.perms.GetAuthorizedIndexList(uinfo.(*authn.UserInfo).Groups, authz.Read) if !h.perms.IsAdmin(uinfo.(*authn.UserInfo).Groups) { - if !isAllowed(parsed.Tables, h.perms.GetAuthorizedIndexList(uinfo.(*authn.UserInfo).Groups, authz.Read)) { + if !isAllowed(parsed.Tables, allowed) { return status.Error(codes.PermissionDenied, "insufficient permissions to access requested tables") } + ctx = context.WithValue(ctx, "indices", allowed) } } start := time.Now() - results, err := h.execSQL(stream.Context(), req.Sql) + results, err := h.execSQL(ctx, req.Sql) duration := time.Since(start) if err != nil { return err @@ -208,6 +210,10 @@ func (h *GRPCHandler) QuerySQL(req *pb.QuerySQLRequest, stream pb.Pilosa_QuerySQ // https://github.com/molecula/pilosa/pull/644 func (h *GRPCHandler) QuerySQLUnary(ctx context.Context, req *pb.QuerySQLRequest) (*pb.TableResponse, error) { start := time.Now() + uinfo := ctx.Value("userinfo") + if uinfo != nil && !h.perms.IsAdmin(uinfo.(*authn.UserInfo).Groups) { + ctx = context.WithValue(ctx, "indices", h.perms.GetAuthorizedIndexList(uinfo.(*authn.UserInfo).Groups, authz.Read)) + } results, err := h.execSQL(ctx, req.Sql) if err != nil { return nil, err diff --git a/server/grpc_test.go b/server/grpc_test.go index 63482f687..da56b02ca 100644 --- a/server/grpc_test.go +++ b/server/grpc_test.go @@ -970,7 +970,7 @@ func TestQuerySQL(t *testing.T) { } } -func TestQuerySQLUnaryWithError(t *testing.T) { +func TestQuerySQLWithError(t *testing.T) { stream := &MockServerTransportStream{} ctx := grpc.NewContextWithServerTransportStream(context.Background(), stream) @@ -1060,14 +1060,13 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"` } validToken = "Bearer " + validToken - adminUser := &authn.UserInfo{ + return &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") @@ -1076,7 +1075,7 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"` "userinfo", user, ) - readuser := makeUser([]authn.Group{{GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: "readers"}}, "admin") + readuser := makeUser([]authn.Group{{GroupID: "dca35310-ecda-4f23-86cd-876aee55906b", GroupName: "readers"}}, "reader") readCtx := context.WithValue( ctx, "userinfo", @@ -1110,6 +1109,16 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"` t.Fatal(err) } }) + + t.Run("test-admin-auth-sql-show", func(t *testing.T) { + mock := &mockPilosa_QuerySQLServer{ctx: adminCtx} + + err := gh.QuerySQL(&pb.QuerySQLRequest{Sql: "show tables"}, 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 { @@ -1143,6 +1152,23 @@ admin: "ac97c9e2-346b-42a2-b6da-18bcb61a32fe"` t.Fatal(err) } }) + t.Run("test-show-tables-unary", func(t *testing.T) { + response, err := gh.QuerySQLUnary(readCtx, &pb.QuerySQLRequest{ + Sql: "show tables", + }) + + if err != nil && len(response.Rows) != 1 { + t.Fatal(err) + } + }) + t.Run("test-show-tables-unary-admin", func(t *testing.T) { + response, err := gh.QuerySQLUnary(adminCtx, &pb.QuerySQLRequest{ + Sql: "show tables", + }) + if err != nil && len(response.Rows) != 3 { + t.Fatal(err) + } + }) t.Run("test-write-with-write-auth-pql", func(t *testing.T) { _, err := gh.QueryPQLUnary(writeCtx, &pb.QueryPQLRequest{ Index: "grouper", diff --git a/sql/show.go b/sql/show.go index 2848a77d7..186d51923 100644 --- a/sql/show.go +++ b/sql/show.go @@ -5,9 +5,11 @@ import ( "context" "fmt" - "github.com/molecula/featurebase/v2" + pilosa "github.com/molecula/featurebase/v2" pproto "github.com/molecula/featurebase/v2/proto" "github.com/pkg/errors" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" "vitess.io/vitess/go/vt/sqlparser" ) @@ -46,16 +48,32 @@ func (s *ShowHandler) execShowTables(ctx context.Context, showStmt *sqlparser.Sh return nil, errors.Wrap(err, "getting schema") } - result := make(pproto.ConstRowser, len(indexInfo)) - for i, ii := range indexInfo { - result[i] = pproto.RowResponse{ + allowed, ok := ctx.Value("indices").([]string) + + result := make(pproto.ConstRowser, 0) + for _, ii := range indexInfo { + if ok { + // if authorization is turned on, allowed will be a list + // so we have to check if the index is in the allowed list + found := false + for _, idx := range allowed { + if ii.Name == idx { + found = true + break + } + } + if !found { + continue + } + } + result = append(result, pproto.RowResponse{ Headers: []*pproto.ColumnInfo{ {Name: "Table", Datatype: "string"}, }, Columns: []*pproto.ColumnResponse{ {ColumnVal: &pproto.ColumnResponse_StringVal{StringVal: ii.Name}}, }, - } + }) } // Sort the result. @@ -64,6 +82,19 @@ func (s *ShowHandler) execShowTables(ctx context.Context, showStmt *sqlparser.Sh func (s *ShowHandler) execShowFields(ctx context.Context, showStmt *sqlparser.Show) (pproto.ToRowser, error) { indexName := showStmt.OnTable.ToViewName().Name.String() + allowed, ok := ctx.Value("indices").([]string) + if ok { + found := false + for _, idx := range allowed { + if idx == indexName { + found = true + break + } + } + if !found { + return nil, status.Error(codes.PermissionDenied, "insufficient permissions to access requested tables") + } + } index, err := s.api.Index(ctx, indexName) if err != nil { return nil, errors.Wrap(err, "getting schema")