diff --git a/cluster.go b/cluster.go index 092233055..8860d8405 100644 --- a/cluster.go +++ b/cluster.go @@ -167,6 +167,7 @@ type Cluster struct { // Close management wg sync.WaitGroup closing chan struct{} + prefect SecurityManager // The writer for any logging. LogOutput io.Writer @@ -185,6 +186,7 @@ func NewCluster() *Cluster { closing: make(chan struct{}), LogOutput: os.Stderr, + prefect: &NopSecurityManager{}, } } @@ -226,6 +228,20 @@ func (c *Cluster) NodeSet() []URI { } func (c *Cluster) setState(state string) { + // Ignore cases where the state hasn't changed. + if state == c.State { + return + } + + switch state { + case ClusterStateResizing: + c.prefect.SetRestricted() + case ClusterStateNormal: + c.prefect.SetNormal() + // Don't change routing for these states: + // - ClusterStateStarting + } + c.State = state } diff --git a/handler.go b/handler.go index 88238b9ea..55b6e91d9 100644 --- a/handler.go +++ b/handler.go @@ -60,7 +60,9 @@ type Handler struct { Cluster *Cluster ClientOptions *ClientOptions - Router *mux.Router + Router *mux.Router + NormalRouter *mux.Router + RestrictedRouter *mux.Router // The execution engine for running queries. Executor interface { @@ -89,14 +91,54 @@ func NewHandler() *Handler { handler := &Handler{ LogOutput: os.Stderr, } - handler.Router = NewRouter(handler) + BuildRouters(handler) return handler } -// NewRouter creates a Gorilla Mux http router. -func NewRouter(handler *Handler) *mux.Router { +// BuildRouters creates Gorilla Mux http routers for both normal and restricted endpoints. +func BuildRouters(handler *Handler) { + // Normal router. router := mux.NewRouter() - router.HandleFunc("/", handler.handleWebUI).Methods("GET") + loadCommon(router, handler) + loadNormal(router, handler) + handler.NormalRouter = router + + // Restricted router. + router = mux.NewRouter() + loadCommon(router, handler) + loadRestricted(router, handler) + handler.RestrictedRouter = router + + handler.SetNormal() +} + +// SetNormal is a method of the SecurityManager interface which provides normal URI routing. +func (h *Handler) SetNormal() { + h.Router = h.NormalRouter +} + +// SetRestricted is a method of the SecurityManager interface which provides restricted URI routing. +func (h *Handler) SetRestricted() { + h.Router = h.RestrictedRouter +} + +func loadCommon(router *mux.Router, handler *Handler) { + router.HandleFunc("/schema", handler.handleGetSchema).Methods("GET") + router.HandleFunc("/status", handler.handleGetStatus).Methods("GET") + router.HandleFunc("/version", handler.handleGetVersion).Methods("GET") + router.PathPrefix("/debug/pprof/").Handler(http.DefaultServeMux).Methods("GET") + router.HandleFunc("/debug/vars", handler.handleExpvar).Methods("GET") + router.HandleFunc("/slices/max", handler.handleGetSlicesMax).Methods("GET") // TODO: deprecate, but it's being used by the client (for backups) + router.HandleFunc("/fragment/data", handler.handleGetFragmentData).Methods("GET") + router.HandleFunc("/hosts", handler.handleGetHosts).Methods("GET") +} + +func loadRestricted(router *mux.Router, handler *Handler) { + router.HandleFunc("/cluster/resize/abort", handler.handlePostClusterResizeAbort).Methods("POST") + router.NotFoundHandler = http.HandlerFunc(handler.reportRestricted) +} + +func loadNormal(router *mux.Router, handler *Handler) { router.HandleFunc("/assets/{file}", handler.handleWebUI).Methods("GET") router.HandleFunc("/cluster/resize/abort", handler.handlePostClusterResizeAbort).Methods("POST") router.PathPrefix("/debug/pprof/").Handler(http.DefaultServeMux).Methods("GET") @@ -104,7 +146,6 @@ func NewRouter(handler *Handler) *mux.Router { router.HandleFunc("/export", handler.handleGetExport).Methods("GET") router.HandleFunc("/fragment/block/data", handler.handleGetFragmentBlockData).Methods("GET") router.HandleFunc("/fragment/blocks", handler.handleGetFragmentBlocks).Methods("GET") - router.HandleFunc("/fragment/data", handler.handleGetFragmentData).Methods("GET") router.HandleFunc("/fragment/data", handler.handlePostFragmentData).Methods("POST") router.HandleFunc("/fragment/nodes", handler.handleGetFragmentNodes).Methods("GET") router.HandleFunc("/import", handler.handlePostImport).Methods("POST") @@ -131,11 +172,6 @@ func NewRouter(handler *Handler) *mux.Router { router.HandleFunc("/index/{index}/input-definition/{input-definition}", handler.handleDeleteInputDefinition).Methods("DELETE") router.HandleFunc("/index/{index}/query", handler.handlePostQuery).Methods("POST") router.HandleFunc("/index/{index}/time-quantum", handler.handlePatchIndexTimeQuantum).Methods("PATCH") - router.HandleFunc("/hosts", handler.handleGetHosts).Methods("GET") - router.HandleFunc("/schema", handler.handleGetSchema).Methods("GET") - router.HandleFunc("/slices/max", handler.handleGetSlicesMax).Methods("GET") // TODO: deprecate, but it's being used by the client (for backups) - router.HandleFunc("/status", handler.handleGetStatus).Methods("GET") - router.HandleFunc("/version", handler.handleGetVersion).Methods("GET") router.HandleFunc("/recalculate-caches", handler.handleRecalculateCaches).Methods("POST") // TODO: Apply MethodNotAllowed statuses to all endpoints. @@ -144,7 +180,10 @@ func NewRouter(handler *Handler) *mux.Router { // For now we just do it for the most commonly used handler, /query router.HandleFunc("/index/{index}/query", handler.methodNotAllowedHandler).Methods("GET") - return router +} + +func (h *Handler) reportRestricted(w http.ResponseWriter, r *http.Request) { + http.Error(w, "not allowed during resize", http.StatusMethodNotAllowed) } func (h *Handler) methodNotAllowedHandler(w http.ResponseWriter, r *http.Request) { diff --git a/security_manager.go b/security_manager.go new file mode 100644 index 000000000..2acf82986 --- /dev/null +++ b/security_manager.go @@ -0,0 +1,22 @@ +package pilosa + +// SecurityManager provides the ability to limit access to restricted endpoints +// during cluster configuration. +type SecurityManager interface { + SetRestricted() + SetNormal() +} + +// NopSecurityManager provides a no-op implementation of the SecurityManager interface. +type NopSecurityManager struct { +} + +// SetRestricted no-op. +func (sdm *NopSecurityManager) SetRestricted() { + +} + +// SetNormal no-op. +func (sdm *NopSecurityManager) SetNormal() { + +}