From 8e6041d0fe36fb9d3e30812f425d76efb062d8f9 Mon Sep 17 00:00:00 2001 From: Yuce Tekol Date: Mon, 19 Feb 2018 13:55:19 +0300 Subject: [PATCH] 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.