From 8e6041d0fe36fb9d3e30812f425d76efb062d8f9 Mon Sep 17 00:00:00 2001 From: Yuce Tekol Date: Mon, 19 Feb 2018 13:55:19 +0300 Subject: [PATCH 1/4] Validate /fragment/nodes arguments. Fixes #652 --- handler.go | 60 ++++++++++++++++++++++++++++++++++++++++++++++++- handler_test.go | 9 ++++++++ 2 files changed, 68 insertions(+), 1 deletion(-) diff --git a/handler.go b/handler.go index e5cce4648..f874b4103 100644 --- a/handler.go +++ b/handler.go @@ -27,6 +27,7 @@ import ( "io/ioutil" "log" "net/http" + "net/url" // Imported for its side-effect of registering pprof endpoints with the server. _ "net/http/pprof" "os" @@ -70,6 +71,9 @@ type Handler struct { // The writer for any logging. LogOutput io.Writer + + // Keeps the query argument validators for each handler + validators map[string]queryValidationSpec } // externalPrefixFlag denotes endpoints that are intended to be exposed to clients. @@ -91,9 +95,19 @@ func NewHandler() *Handler { LogOutput: os.Stderr, } handler.Router = NewRouter(handler) + handler.populateValidators() return handler } +func (h *Handler) populateValidators() { + h.validators = map[string]queryValidationSpec{} + // GetFragmentNodes validator + validator := NewQueryValidationSpec() + validator.Required("slice") + validator.Optional("index") + h.validators["GetFragmentNodes"] = validator +} + // NewRouter creates a Gorilla Mux http router. func NewRouter(handler *Handler) *mux.Router { router := mux.NewRouter() @@ -1342,12 +1356,18 @@ func (h *Handler) handleGetExportCSV(w http.ResponseWriter, r *http.Request) { // handleGetFragmentNodes handles /fragment/nodes requests. func (h *Handler) handleGetFragmentNodes(w http.ResponseWriter, r *http.Request) { q := r.URL.Query() + validator := h.validators["GetFragmentNodes"] + err := validator.validate(q) + if err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } index := q.Get("index") // Read slice parameter. slice, err := strconv.ParseUint(q.Get("slice"), 10, 64) if err != nil { - http.Error(w, "slice required", http.StatusBadRequest) + http.Error(w, "slice should be an unsigned integer", http.StatusBadRequest) return } @@ -2057,3 +2077,41 @@ func (h *Handler) handleGetID(w http.ResponseWriter, r *http.Request) { } type defaultClusterMessageResponse struct{} + +type queryValidationSpec struct { + required []string + args map[string]bool +} + +func NewQueryValidationSpec() queryValidationSpec { + return queryValidationSpec{ + args: map[string]bool{}, + } +} + +func (s *queryValidationSpec) Required(args ...string) { + s.required = args + for _, arg := range args { + s.args[arg] = true + } +} + +func (s *queryValidationSpec) Optional(args ...string) { + for _, arg := range args { + s.args[arg] = true + } +} + +func (s queryValidationSpec) validate(query url.Values) error { + for _, req := range s.required { + if query.Get(req) == "" { + return errors.New(fmt.Sprintf("%s is required", req)) + } + } + for k, _ := range query { + if _, ok := s.args[k]; !ok { + return errors.New(fmt.Sprintf("%s is not a valid argument", k)) + } + } + return nil +} diff --git a/handler_test.go b/handler_test.go index 585230b15..636665022 100644 --- a/handler_test.go +++ b/handler_test.go @@ -1198,6 +1198,15 @@ func TestHandler_Fragment_Nodes(t *testing.T) { } else if w.Body.String() != `[{"scheme":"http","host":"host2"},{"scheme":"http","host":"host0"}]`+"\n" { t.Fatalf("unexpected body: %q", w.Body.String()) } + + // invalid argument should return BadRequest + w = httptest.NewRecorder() + r = test.MustNewHTTPRequest("GET", "/fragment/nodes?db=X&slice=0", nil) + h.ServeHTTP(w, r) + if w.Code != http.StatusBadRequest { + t.Fatalf("unexpected status code: %d", w.Code) + } + } // Ensure the handler can return expvars without panicking. From c1f9dd4e593e0229c196fd545b09d017abe7f43a Mon Sep 17 00:00:00 2001 From: Yuce Tekol Date: Tue, 20 Feb 2018 15:38:58 +0300 Subject: [PATCH 2/4] Added query arg validator middleware; updated Gorilla/mux --- Gopkg.lock | 4 ++-- handler.go | 55 ++++++++++++++++++++++++++++++------------------------ 2 files changed, 33 insertions(+), 26 deletions(-) diff --git a/Gopkg.lock b/Gopkg.lock index a8071c915..c77dc951a 100644 --- a/Gopkg.lock +++ b/Gopkg.lock @@ -88,8 +88,8 @@ [[projects]] name = "github.com/gorilla/mux" packages = ["."] - revision = "7f08801859139f86dfafd1c296e2cba9a80d292e" - version = "v1.6.0" + revision = "53c1911da2b537f792e7cafcb446b05ffe33b996" + version = "v1.6.1" [[projects]] branch = "master" diff --git a/handler.go b/handler.go index f874b4103..895141135 100644 --- a/handler.go +++ b/handler.go @@ -73,7 +73,7 @@ type Handler struct { LogOutput io.Writer // Keeps the query argument validators for each handler - validators map[string]queryValidationSpec + validators map[string]*queryValidationSpec } // externalPrefixFlag denotes endpoints that are intended to be exposed to clients. @@ -100,12 +100,23 @@ func NewHandler() *Handler { } func (h *Handler) populateValidators() { - h.validators = map[string]queryValidationSpec{} - // GetFragmentNodes validator - validator := NewQueryValidationSpec() - validator.Required("slice") - validator.Optional("index") - h.validators["GetFragmentNodes"] = validator + h.validators = map[string]*queryValidationSpec{} + h.validators["GET /fragment/nodes"] = QueryValidationSpecRequired("slice").Optional("index") +} + +func (h *Handler) queryArgValidator(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + key := fmt.Sprintf("%s %s", r.Method, r.URL.Path) + if validator, ok := h.validators[key]; ok { + q := r.URL.Query() + err := validator.validate(q) + if err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + } + next.ServeHTTP(w, r) + }) } // NewRouter creates a Gorilla Mux http router. @@ -160,6 +171,8 @@ 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") + router.Use(handler.queryArgValidator) + return router } @@ -1356,12 +1369,6 @@ func (h *Handler) handleGetExportCSV(w http.ResponseWriter, r *http.Request) { // handleGetFragmentNodes handles /fragment/nodes requests. func (h *Handler) handleGetFragmentNodes(w http.ResponseWriter, r *http.Request) { q := r.URL.Query() - validator := h.validators["GetFragmentNodes"] - err := validator.validate(q) - if err != nil { - http.Error(w, err.Error(), http.StatusBadRequest) - return - } index := q.Get("index") // Read slice parameter. @@ -2083,23 +2090,23 @@ type queryValidationSpec struct { args map[string]bool } -func NewQueryValidationSpec() queryValidationSpec { - return queryValidationSpec{ - args: map[string]bool{}, +func QueryValidationSpecRequired(requiredArgs ...string) *queryValidationSpec { + args := map[string]bool{} + for _, arg := range requiredArgs { + args[arg] = true + } + + return &queryValidationSpec{ + required: requiredArgs, + args: args, } } -func (s *queryValidationSpec) Required(args ...string) { - s.required = args - for _, arg := range args { - s.args[arg] = true - } -} - -func (s *queryValidationSpec) Optional(args ...string) { +func (s *queryValidationSpec) Optional(args ...string) *queryValidationSpec { for _, arg := range args { s.args[arg] = true } + return s } func (s queryValidationSpec) validate(query url.Values) error { From 0417439afaf2f5d3554dd318f027bc804afb750c Mon Sep 17 00:00:00 2001 From: Yuce Tekol Date: Wed, 21 Feb 2018 18:03:05 +0300 Subject: [PATCH 3/4] Validate all relevant endpoints; Changed how validation key is composed --- handler.go | 46 +++++++++++++++++++++++++++++----------------- handler_test.go | 2 +- 2 files changed, 30 insertions(+), 18 deletions(-) diff --git a/handler.go b/handler.go index 895141135..024c222a7 100644 --- a/handler.go +++ b/handler.go @@ -89,6 +89,10 @@ var externalPrefixFlag = map[string]bool{ "version": true, } +type errorResponse struct { + Error string `json:"error"` +} + // NewHandler returns a new instance of Handler with a default logger. func NewHandler() *Handler { handler := &Handler{ @@ -101,17 +105,31 @@ func NewHandler() *Handler { func (h *Handler) populateValidators() { h.validators = map[string]*queryValidationSpec{} - h.validators["GET /fragment/nodes"] = QueryValidationSpecRequired("slice").Optional("index") + h.validators["GetFragmentNodes"] = QueryValidationSpecRequired("slice").Optional("index") + h.validators["GetSliceMax"] = QueryValidationSpecRequired().Optional("inverse") + h.validators["PostQuery"] = QueryValidationSpecRequired().Optional("slices", "columnAttrs", "excludeAttrs", "excludeBits") + h.validators["GetExport"] = QueryValidationSpecRequired("index", "frame", "view", "slice") + h.validators["GetFragmentData"] = QueryValidationSpecRequired("index", "frame", "view", "slice") + h.validators["PostFragmentData"] = QueryValidationSpecRequired("index", "frame", "view", "slice") + h.validators["GetFragmentBlocks"] = QueryValidationSpecRequired("index", "frame", "view", "slice") + h.validators["PostFrameRestore"] = QueryValidationSpecRequired("host") } func (h *Handler) queryArgValidator(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - key := fmt.Sprintf("%s %s", r.Method, r.URL.Path) + key := mux.CurrentRoute(r).GetName() if validator, ok := h.validators[key]; ok { q := r.URL.Query() err := validator.validate(q) if err != nil { - http.Error(w, err.Error(), http.StatusBadRequest) + // TODO: Return the response depending on the Accept header + response := errorResponse{Error: err.Error()} + body, err := json.Marshal(response) + if err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + http.Error(w, string(body), http.StatusBadRequest) return } } @@ -126,12 +144,12 @@ func NewRouter(handler *Handler) *mux.Router { router.HandleFunc("/assets/{file}", handler.handleWebUI).Methods("GET") router.PathPrefix("/debug/pprof/").Handler(http.DefaultServeMux).Methods("GET") router.HandleFunc("/debug/vars", handler.handleExpvar).Methods("GET") - router.HandleFunc("/export", handler.handleGetExport).Methods("GET") + router.HandleFunc("/export", handler.handleGetExport).Methods("GET").Name("GetExport") 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("/fragment/blocks", handler.handleGetFragmentBlocks).Methods("GET").Name("GetFragmentBlocks") + router.HandleFunc("/fragment/data", handler.handleGetFragmentData).Methods("GET").Name("GetFragmentData") + router.HandleFunc("/fragment/data", handler.handlePostFragmentData).Methods("POST").Name("PostFragmentData") + router.HandleFunc("/fragment/nodes", handler.handleGetFragmentNodes).Methods("GET").Name("GetFragmentNodes") router.HandleFunc("/import", handler.handlePostImport).Methods("POST") router.HandleFunc("/import-value", handler.handlePostImportValue).Methods("POST") router.HandleFunc("/index", handler.handleGetIndexes).Methods("GET") @@ -143,7 +161,7 @@ func NewRouter(handler *Handler) *mux.Router { router.HandleFunc("/index/{index}/frame/{frame}", handler.handlePostFrame).Methods("POST") router.HandleFunc("/index/{index}/frame/{frame}", handler.handleDeleteFrame).Methods("DELETE") router.HandleFunc("/index/{index}/frame/{frame}/attr/diff", handler.handlePostFrameAttrDiff).Methods("POST") - router.HandleFunc("/index/{index}/frame/{frame}/restore", handler.handlePostFrameRestore).Methods("POST") + router.HandleFunc("/index/{index}/frame/{frame}/restore", handler.handlePostFrameRestore).Methods("POST").Name("PostFrameRestore") router.HandleFunc("/index/{index}/frame/{frame}/time-quantum", handler.handlePatchFrameTimeQuantum).Methods("PATCH") router.HandleFunc("/index/{index}/frame/{frame}/field/{field}", handler.handlePostFrameField).Methods("POST") router.HandleFunc("/index/{index}/frame/{frame}/fields", handler.handleGetFrameFields).Methods("GET") @@ -154,11 +172,11 @@ func NewRouter(handler *Handler) *mux.Router { router.HandleFunc("/index/{index}/input-definition/{input-definition}", handler.handleGetInputDefinition).Methods("GET") router.HandleFunc("/index/{index}/input-definition/{input-definition}", handler.handlePostInputDefinition).Methods("POST") 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}/query", handler.handlePostQuery).Methods("POST").Name("PostQuery") 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.handleGetSliceMax).Methods("GET") + router.HandleFunc("/slices/max", handler.handleGetSliceMax).Methods("GET").Name("GetSliceMax") router.HandleFunc("/status", handler.handleGetStatus).Methods("GET") router.HandleFunc("/version", handler.handleGetVersion).Methods("GET") router.HandleFunc("/recalculate-caches", handler.handleRecalculateCaches).Methods("POST") @@ -1097,12 +1115,6 @@ func (h *Handler) readProtobufQueryRequest(r *http.Request) (*QueryRequest, erro // readURLQueryRequest parses query parameters from URL parameters from r. func (h *Handler) readURLQueryRequest(r *http.Request) (*QueryRequest, error) { q := r.URL.Query() - validQuery := validOptions(QueryRequest{}) - for key := range q { - if _, ok := validQuery[key]; !ok { - return nil, errors.New("invalid query params") - } - } // Parse query string. buf, err := ioutil.ReadAll(r.Body) diff --git a/handler_test.go b/handler_test.go index 636665022..f1c533691 100644 --- a/handler_test.go +++ b/handler_test.go @@ -306,7 +306,7 @@ func TestHandler_Query_Params_Err(t *testing.T) { test.NewHandler().ServeHTTP(w, test.MustNewHTTPRequest("POST", "/index/idx0/query?slices=0,1&db=sample", strings.NewReader("Bitmap(id=100)"))) if w.Code != http.StatusBadRequest { t.Fatalf("unexpected status code: %d", w.Code) - } else if body := w.Body.String(); body != `{"error":"invalid query params"}`+"\n" { + } else if body := w.Body.String(); body != `{"error":"db is not a valid argument"}`+"\n" { t.Fatalf("unexpected body: %q", body) } From 25e244ebdfd90b434fa04eec894e7164b6c60267 Mon Sep 17 00:00:00 2001 From: Yuce Tekol Date: Thu, 22 Feb 2018 17:33:41 +0300 Subject: [PATCH 4/4] Updates --- handler.go | 12 +++++------- 1 file changed, 5 insertions(+), 7 deletions(-) diff --git a/handler.go b/handler.go index 024c222a7..34075b3d4 100644 --- a/handler.go +++ b/handler.go @@ -119,9 +119,7 @@ func (h *Handler) queryArgValidator(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { key := mux.CurrentRoute(r).GetName() if validator, ok := h.validators[key]; ok { - q := r.URL.Query() - err := validator.validate(q) - if err != nil { + if err := validator.validate(r.URL.Query()); err != nil { // TODO: Return the response depending on the Accept header response := errorResponse{Error: err.Error()} body, err := json.Marshal(response) @@ -2099,13 +2097,13 @@ type defaultClusterMessageResponse struct{} type queryValidationSpec struct { required []string - args map[string]bool + args map[string]struct{} } func QueryValidationSpecRequired(requiredArgs ...string) *queryValidationSpec { - args := map[string]bool{} + args := map[string]struct{}{} for _, arg := range requiredArgs { - args[arg] = true + args[arg] = struct{}{} } return &queryValidationSpec{ @@ -2116,7 +2114,7 @@ func QueryValidationSpecRequired(requiredArgs ...string) *queryValidationSpec { func (s *queryValidationSpec) Optional(args ...string) *queryValidationSpec { for _, arg := range args { - s.args[arg] = true + s.args[arg] = struct{}{} } return s }