Added query arg validator middleware; updated Gorilla/mux

This commit is contained in:
Yuce Tekol 2018-02-20 15:38:58 +03:00
parent 8e6041d0fe
commit c1f9dd4e59
No known key found for this signature in database
GPG key ID: CB59E46D2FB90573
2 changed files with 33 additions and 26 deletions

4
Gopkg.lock generated
View file

@ -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"

View file

@ -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 {