From c1f9dd4e593e0229c196fd545b09d017abe7f43a Mon Sep 17 00:00:00 2001 From: Yuce Tekol Date: Tue, 20 Feb 2018 15:38:58 +0300 Subject: [PATCH] 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 {