From 93542dbbecfa2230ecef486d2541a722cf624cd8 Mon Sep 17 00:00:00 2001 From: NIkhil Bhatia Date: Tue, 23 Oct 2018 10:27:26 -0700 Subject: [PATCH] Post query (#1) * add POST method for /query Signed-off-by: Nikhil Bhatia --- server/server.go | 47 ++++++++++++++++++++++++++- server/server_test.go | 75 +++++++++++++++++++++++++++++-------------- server/types/types.go | 5 +++ 3 files changed, 102 insertions(+), 25 deletions(-) diff --git a/server/server.go b/server/server.go index 4bf538c014..9e5a7b8c67 100644 --- a/server/server.go +++ b/server/server.go @@ -385,6 +385,7 @@ func (s *Server) initRouter() { s.registerHandler(router, 1, "/policies/{path:.+}", http.MethodGet, promhttp.InstrumentHandlerDuration(v1PoliciesDur, http.HandlerFunc(s.v1PoliciesGet))) s.registerHandler(router, 1, "/policies/{path:.+}", http.MethodPut, promhttp.InstrumentHandlerDuration(v1PoliciesDur, http.HandlerFunc(s.v1PoliciesPut))) s.registerHandler(router, 1, "/query", http.MethodGet, promhttp.InstrumentHandlerDuration(v1QueryDur, http.HandlerFunc(s.v1QueryGet))) + s.registerHandler(router, 1, "/query", http.MethodPost, promhttp.InstrumentHandlerDuration(v1QueryDur, http.HandlerFunc(s.v1QueryPost))) s.registerHandler(router, 1, "/compile", http.MethodPost, promhttp.InstrumentHandlerDuration(v1CompileDur, http.HandlerFunc(s.v1CompilePost))) router.HandleFunc("/", promhttp.InstrumentHandlerDuration(indexDur, http.HandlerFunc(s.unversionedPost))).Methods(http.MethodPost) router.HandleFunc("/", promhttp.InstrumentHandlerDuration(indexDur, http.HandlerFunc(s.indexGet))).Methods(http.MethodGet) @@ -410,7 +411,7 @@ func (s *Server) initRouter() { router.HandleFunc("/v1/query/{path:.*}", promhttp.InstrumentHandlerDuration(catchAllDur, http.HandlerFunc(writer.HTTPStatus(405)))).Methods(http.MethodHead, http.MethodConnect, http.MethodDelete, http.MethodOptions, http.MethodTrace, http.MethodPost, http.MethodPut, http.MethodPatch) router.HandleFunc("/v1/query", promhttp.InstrumentHandlerDuration(catchAllDur, http.HandlerFunc(writer.HTTPStatus(405)))).Methods(http.MethodHead, - http.MethodConnect, http.MethodDelete, http.MethodOptions, http.MethodTrace, http.MethodPost, http.MethodPut, http.MethodPatch) + http.MethodConnect, http.MethodDelete, http.MethodOptions, http.MethodTrace, http.MethodPut, http.MethodPatch) s.Handler = router } @@ -1379,6 +1380,50 @@ func (s *Server) v1QueryGet(w http.ResponseWriter, r *http.Request) { writer.JSON(w, 200, results, pretty) } +func (s *Server) v1QueryPost(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + + var request types.QueryRequestV1 + err := util.NewJSONDecoder(r.Body).Decode(&request) + if err != nil { + writer.Error(w, http.StatusBadRequest, types.NewErrorV1(types.CodeInvalidParameter, "error(s) occurred while decoding request: %v", err.Error())) + return + } + qStr := request.Query + unsafeBuiltins, err := validateQuery(qStr) + if err != nil { + writer.ErrorAuto(w, err) + return + } else if len(unsafeBuiltins) > 0 { + writer.Error(w, http.StatusBadRequest, types.NewErrorV1(types.CodeInvalidParameter, "unsafe built-in function calls in query: %v", strings.Join(unsafeBuiltins, ","))) + return + } + + watch := getWatch(r.URL.Query()[types.ParamWatchV1]) + if watch { + s.watchQuery(qStr, w, r, false) + return + } + + pretty := getBoolParam(r.URL, types.ParamPrettyV1, true) + explainMode := getExplain(r.URL.Query()["explain"], types.ExplainOffV1) + includeMetrics := getBoolParam(r.URL, types.ParamMetricsV1, true) + includeInstrumentation := getBoolParam(r.URL, types.ParamInstrumentV1, true) + + results, err := s.execQuery(ctx, r, qStr, nil, explainMode, includeMetrics, includeInstrumentation, pretty) + if err != nil { + switch err := err.(type) { + case ast.Errors: + writer.Error(w, http.StatusBadRequest, types.NewErrorV1(types.CodeInvalidParameter, types.MsgCompileQueryError).WithASTErrors(err)) + default: + writer.ErrorAuto(w, err) + } + return + } + + writer.JSON(w, 200, results, pretty) +} + func (s *Server) watchQuery(query string, w http.ResponseWriter, r *http.Request, data bool) { pretty := getBoolParam(r.URL, types.ParamPrettyV1, true) explainMode := getExplain(r.URL.Query()["explain"], types.ExplainOffV1) diff --git a/server/server_test.go b/server/server_test.go index bef055baba..1479754dc2 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -140,7 +140,6 @@ func Test405StatusCodev1(t *testing.T) { {http.MethodDelete, "/query", "", 405, ""}, {http.MethodOptions, "/query", "", 405, ""}, {http.MethodTrace, "/query", "", 405, ""}, - {http.MethodPost, "/query", "", 405, ""}, {http.MethodPut, "/query", "", 405, ""}, {http.MethodPatch, "/query", "", 405, ""}, }}, @@ -1527,9 +1526,30 @@ func TestPoliciesPathSlashes(t *testing.T) { } } -func TestQueryWatchBasic(t *testing.T) { +func TestQueryPostBasic(t *testing.T) { f := newFixture(t) + f.server, _ = New(). + WithAddresses([]string{":8182"}). + WithStore(f.server.store). + WithManager(f.server.manager). + WithDiagnosticsBuffer(NewBoundedBuffer(8)). + Init(context.Background()) + setup := []tr{ + {http.MethodPost, "/query", `{"query": "a=data.k.x with data.k as {\"x\" : 7}"}`, 200, `{"result":[{"a":7}]}`}, + } + + for _, tr := range setup { + req := newReqV1(tr.method, tr.path, tr.body) + req.RemoteAddr = "testaddr" + + if err := f.executeRequest(req, tr.code, tr.resp); err != nil { + t.Fatal(err) + } + } +} + +func TestQueryWatchBasic(t *testing.T) { // Test basic watch results. exp := strings.Join([]string{ "HTTP/1.1 200 OK\nContent-Type: application/json\nTransfer-Encoding: chunked\n\n10", @@ -1546,32 +1566,39 @@ func TestQueryWatchBasic(t *testing.T) { `, ``, }, "\r\n") - recorder := newMockConn() - get := newReqV1(http.MethodGet, `/query?q=a=data.x&watch`, "") - go f.server.Handler.ServeHTTP(recorder, get) - <-recorder.hijacked - <-recorder.write - - tests := []trw{ - {tr{http.MethodPut, "/data/x", `{"a":1,"b":2}`, 204, ""}, recorder.write}, - {tr{http.MethodPut, "/data/x", `"foo"`, 204, ""}, recorder.write}, - {tr{http.MethodPut, "/data/x", `7`, 204, ""}, recorder.write}, + requests := []*http.Request{ + newReqV1(http.MethodGet, `/query?q=a=data.x&watch`, ""), + newReqV1(http.MethodPost, `/query?&watch`, `{"query": "a=data.x"}`), } - for _, test := range tests { - tr := test.tr - if err := f.v1(tr.method, tr.path, tr.body, tr.code, tr.resp); err != nil { - t.Fatal(err) - } - if test.wait != nil { - <-test.wait - } - } - recorder.Close() + for _, get := range requests { + f := newFixture(t) + recorder := newMockConn() + go f.server.Handler.ServeHTTP(recorder, get) + <-recorder.hijacked + <-recorder.write - if result := recorder.buf.String(); result != exp { - t.Fatalf("Expected stream to equal %s, got %s", exp, result) + tests := []trw{ + {tr{http.MethodPut, "/data/x", `{"a":1,"b":2}`, 204, ""}, recorder.write}, + {tr{http.MethodPut, "/data/x", `"foo"`, 204, ""}, recorder.write}, + {tr{http.MethodPut, "/data/x", `7`, 204, ""}, recorder.write}, + } + + for _, test := range tests { + tr := test.tr + if err := f.v1(tr.method, tr.path, tr.body, tr.code, tr.resp); err != nil { + t.Fatal(err) + } + if test.wait != nil { + <-test.wait + } + } + recorder.Close() + + if result := recorder.buf.String(); result != exp { + t.Fatalf("Expected stream to equal %s, got %s", exp, result) + } } } diff --git a/server/types/types.go b/server/types/types.go index f1d16d99a9..fcdd37c0a0 100644 --- a/server/types/types.go +++ b/server/types/types.go @@ -376,6 +376,11 @@ type PartialEvaluationResultV1 struct { Support []*ast.Module `json:"support,omitempty"` } +// QueryRequestV1 models the request message for Query API operations. +type QueryRequestV1 struct { + Query string `json:"query"` +} + const ( // ParamQueryV1 defines the name of the HTTP URL parameter that specifies // values for the request query.