// Copyright 2016 The OPA Authors. All rights reserved. // Use of this source code is governed by an Apache2 // license that can be found in the LICENSE file. package server import ( "context" "encoding/json" "fmt" "net/http" "net/http/httptest" "reflect" "sort" "strings" "testing" "time" "github.com/pkg/errors" "github.com/open-policy-agent/opa/ast" "github.com/open-policy-agent/opa/bundle" "github.com/open-policy-agent/opa/metrics" "github.com/open-policy-agent/opa/plugins" pluginBundle "github.com/open-policy-agent/opa/plugins/bundle" "github.com/open-policy-agent/opa/server/identifier" "github.com/open-policy-agent/opa/server/types" "github.com/open-policy-agent/opa/storage" "github.com/open-policy-agent/opa/storage/inmem" "github.com/open-policy-agent/opa/util" "github.com/open-policy-agent/opa/util/test" "github.com/open-policy-agent/opa/version" ) type tr struct { method string path string body string code int resp string } type trw struct { tr tr wait chan struct{} } func TestUnversionedGetHealth(t *testing.T) { f := newFixture(t) req := newReqUnversioned(http.MethodGet, "/health", "") validateDiagnosticRequest(t, f, req, 200, `{}`) } func TestUnversionedGetHealthBundleNoBundleSet(t *testing.T) { f := newFixture(t) req := newReqUnversioned(http.MethodGet, "/health?bundles=true", "") validateDiagnosticRequest(t, f, req, 200, `{}`) } func TestUnversionedGetHealthCheckOnlyBundlePlugin(t *testing.T) { f := newFixture(t) // Initialize the server as if a bundle plugin was // configured on the manager. f.server.manager.UpdatePluginStatus("bundle", &plugins.Status{State: plugins.StateNotReady}) // The bundle hasn't been activated yet, expect the health check to fail req := newReqUnversioned(http.MethodGet, "/health?bundles=true", "") validateDiagnosticRequest(t, f, req, 500, `{}`) // Set the bundle to be activated. f.server.manager.UpdatePluginStatus("bundle", &plugins.Status{State: plugins.StateOK}) // The heath check should now respond as healthy req = newReqUnversioned(http.MethodGet, "/health?bundles=true", "") validateDiagnosticRequest(t, f, req, 200, `{}`) } func TestUnversionedGetHealthCheckDiscoveryWithBundle(t *testing.T) { f := newFixture(t) // Initialize the server as if a discovery bundle is configured f.server.manager.UpdatePluginStatus("discovery", &plugins.Status{State: plugins.StateNotReady}) // The discovery bundle hasn't been activated yet, expect the health check to fail req := newReqUnversioned(http.MethodGet, "/health?bundles=true", "") validateDiagnosticRequest(t, f, req, 500, `{}`) // Set the bundle to be not ready (plugin configured and created, but hasn't activated all bundles yet). f.server.manager.UpdatePluginStatus("discovery", &plugins.Status{State: plugins.StateOK}) f.server.manager.UpdatePluginStatus("bundle", &plugins.Status{State: plugins.StateNotReady}) // The discovery bundle is OK, but the newly configured bundle hasn't been activated yet, expect the health check to fail req = newReqUnversioned(http.MethodGet, "/health?bundles=true", "") validateDiagnosticRequest(t, f, req, 500, `{}`) // Set the bundle to be activated. f.server.manager.UpdatePluginStatus("bundle", &plugins.Status{State: plugins.StateOK}) // The heath check should now respond as healthy req = newReqUnversioned(http.MethodGet, "/health?bundles=true", "") validateDiagnosticRequest(t, f, req, 200, `{}`) } func TestUnversionedGetHealthCheckBundleActivationSingleLegacy(t *testing.T) { // Initialize the server as if there is no bundle plugin f := newFixture(t) ctx := context.Background() // The server doesn't know about any bundles, so return a healthy status req := newReqUnversioned(http.MethodGet, "/health?bundle=true", "") validateDiagnosticRequest(t, f, req, 200, `{}`) err := storage.Txn(ctx, f.server.store, storage.WriteParams, func(txn storage.Transaction) error { return bundle.LegacyWriteManifestToStore(ctx, f.server.store, txn, bundle.Manifest{ Revision: "a", }) }) if err != nil { t.Fatalf("Unexpected error: %s", err) } // The heath check still respond as healthy with a legacy bundle found in storage req = newReqUnversioned(http.MethodGet, "/health?bundle=true", "") validateDiagnosticRequest(t, f, req, 200, `{}`) } func TestBundlesReady(t *testing.T) { cases := []struct { note string status map[string]*plugins.Status ready bool }{ { note: "nil status", status: nil, ready: true, }, { note: "empty status", status: map[string]*plugins.Status{}, ready: true, }, { note: "discovery not ready - bundle missing", status: map[string]*plugins.Status{ "discovery": {State: plugins.StateNotReady}, }, ready: false, }, { note: "discovery ok - bundle missing", status: map[string]*plugins.Status{ "discovery": {State: plugins.StateOK}, }, ready: true, // bundles aren't enabled, only discovery plugin configured }, { note: "discovery missing - bundle not ready", status: map[string]*plugins.Status{ "bundle": {State: plugins.StateNotReady}, }, ready: false, }, { note: "discovery missing - bundle ok", status: map[string]*plugins.Status{ "bundle": {State: plugins.StateOK}, }, ready: true, // discovery isn't enabled, only bundle plugin configured }, { note: "discovery not ready - bundle not ready", status: map[string]*plugins.Status{ "discovery": {State: plugins.StateNotReady}, "bundle": {State: plugins.StateNotReady}, }, ready: false, }, { note: "discovery ok - bundle not ready", status: map[string]*plugins.Status{ "discovery": {State: plugins.StateOK}, "bundle": {State: plugins.StateNotReady}, }, ready: false, }, { note: "discovery not ready - bundle ok", status: map[string]*plugins.Status{ "discovery": {State: plugins.StateNotReady}, "bundle": {State: plugins.StateOK}, }, ready: false, }, { note: "discovery ok - bundle ok", status: map[string]*plugins.Status{ "discovery": {State: plugins.StateOK}, "bundle": {State: plugins.StateOK}, }, ready: true, }, } for _, tc := range cases { t.Run(tc.note, func(t *testing.T) { f := newFixture(t) actual := f.server.bundlesReady(tc.status) if actual != tc.ready { t.Errorf("Expected %t got %t", tc.ready, actual) } }) } } func TestUnversionedGetHealthCheckDiscoveryWithPlugins(t *testing.T) { // Use the same server through the cases, the status updates apply incrementally to it. f := newFixture(t) cases := []struct { note string statusUpdates map[string]*plugins.Status exp int }{ { note: "no plugins configured", statusUpdates: nil, exp: 200, }, { note: "one plugin configured - not ready", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateNotReady}, }, exp: 500, }, { note: "one plugin configured - ready", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateOK}, }, exp: 200, }, { note: "one plugin configured - error state", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateErr}, }, exp: 500, }, { note: "one plugin configured - recovered from error", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateOK}, }, exp: 200, }, { note: "add second plugin - not ready", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateOK}, "p2": {State: plugins.StateNotReady}, }, exp: 500, }, { note: "add third plugin - not ready", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateOK}, "p2": {State: plugins.StateNotReady}, "p3": {State: plugins.StateNotReady}, }, exp: 500, }, { note: "mixed states - not ready", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateOK}, "p2": {State: plugins.StateErr}, "p3": {State: plugins.StateNotReady}, }, exp: 500, }, { note: "mixed states - still not ready", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateOK}, "p2": {State: plugins.StateErr}, "p3": {State: plugins.StateOK}, }, exp: 500, }, { note: "all plugins ready", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateOK}, "p2": {State: plugins.StateOK}, "p3": {State: plugins.StateOK}, }, exp: 200, }, { note: "one plugins fails", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateErr}, "p2": {State: plugins.StateOK}, "p3": {State: plugins.StateOK}, }, exp: 500, }, { note: "all plugins ready - recovery", statusUpdates: map[string]*plugins.Status{ "p1": {State: plugins.StateOK}, "p2": {State: plugins.StateOK}, "p3": {State: plugins.StateOK}, }, exp: 200, }, { note: "nil plugin status", statusUpdates: map[string]*plugins.Status{ "p1": nil, }, exp: 200, }, } for _, tc := range cases { t.Run(tc.note, func(t *testing.T) { for name, status := range tc.statusUpdates { f.server.manager.UpdatePluginStatus(name, status) } req := newReqUnversioned(http.MethodGet, "/health?plugins", "") validateDiagnosticRequest(t, f, req, tc.exp, `{}`) }) } } func TestUnversionedGetHealthCheckBundleAndPlugins(t *testing.T) { cases := []struct { note string statuses map[string]*plugins.Status exp int }{ { note: "no plugins configured", statuses: nil, exp: 200, }, { note: "only bundle plugin configured - not ready", statuses: map[string]*plugins.Status{ "bundle": {State: plugins.StateNotReady}, }, exp: 500, }, { note: "only bundle plugin configured - ok", statuses: map[string]*plugins.Status{ "bundle": {State: plugins.StateOK}, }, exp: 200, }, { note: "only custom plugin configured - not ready", statuses: map[string]*plugins.Status{ "p1": {State: plugins.StateNotReady}, }, exp: 500, }, { note: "only custom plugin configured - ok", statuses: map[string]*plugins.Status{ "p1": {State: plugins.StateOK}, }, exp: 200, }, { note: "both configured - bundle not ready", statuses: map[string]*plugins.Status{ "bundle": {State: plugins.StateNotReady}, "p1": {State: plugins.StateOK}, }, exp: 500, }, { note: "both configured - custom plugin not ready", statuses: map[string]*plugins.Status{ "bundle": {State: plugins.StateOK}, "p1": {State: plugins.StateNotReady}, }, exp: 500, }, { note: "both configured - both ready", statuses: map[string]*plugins.Status{ "bundle": {State: plugins.StateOK}, "p1": {State: plugins.StateOK}, }, exp: 200, }, } for _, tc := range cases { t.Run(tc.note, func(t *testing.T) { f := newFixture(t) for name, status := range tc.statuses { f.server.manager.UpdatePluginStatus(name, status) } req := newReqUnversioned(http.MethodGet, "/health?plugins&bundles", "") validateDiagnosticRequest(t, f, req, tc.exp, `{}`) }) } } func TestDataV0(t *testing.T) { testMod1 := `package test p = "hello" q = { "foo": [1,2,3,4] } { input.flag = true } ` f := newFixture(t) if err := f.v1(http.MethodPut, "/policies/test", testMod1, 200, ""); err != nil { t.Fatalf("Unexpected error while creating policy: %v", err) } if err := f.v0(http.MethodPost, "/data/test/p", "", 200, `"hello"`); err != nil { t.Fatalf("Expected response hello but got: %v", err) } if err := f.v0(http.MethodPost, "/data/test/q/foo", `{"flag": true}`, 200, `[1,2,3,4]`); err != nil { t.Fatalf("Expected response [1,2,3,4] but got: %v", err) } req := newReqV0(http.MethodPost, "/data/test/q", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) if f.recorder.Code != 404 { t.Fatalf("Expected HTTP 404 but got: %v", f.recorder) } var resp types.ErrorV1 if err := util.NewJSONDecoder(f.recorder.Body).Decode(&resp); err != nil { t.Fatalf("Unexpected error while deserializing response: %v", err) } if resp.Code != types.CodeUndefinedDocument { t.Fatalf("Expected undefiend code but got: %v", resp) } } // Tests that the responses for (theoretically) valid resources but with forbidden methods return the proper status code func Test405StatusCodev1(t *testing.T) { tests := []struct { note string reqs []tr }{ {"v1 data one level 405", []tr{ {http.MethodHead, "/data/lvl1", "", 405, ""}, {http.MethodConnect, "/data/lvl1", "", 405, ""}, {http.MethodOptions, "/data/lvl1", "", 405, ""}, {http.MethodTrace, "/data/lvl1", "", 405, ""}, }}, {"v1 data 405", []tr{ {http.MethodHead, "/data", "", 405, ""}, {http.MethodConnect, "/data", "", 405, ""}, {http.MethodOptions, "/data", "", 405, ""}, {http.MethodTrace, "/data", "", 405, ""}, {http.MethodDelete, "/data", "", 405, ""}, }}, {"v1 policies 405", []tr{ {http.MethodHead, "/policies", "", 405, ""}, {http.MethodConnect, "/policies", "", 405, ""}, {http.MethodDelete, "/policies", "", 405, ""}, {http.MethodOptions, "/policies", "", 405, ""}, {http.MethodTrace, "/policies", "", 405, ""}, {http.MethodPost, "/policies", "", 405, ""}, {http.MethodPut, "/policies", "", 405, ""}, {http.MethodPatch, "/policies", "", 405, ""}, }}, {"v1 policies one level 405", []tr{ {http.MethodHead, "/policies/lvl1", "", 405, ""}, {http.MethodConnect, "/policies/lvl1", "", 405, ""}, {http.MethodOptions, "/policies/lvl1", "", 405, ""}, {http.MethodTrace, "/policies/lvl1", "", 405, ""}, {http.MethodPost, "/policies/lvl1", "", 405, ""}, }}, {"v1 query one level 405", []tr{ {http.MethodHead, "/query/lvl1", "", 405, ""}, {http.MethodConnect, "/query/lvl1", "", 405, ""}, {http.MethodDelete, "/query/lvl1", "", 405, ""}, {http.MethodOptions, "/query/lvl1", "", 405, ""}, {http.MethodTrace, "/query/lvl1", "", 405, ""}, {http.MethodPost, "/query/lvl1", "", 405, ""}, {http.MethodPut, "/query/lvl1", "", 405, ""}, {http.MethodPatch, "/query/lvl1", "", 405, ""}, }}, {"v1 query 405", []tr{ {http.MethodHead, "/query", "", 405, ""}, {http.MethodConnect, "/query", "", 405, ""}, {http.MethodDelete, "/query", "", 405, ""}, {http.MethodOptions, "/query", "", 405, ""}, {http.MethodTrace, "/query", "", 405, ""}, {http.MethodPut, "/query", "", 405, ""}, {http.MethodPatch, "/query", "", 405, ""}, }}, } for _, tc := range tests { test.Subtest(t, tc.note, func(t *testing.T) { executeRequests(t, tc.reqs) }) } } // Tests that the responses for (theoretically) valid resources but with forbidden methods return the proper status code func Test405StatusCodev0(t *testing.T) { tests := []struct { note string reqs []tr }{ {"v0 data one levels 405", []tr{ {http.MethodHead, "/data/lvl2", "", 405, ""}, {http.MethodConnect, "/data/lvl2", "", 405, ""}, {http.MethodDelete, "/data/lvl2", "", 405, ""}, {http.MethodOptions, "/data/lvl2", "", 405, ""}, {http.MethodTrace, "/data/lvl2", "", 405, ""}, {http.MethodGet, "/data/lvl2", "", 405, ""}, {http.MethodPatch, "/data/lvl2", "", 405, ""}, {http.MethodPut, "/data/lvl2", "", 405, ""}, }}, {"v0 data 405", []tr{ {http.MethodHead, "/data", "", 405, ""}, {http.MethodConnect, "/data", "", 405, ""}, {http.MethodDelete, "/data", "", 405, ""}, {http.MethodOptions, "/data", "", 405, ""}, {http.MethodTrace, "/data", "", 405, ""}, {http.MethodGet, "/data", "", 405, ""}, {http.MethodPatch, "/data", "", 405, ""}, {http.MethodPut, "/data", "", 405, ""}, }}, } for _, tc := range tests { test.Subtest(t, tc.note, func(t *testing.T) { executeRequestsv0(t, tc.reqs) }) } } func TestCompileV1(t *testing.T) { mod := `package test p { input.x = 1 } q { data.a[i] = input.x } default r = true r { input.x = 1 } ` expQuery := func(s string) string { return fmt.Sprintf(`{"result": {"queries": [%v]}}`, string(util.MustMarshalJSON(ast.MustParseBody(s)))) } expQueryAndSupport := func(q string, m string) string { return fmt.Sprintf(`{"result": {"queries": [%v], "support": [%v]}}`, string(util.MustMarshalJSON(ast.MustParseBody(q))), string(util.MustMarshalJSON(ast.MustParseModule(m)))) } tests := []struct { note string trs []tr }{ { note: "basic", trs: []tr{ {http.MethodPut, "/policies/test", mod, 200, ""}, {http.MethodPost, "/compile", `{ "unknowns": ["input"], "query": "data.test.p = true" }`, 200, expQuery("input.x = 1")}, }, }, { note: "subtree", trs: []tr{ {http.MethodPost, "/compile", `{ "unknowns": ["input.x"], "input": {"y": 1}, "query": "input.x > input.y" }`, 200, expQuery("input.x > 1")}, }, }, { note: "data", trs: []tr{ {http.MethodPut, "/policies/test", mod, 200, ""}, {http.MethodPost, "/compile", `{ "unknowns": ["data.a"], "input": { "x": 1 }, "query": "data.test.q = true" }`, 200, expQuery("1 = data.a[i1]")}, }, }, { note: "escaped string", trs: []tr{ {http.MethodPost, "/compile", `{ "query": "input[\"x\"] = 1" }`, 200, expQuery("input.x = 1")}, }, }, { note: "support", trs: []tr{ {http.MethodPut, "/policies/test", mod, 200, ""}, {http.MethodPost, "/compile", `{ "query": "data.test.r = true" }`, 200, expQueryAndSupport( `data.partial.test.r = true`, `package partial.test r { input.x = 1 } default r = true `)}, }, }, { note: "empty unknowns", trs: []tr{ {http.MethodPost, "/compile", `{"query": "input.x > 1", "unknowns": []}`, 200, `{"result": {}}`}, }, }, { note: "never defined", trs: []tr{ {http.MethodPost, "/compile", `{"query": "1 = 2"}`, 200, `{"result": {}}`}, }, }, { note: "always defined", trs: []tr{ {http.MethodPost, "/compile", `{"query": "1 = 1"}`, 200, `{"result": {"queries": [[]]}}`}, }, }, { note: "error: bad request", trs: []tr{{http.MethodPost, "/compile", `{"input": [{]}`, 400, ``}}, }, { note: "error: empty query", trs: []tr{{http.MethodPost, "/compile", `{}`, 400, ""}}, }, { note: "error: bad query", trs: []tr{{http.MethodPost, "/compile", `{"query": "x %!> 9"}`, 400, ""}}, }, { note: "error: bad unknown", trs: []tr{{http.MethodPost, "/compile", `{"unknowns": ["input."], "query": "true"}`, 400, ""}}, }, } for _, tc := range tests { test.Subtest(t, tc.note, func(t *testing.T) { executeRequests(t, tc.trs) }) } } func TestCompileV1Observability(t *testing.T) { f := newFixture(t) f.v1(http.MethodPut, "/policies/test", `package test p { input.x = 1 }`, 200, "") compileReq := newReqV1(http.MethodPost, "/compile?metrics&explain=full", `{ "query": "data.test.p = true" }`) f.reset() f.server.Handler.ServeHTTP(f.recorder, compileReq) var response types.CompileResponseV1 if err := json.NewDecoder(f.recorder.Body).Decode(&response); err != nil { t.Fatal(err) } if len(response.Explanation) == 0 { t.Fatal("Expected non-empty explanation") } if _, ok := response.Metrics["timer_rego_partial_eval_ns"]; !ok { t.Fatal("Expected partial evaluation latency") } } func TestCompileV1UnsafeBuiltin(t *testing.T) { f := newFixture(t) query := `{"query": "http.send({\"method\": \"get\", \"url\": \"foo.com\"}, x)"}` expResp := `{ "code": "invalid_parameter", "message": "error(s) occurred while compiling module(s)", "errors": [ { "code": "rego_type_error", "message": "unsafe built-in function calls in expression: http.send", "location": { "file": "", "row": 1, "col": 1 } } ] }` if err := f.v1(http.MethodPost, `/compile`, query, 400, expResp); err != nil { t.Fatalf("Expected bad request but got %v", f.recorder) } } func TestDataV1Redirection(t *testing.T) { f := newFixture(t) // Testing redirect at the root level if err := f.v1(http.MethodPut, "/data/", `{"foo": [1,2,3]}`, 301, ""); err != nil { t.Fatalf("Unexpected error from PUT: %v", err) } locHdr := f.recorder.Header().Get("Location") if strings.Compare(locHdr, "/v1/data") != 0 { t.Fatalf("Unexpected error Location header value: %v", locHdr) } RedirectedPath := strings.SplitAfter(locHdr, "/v1")[1] if err := f.v1(http.MethodPut, RedirectedPath, `{"foo": [1,2,3]}`, 204, ""); err != nil { t.Fatalf("Unexpected error from PUT: %v", err) } if err := f.v1(http.MethodGet, RedirectedPath, "", 200, `{"result": {"foo": [1,2,3]}}`); err != nil { t.Fatalf("Unexpected error from GET: %v", err) } // Now we test redirection a few levels down if err := f.v1(http.MethodPut, "/data/a/b/c/", `{"foo": [1,2,3]}`, 301, ""); err != nil { t.Fatalf("Unexpected error from PUT: %v", err) } locHdrLv := f.recorder.Header().Get("Location") if strings.Compare(locHdrLv, "/v1/data/a/b/c") != 0 { t.Fatalf("Unexpected error Location header value: %v", locHdrLv) } RedirectedPathLvl := strings.SplitAfter(locHdrLv, "/v1")[1] if err := f.v1(http.MethodPut, RedirectedPathLvl, `{"foo": [1,2,3]}`, 204, ""); err != nil { t.Fatalf("Unexpected error from PUT: %v", err) } if err := f.v1(http.MethodGet, RedirectedPathLvl, "", 200, `{"result": {"foo": [1,2,3]}}`); err != nil { t.Fatalf("Unexpected error from GET: %v", err) } } func TestDataV1(t *testing.T) { testMod1 := `package testmod import input.req1 import input.req2 as reqx import input.req3.attr1 p[x] { q[x]; not r[x] } q[x] { data.x.y[i] = x } r[x] { data.x.z[i] = x } g = true { req1.a[0] = 1; reqx.b[i] = 1 } h = true { attr1[i] > 1 } gt1 = true { req1 > 1 } arr = [1, 2, 3, 4] { true } undef = true { false }` testMod2 := `package testmod p = [1, 2, 3, 4] { true } q = {"a": 1, "b": 2} { true }` testMod4 := `package testmod p = true { true } p = false { true }` testMod5 := `package testmod.empty.mod` testMod6 := `package testmod.all.undefined p = true { false }` testMod7 := `package testmod default p = false p { q[x]; not r[x] } q[1] { input.x = 1 } q[2] { input.y = 2 } r[1] { input.z = 3 }` testMod7Modified := `package testmod default p = false p { q[x]; not r[x] } q[1] { input.x = 1 } q[2] { input.y = 2 } r[1] { input.z = 3 } r[2] { input.z = 3 }` testMod8 := `package testmod p { data.x = 1 }` tests := []struct { note string reqs []tr }{ {"add root", []tr{ {http.MethodPatch, "/data/x", `[{"op": "add", "path": "/", "value": {"a": 1}}]`, 204, ""}, {http.MethodGet, "/data/x/a", "", 200, `{"result": 1}`}, }}, {"append array", []tr{ {http.MethodPatch, "/data/x", `[{"op": "add", "path": "/", "value": []}]`, 204, ""}, {http.MethodPatch, "/data/x", `[{"op": "add", "path": "-", "value": {"a": 1}}]`, 204, ""}, {http.MethodPatch, "/data/x", `[{"op": "add", "path": "-", "value": {"a": 2}}]`, 204, ""}, {http.MethodGet, "/data/x/0/a", "", 200, `{"result": 1}`}, {http.MethodGet, "/data/x/1/a", "", 200, `{"result": 2}`}, }}, {"append array one-shot", []tr{ {http.MethodPatch, "/data/x", `[ {"op": "add", "path": "/", "value": []}, {"op": "add", "path": "-", "value": {"a": 1}}, {"op": "add", "path": "-", "value": {"a": 2}} ]`, 204, ""}, {http.MethodGet, "/data/x/1/a", "", 200, `{"result": 2}`}, }}, {"insert array", []tr{ {http.MethodPatch, "/data/x", `[{"op": "add", "path": "/", "value": { "y": [ {"z": [1,2,3]}, {"z": [4,5,6]} ] }}]`, 204, ""}, {http.MethodGet, "/data/x/y/1/z/2", "", 200, `{"result": 6}`}, {http.MethodPatch, "/data/x/y/1", `[{"op": "add", "path": "/z/1", "value": 100}]`, 204, ""}, {http.MethodGet, "/data/x/y/1/z", "", 200, `{"result": [4, 100, 5, 6]}`}, }}, {"patch root", []tr{ {http.MethodPatch, "/data", `[ {"op": "add", "path": "/", "value": {"a": 1, "b": 2} } ]`, 204, ""}, {http.MethodGet, "/data", "", 200, `{"result": {"a": 1, "b": 2}}`}, }}, {"patch invalid", []tr{ {http.MethodPatch, "/data", `[ { "op": "remove", "path": "/" } ]`, 400, ""}, }}, {"patch abort", []tr{ {http.MethodPatch, "/data", `[ {"op": "add", "path": "/foo", "value": "hello"}, {"op": "add", "path": "/bar", "value": "world"}, {"op": "add", "path": "/foo/bad", "value": "deadbeef"} ]`, 404, ""}, {http.MethodGet, "/data", "", 200, `{"result": {}}`}, }}, {"put root", []tr{ {http.MethodPut, "/data", `{"foo": [1,2,3]}`, 204, ""}, {http.MethodGet, "/data", "", 200, `{"result": {"foo": [1,2,3]}}`}, }}, {"put deep makedir", []tr{ {http.MethodPut, "/data/a/b/c/d", `1`, 204, ""}, {http.MethodGet, "/data/a/b/c", "", 200, `{"result": {"d": 1}}`}, }}, {"put deep makedir partial", []tr{ {http.MethodPut, "/data/a/b", `{}`, 204, ""}, {http.MethodPut, "/data/a/b/c/d", `0`, 204, ""}, {http.MethodGet, "/data/a/b/c", "", 200, `{"result": {"d": 0}}`}, }}, {"put exists overwrite", []tr{ {http.MethodPut, "/data/a/b/c", `"hello"`, 204, ""}, {http.MethodPut, "/data/a/b", `"goodbye"`, 204, ""}, {http.MethodGet, "/data/a", "", 200, `{"result": {"b": "goodbye"}}`}, }}, {"put base write conflict", []tr{ {http.MethodPut, "/data/a/b", `[1,2,3,4]`, 204, ""}, {http.MethodPut, "/data/a/b/c/d", "0", 404, `{ "code": "resource_conflict", "message": "storage_write_conflict_error: /a/b" }`}, }}, {"put base/virtual conflict", []tr{ {http.MethodPut, "/policies/testmod", "package x.y\np = 1\nq = 2", 200, ""}, {http.MethodPut, "/data/x", `{"y": {"p": "xxx"}}`, 400, `{ "code": "invalid_parameter", "message": "1 error occurred: testmod:2: rego_compile_error: conflicting rule for data path x/y/p found" }`}, {http.MethodPut, "/data/x/y", `{"p": "xxx"}`, 400, ``}, {http.MethodPut, "/data/x/y/p", `"xxx"`, 400, ``}, {http.MethodPut, "/data/x/y/p/a", `1`, 400, ``}, {http.MethodDelete, "/policies/testmod", "", 200, ""}, {http.MethodPut, "/data/x/y/p/a", `1`, 204, ``}, {http.MethodPut, "/policies/testmod", "package x.y\np = 1\nq = 2", 400, `{ "code": "invalid_parameter", "message": "error(s) occurred while compiling module(s)", "errors": [ { "code": "rego_compile_error", "message": "conflicting rule for data path x/y/p found", "location": { "file": "testmod", "row": 2, "col": 1 } } ] }`}, }}, {"get virtual", []tr{ {http.MethodPut, "/policies/test", testMod1, 200, ""}, {http.MethodPatch, "/data/x", `[{"op": "add", "path": "/", "value": {"y": [1,2,3,4], "z": [3,4,5,6]}}]`, 204, ""}, {http.MethodGet, "/data/testmod/p", "", 200, `{"result": [1,2]}`}, }}, {"get with input", []tr{ {http.MethodPut, "/policies/test", testMod1, 200, ""}, {http.MethodGet, "/data/testmod/g?input=%7B%22req1%22%3A%7B%22a%22%3A%5B1%5D%7D%2C+%22req2%22%3A%7B%22b%22%3A%5B0%2C1%5D%7D%7D", "", 200, `{"result": true}`}, }}, {"get with input (missing input value)", []tr{ {http.MethodPut, "/policies/test", testMod1, 200, ""}, {http.MethodGet, "/data/testmod/g?input=%7B%22req1%22%3A%7B%22a%22%3A%5B1%5D%7D%7D", "", 200, "{}"}, // req2 not specified }}, {"get with input (namespaced)", []tr{ {http.MethodPut, "/policies/test", testMod1, 200, ""}, {http.MethodGet, "/data/testmod/h?input=%7B%22req3%22%3A%7B%22attr1%22%3A%5B4%2C3%2C2%2C1%5D%7D%7D", "", 200, `{"result": true}`}, }}, {"get with input (root)", []tr{ {http.MethodPut, "/policies/test", testMod1, 200, ""}, {http.MethodGet, `/data/testmod/gt1?input={"req1":2}`, "", 200, `{"result": true}`}, }}, {"get with input (bad format)", []tr{ {http.MethodGet, "/data/deadbeef?input", "", 400, `{ "code": "invalid_parameter", "message": "parameter contains malformed input document: EOF" }`}, {http.MethodGet, "/data/deadbeef?input=", "", 400, `{ "code": "invalid_parameter", "message": "parameter contains malformed input document: EOF" }`}, {http.MethodGet, `/data/deadbeef?input="foo`, "", 400, `{ "code": "invalid_parameter", "message": "parameter contains malformed input document: unexpected EOF" }`}, }}, {"get with input (path error)", []tr{ {http.MethodGet, `/data/deadbeef?input={"foo:1}`, "", 400, `{ "code": "invalid_parameter", "message": "parameter contains malformed input document: unexpected EOF" }`}, }}, {"get empty and undefined", []tr{ {http.MethodPut, "/policies/test", testMod1, 200, ""}, {http.MethodPut, "/policies/test2", testMod5, 200, ""}, {http.MethodPut, "/policies/test3", testMod6, 200, ""}, {http.MethodGet, "/data/testmod/undef", "", 200, "{}"}, {http.MethodGet, "/data/doesnot/exist", "", 200, "{}"}, {http.MethodGet, "/data/testmod/empty/mod", "", 200, `{ "result": {} }`}, {http.MethodGet, "/data/testmod/all/undefined", "", 200, `{ "result": {} }`}, }}, {"get root", []tr{ {http.MethodPut, "/policies/test", testMod2, 200, ""}, {http.MethodPatch, "/data/x", `[{"op": "add", "path": "/", "value": [1,2,3,4]}]`, 204, ""}, {http.MethodGet, "/data", "", 200, `{"result": {"testmod": {"p": [1,2,3,4], "q": {"a":1, "b": 2}}, "x": [1,2,3,4]}}`}, }}, {"post root", []tr{ {http.MethodPost, "/data", "", 200, `{"result": {}}`}, {http.MethodPut, "/policies/test", testMod2, 200, ""}, {http.MethodPost, "/data", "", 200, `{"result": {"testmod": {"p": [1,2,3,4], "q": {"b": 2, "a": 1}}}}`}, }}, {"post input", []tr{ {http.MethodPut, "/policies/test", testMod1, 200, ""}, {http.MethodPost, "/data/testmod/gt1", `{"input": {"req1": 2}}`, 200, `{"result": true}`}, }}, {"post malformed input", []tr{ {http.MethodPost, "/data/deadbeef", `{"input": @}`, 400, `{ "code": "invalid_parameter", "message": "body contains malformed input document: invalid character '@' looking for beginning of value" }`}, }}, {"post empty object", []tr{ {http.MethodPost, "/data", `{}`, 200, `{"result": {}}`}, }}, {"post partial", []tr{ {http.MethodPut, "/policies/test", testMod7, 200, ""}, {http.MethodPost, "/data/testmod/p?partial", `{"input": {"x": 1, "y": 2, "z": 9999}}`, 200, `{"result": true}`}, {http.MethodPost, "/data/testmod/p?partial", `{"input": {"x": 1, "z": 3}}`, 200, `{"result": false}`}, {http.MethodPost, "/data/testmod/p", `{"input": {"x": 1, "y": 2, "z": 9999}}`, 200, `{"result": true}`}, {http.MethodPost, "/data/testmod/p", `{"input": {"x": 1, "z": 3}}`, 200, `{"result": false}`}, }}, {"post partial idempotent", []tr{ {http.MethodPut, "/policies/test", testMod7, 200, ""}, {http.MethodPost, "/data/testmod/p?partial", `{"input": {"x": 1, "y": 2, "z": 9999}}`, 200, `{"result": true}`}, {http.MethodPost, "/data/testmod/q?partial", `{"input": {"x": 1, "z": 3}}`, 200, `{"result": [1]}`}, {http.MethodPost, "/data/testmod/p?partial", `{"input": {"x": 1, "y": 2, "z": 9999}}`, 200, `{"result": true}`}, }}, {"partial invalidate policy", []tr{ {http.MethodPut, "/policies/test", testMod7, 200, ""}, {http.MethodPost, "/data/testmod/p?partial", `{"input": {"x": 1, "y": 2, "z": 3}}`, 200, `{"result": true}`}, {http.MethodPut, "/policies/test", testMod7Modified, 200, ""}, {http.MethodPost, "/data/testmod/p?partial", `{"input": {"x": 1, "y": 2, "z": 3}}`, 200, `{"result": false}`}, }}, {"partial invalidate data", []tr{ {http.MethodPut, "/policies/test", testMod8, 200, ""}, {http.MethodPost, "/data/testmod/p?partial", "", 200, `{}`}, {http.MethodPut, "/data/x", `1`, 204, ""}, {http.MethodPost, "/data/testmod/p?partial", "", 200, `{"result": true}`}, }}, {"partial ineffective fallback to normal", []tr{ {http.MethodPut, "/policies/test", testMod7, 200, ""}, {http.MethodPost, "/data?partial", "", 200, `{ "result": { "testmod": { "p": false, "q": [], "r": [] } } }`}, {http.MethodPost, "/data", "", 200, `{ "result": { "testmod": { "p": false, "q": [], "r": [] } } }`}, }}, {"evaluation conflict", []tr{ {http.MethodPut, "/policies/test", testMod4, 200, ""}, {http.MethodPost, "/data/testmod/p", "", 500, `{ "code": "internal_error", "errors": [ { "code": "eval_conflict_error", "location": { "col": 1, "file": "test", "row": 4 }, "message": "complete rules must not produce multiple outputs" } ], "message": "error(s) occurred while evaluating query" }`}, }}, {"query wildcards omitted", []tr{ {http.MethodPatch, "/data/x", `[{"op": "add", "path": "/", "value": [1,2,3,4]}]`, 204, ""}, {http.MethodGet, "/query?q=data.x[_]%20=%20x", "", 200, `{"result": [{"x": 1}, {"x": 2}, {"x": 3}, {"x": 4}]}`}, }}, {"query undefined", []tr{ {http.MethodGet, "/query?q=a=1%3Bb=2%3Ba=b", "", 200, `{}`}, }}, {"query compiler error", []tr{ {http.MethodGet, "/query?q=x", "", 400, ""}, // Subsequent query should not fail. {http.MethodGet, "/query?q=x=1", "", 200, `{"result": [{"x": 1}]}`}, }}, {"delete and check", []tr{ {http.MethodDelete, "/data/a/b", "", 404, ""}, {http.MethodPut, "/data/a/b/c/d", `1`, 204, ""}, {http.MethodGet, "/data/a/b/c", "", 200, `{"result": {"d": 1}}`}, {http.MethodDelete, "/data/a/b", "", 204, ""}, {http.MethodGet, "/data/a/b/c/d", "", 200, `{}`}, {http.MethodGet, "/data/a", "", 200, `{"result": {}}`}, {http.MethodGet, "/data/a/b/c", "", 200, `{}`}, }}, {"escaped paths", []tr{ {http.MethodPut, "/data/a%2Fb", `{"c/d": 1}`, 204, ""}, {http.MethodGet, "/data", "", 200, `{"result": {"a/b": {"c/d": 1}}}`}, {http.MethodGet, "/data/a%2Fb/c%2Fd", "", 200, `{"result": 1}`}, {http.MethodGet, "/data/a/b", "", 200, `{}`}, {http.MethodPost, "/data/a%2Fb/c%2Fd", "", 200, `{"result": 1}`}, {http.MethodPost, "/data/a/b", "", 200, `{}`}, {http.MethodPatch, "/data/a%2Fb", `[{"op": "add", "path": "/e%2Ff", "value": 2}]`, 204, ""}, {http.MethodPost, "/data", "", 200, `{"result": {"a/b": {"c/d": 1, "e/f": 2}}}`}, }}, {"strict-builtin-errors", []tr{ {http.MethodPut, "/policies/test", ` package test default p = false p { 1/0 } `, 200, ""}, {http.MethodGet, "/data/test/p", "", 200, `{"result": false}`}, {http.MethodGet, "/data/test/p?strict-builtin-errors", "", 500, `{ "code": "internal_error", "message": "error(s) occurred while evaluating query", "errors": [ { "code": "eval_builtin_error", "message": "div: divide by zero", "location": { "file": "test", "row": 6, "col": 9 } } ] }`}, {http.MethodPost, "/data/test/p", "", 200, `{"result": false}`}, {http.MethodPost, "/data/test/p?strict-builtin-errors", "", 500, `{ "code": "internal_error", "message": "error(s) occurred while evaluating query", "errors": [ { "code": "eval_builtin_error", "message": "div: divide by zero", "location": { "file": "test", "row": 6, "col": 9 } } ] }`}, }}, } for _, tc := range tests { test.Subtest(t, tc.note, func(t *testing.T) { executeRequests(t, tc.reqs) }) } } func TestDataYAML(t *testing.T) { testMod1 := `package testmod import input.req1 gt1 = true { req1 > 1 }` inputYaml1 := ` --- input: req1: 2` inputYaml2 := ` --- req1: 2` f := newFixture(t) if err := f.v1(http.MethodPut, "/policies/test", testMod1, 200, ""); err != nil { t.Fatalf("Unexpected error from PUT /policies/test: %v", err) } // First JSON and then later yaml to make sure both work if err := f.v1(http.MethodPost, "/data/testmod/gt1", `{"input": {"req1": 2}}`, 200, `{"result": true}`); err != nil { t.Fatalf("Unexpected error from PUT /policies/test: %v", err) } req := newReqV1(http.MethodPost, "/data/testmod/gt1", inputYaml1) req.Header.Set("Content-Type", "application/x-yaml") if err := f.executeRequest(req, 200, `{"result": true}`); err != nil { t.Fatalf("Unexpected error from POST with yaml: %v", err) } req = newReqV0(http.MethodPost, "/data/testmod/gt1", inputYaml2) req.Header.Set("Content-Type", "application/x-yaml") if err := f.executeRequest(req, 200, `true`); err != nil { t.Fatalf("Unexpected error from POST with yaml: %v", err) } if err := f.v1(http.MethodPut, "/policies/test2", `package system main = data.testmod.gt1`, 200, ""); err != nil { t.Fatalf("Unexpected error from PUT /policies/test: %v", err) } req = newReqUnversioned(http.MethodPost, "/", inputYaml2) req.Header.Set("Content-Type", "application/x-yaml") if err := f.executeRequest(req, 200, `true`); err != nil { t.Fatalf("Unexpected error from POST with yaml: %v", err) } } func TestDataPutV1IfNoneMatch(t *testing.T) { f := newFixture(t) if err := f.v1(http.MethodPut, "/data/a/b/c", "0", 204, ""); err != nil { t.Fatalf("Unexpected error from PUT /data/a/b/c: %v", err) } req := newReqV1(http.MethodPut, "/data/a/b/c", "1") req.Header.Set("If-None-Match", "*") if err := f.executeRequest(req, 304, ""); err != nil { t.Fatalf("Unexpected error from PUT with If-None-Match=*: %v", err) } } func TestParsePatchPathEscaped(t *testing.T) { tests := []struct { note string path string expectedPath storage.Path expectedOK bool }{ // success-path tests { note: "single-level", path: "/single-level", expectedPath: storage.Path{"single-level"}, expectedOK: true, }, { note: "multi-level", path: "/a/multi-level/path", expectedPath: storage.Path{"a", "multi-level", "path"}, expectedOK: true, }, { note: "end", path: "/-", expectedPath: storage.Path{"-"}, expectedOK: true, }, { // not strictly correct but included for backwards compatibility with existing OPA note: "url-escaped forward slash", path: "/github.com%2Fopen-policy-agent", expectedPath: storage.Path{"github.com/open-policy-agent"}, expectedOK: true, }, { note: "json-pointer-escaped forward slash", path: "/github.com~1open-policy-agent", expectedPath: storage.Path{"github.com/open-policy-agent"}, expectedOK: true, }, { note: "json-pointer-escaped tilde", path: "/~0opa", expectedPath: storage.Path{"~opa"}, expectedOK: true, }, { note: "json-pointer-escape correctness", path: "/~01", expectedPath: storage.Path{"~1"}, expectedOK: true, }, // failure-path tests { // not possible with existing callers but for completeness... note: "empty string", path: "", expectedOK: false, }, { // not possible with existing callers but for completeness... note: "string that doesn't start with /", path: "foo", expectedOK: false, }, } for _, tc := range tests { t.Run(tc.note, func(t *testing.T) { actualPath, actualOK := parsePatchPathEscaped(tc.path) if tc.expectedOK != actualOK { t.Fatalf("Expected ok to be %v but was %v", tc.expectedOK, actualOK) } if !reflect.DeepEqual(tc.expectedPath, actualPath) { t.Fatalf("Expected path to be %v but was %v", tc.expectedPath, actualPath) } }) } } func TestBundleScope(t *testing.T) { ctx := context.Background() f := newFixture(t) txn := storage.NewTransactionOrDie(ctx, f.server.store, storage.WriteParams) if err := bundle.WriteManifestToStore(ctx, f.server.store, txn, "test-bundle", bundle.Manifest{ Revision: "AAAAA", Roots: &[]string{"a/b/c", "x/y"}, }); err != nil { t.Fatal(err) } if err := f.server.store.UpsertPolicy(ctx, txn, "someid", []byte(`package x.y.z`)); err != nil { t.Fatal(err) } if err := f.server.store.Commit(ctx, txn); err != nil { t.Fatal(err) } cases := []tr{ { method: "PUT", path: "/data/a/b", body: "1", code: http.StatusBadRequest, resp: `{"code": "invalid_parameter", "message": "path a/b is owned by bundle \"test-bundle\""}`, }, { method: "PUT", path: "/data/a/b/c", body: "1", code: http.StatusBadRequest, resp: `{"code": "invalid_parameter", "message": "path a/b/c is owned by bundle \"test-bundle\""}`, }, { method: "PUT", path: "/data/a/b/c/d", body: "1", code: http.StatusBadRequest, resp: `{"code": "invalid_parameter", "message": "path a/b/c/d is owned by bundle \"test-bundle\""}`, }, { method: "PATCH", path: "/data/a", body: `[{"path": "/b/c", "op": "add", "value": 1}]`, code: http.StatusBadRequest, resp: `{"code": "invalid_parameter", "message": "path a/b/c is owned by bundle \"test-bundle\""}`, }, { method: "DELETE", path: "/data/a", code: http.StatusBadRequest, resp: `{"code": "invalid_parameter", "message": "path a is owned by bundle \"test-bundle\""}`, }, { method: "PUT", path: "/policies/test1", body: `package a.b`, code: http.StatusBadRequest, resp: `{"code": "invalid_parameter", "message": "path a/b is owned by bundle \"test-bundle\""}`, }, { method: "DELETE", path: "/policies/someid", code: http.StatusBadRequest, resp: `{"code": "invalid_parameter", "message": "path x/y/z is owned by bundle \"test-bundle\""}`, }, { method: "PUT", path: "/data/foo/bar", body: "1", code: http.StatusNoContent, }, } if err := f.v1TestRequests(cases); err != nil { t.Fatal(err) } } func TestBundleScopeMultiBundle(t *testing.T) { ctx := context.Background() f := newFixture(t) txn := storage.NewTransactionOrDie(ctx, f.server.store, storage.WriteParams) if err := bundle.WriteManifestToStore(ctx, f.server.store, txn, "test-bundle1", bundle.Manifest{ Revision: "AAAAA", Roots: &[]string{"a/b/c", "x/y"}, }); err != nil { t.Fatal(err) } if err := bundle.WriteManifestToStore(ctx, f.server.store, txn, "test-bundle2", bundle.Manifest{ Revision: "AAAAA", Roots: &[]string{"a/b/d"}, }); err != nil { t.Fatal(err) } if err := bundle.WriteManifestToStore(ctx, f.server.store, txn, "test-bundle3", bundle.Manifest{ Revision: "AAAAA", Roots: &[]string{"a/b/e", "a/b/f"}, }); err != nil { t.Fatal(err) } if err := f.server.store.UpsertPolicy(ctx, txn, "someid", []byte(`package x.y.z`)); err != nil { t.Fatal(err) } if err := f.server.store.Commit(ctx, txn); err != nil { t.Fatal(err) } cases := []tr{ { method: "PUT", path: "/data/x/y", body: "1", code: http.StatusBadRequest, resp: `{"code": "invalid_parameter", "message": "path x/y is owned by bundle \"test-bundle1\""}`, }, { method: "PUT", path: "/data/a/b/d", body: "1", code: http.StatusBadRequest, resp: `{"code": "invalid_parameter", "message": "path a/b/d is owned by bundle \"test-bundle2\""}`, }, { method: "PUT", path: "/data/foo/bar", body: "1", code: http.StatusNoContent, }, } if err := f.v1TestRequests(cases); err != nil { t.Fatal(err) } } func TestDataGetExplainFull(t *testing.T) { f := newFixture(t) f.v1(http.MethodPut, "/data/x", `{"a":1,"b":2}`, 204, "") req := newReqV1(http.MethodGet, "/data/x?explain=full", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) var result types.DataResponseV1 if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } explain := mustUnmarshalTrace(result.Explanation) nexpect := 5 if len(explain) != nexpect { t.Fatalf("Expected exactly %d events but got %d", nexpect, len(explain)) } _, ok := explain[2].Node.(ast.Body) if !ok { t.Fatalf("Expected body for node but got: %v", explain[2].Node) } if len(explain[2].Locals) != 1 { t.Fatalf("Expected one binding but got: %v", explain[2].Locals) } req = newReqV1(http.MethodGet, "/data/deadbeef?explain=full", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) result = types.DataResponseV1{} if f.recorder.Code != 200 { t.Fatalf("Expected status code to be 200 but got: %v", f.recorder.Code) } if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } explain = mustUnmarshalTrace(result.Explanation) if len(explain) != 3 { t.Fatalf("Expected exactly 3 events but got %d", len(explain)) } if explain[2].Op != "fail" { t.Fatalf("Expected last event to be 'fail' but got: %v", explain[2]) } req = newReqV1(http.MethodGet, "/data/x?explain=full&pretty=true", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) result = types.DataResponseV1{} if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } exp := []interface{}{ `query:1 Enter data.x = _`, `query:1 | Eval data.x = _`, `query:1 | Exit data.x = _`, `query:1 Redo data.x = _`, `query:1 | Redo data.x = _`} actual := util.MustUnmarshalJSON(result.Explanation).([]interface{}) if !reflect.DeepEqual(actual, exp) { t.Fatalf(`Expected pretty explanation to be %v, got %v`, exp, actual) } } func TestDataPostExplain(t *testing.T) { f := newFixture(t) f.v1(http.MethodPut, "/policies/test", `package test p = [1, 2, 3, 4] { true }`, 200, "") req := newReqV1(http.MethodPost, "/data/test/p?explain=full", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) var result types.DataResponseV1 if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } explain := mustUnmarshalTrace(result.Explanation) nexpect := 11 if len(explain) != nexpect { t.Fatalf("Expected exactly %d events but got %d", nexpect, len(explain)) } var expected interface{} if err := util.UnmarshalJSON([]byte(`[1,2,3,4]`), &expected); err != nil { panic(err) } if result.Result == nil || !reflect.DeepEqual(*result.Result, expected) { t.Fatalf("Expected %v but got: %v", expected, result.Result) } } func TestDataPostExplainNotes(t *testing.T) { f := newFixture(t) f.v1(http.MethodPut, "/policies/test", ` package test p { data.a[i] = x; x > 1 trace(sprintf("found x = %d", [x])) }`, 200, "") f.v1(http.MethodPut, "/data/a", `[1,2,3]`, 200, "") f.reset() req := newReqV1(http.MethodPost, "/data/test/p?explain=notes", "") f.server.Handler.ServeHTTP(f.recorder, req) var result types.DataResponseV1 if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode err: %v", err) } var trace types.TraceV1Raw if err := trace.UnmarshalJSON(result.Explanation); err != nil { t.Fatal(err) } if len(trace) != 6 || trace[2].Op != "note" || trace[5].Op != "note" { t.Logf("Found %d events in trace", len(trace)) for i := range trace { t.Logf("Event #%d: %v\n", i, trace[i]) } t.Fatal("Unexpected trace") } } func TestDataProvenanceSingleBundle(t *testing.T) { f := newFixture(t) // Dummy up since we are not using ld... // Note: No bundle 'revision'... version.Version = "0.10.7" version.Vcs = "ac23eb45" version.Timestamp = "today" version.Hostname = "foo.bar.com" // Initialize as if a bundle plugin is running bp := pluginBundle.New(&pluginBundle.Config{Name: "b1"}, f.server.manager) f.server.manager.Register(pluginBundle.Name, bp) req := newReqV1(http.MethodPost, "/data?provenance", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) var result types.DataResponseV1 if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } if result.Provenance == nil { t.Fatalf("Expected non-nil provenance: %v", result.Provenance) } expectedProvenance := &types.ProvenanceV1{ Version: version.Version, Vcs: version.Vcs, Timestamp: version.Timestamp, Hostname: version.Hostname, } if !reflect.DeepEqual(result.Provenance, expectedProvenance) { t.Errorf("Unexpected provenance data: \n\n%+v\n\nExpected:\n%+v\n\n", result.Provenance, expectedProvenance) } // Update bundle revision and request again f.server.revisions["b1"] = "r1" f.server.legacyRevision = "r1" req = newReqV1(http.MethodPost, "/data?provenance", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) result = types.DataResponseV1{} if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } if result.Provenance == nil { t.Fatalf("Expected non-nil provenance: %v", result.Provenance) } expectedProvenance.Revision = "r1" if !reflect.DeepEqual(result.Provenance, expectedProvenance) { t.Errorf("Unexpected provenance data: \n\n%+v\n\nExpected:\n%+v\n\n", result.Provenance, expectedProvenance) } } func TestDataProvenanceSingleFileBundle(t *testing.T) { f := newFixture(t) // Dummy up since we are not using ld... // Note: No bundle 'revision'... version.Version = "0.10.7" version.Vcs = "ac23eb45" version.Timestamp = "today" version.Hostname = "foo.bar.com" // No bundle plugin initialized, just a legacy revision set f.server.legacyRevision = "r1" req := newReqV1(http.MethodPost, "/data?provenance", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) result := types.DataResponseV1{} if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } if result.Provenance == nil { t.Fatalf("Expected non-nil provenance: %v", result.Provenance) } expectedProvenance := &types.ProvenanceV1{ Version: version.Version, Vcs: version.Vcs, Timestamp: version.Timestamp, Hostname: version.Hostname, } expectedProvenance.Revision = "r1" if !reflect.DeepEqual(result.Provenance, expectedProvenance) { t.Errorf("Unexpected provenance data: \n\n%+v\n\nExpected:\n%+v\n\n", result.Provenance, expectedProvenance) } } func TestDataProvenanceMultiBundle(t *testing.T) { f := newFixture(t) // Dummy up since we are not using ld... version.Version = "0.10.7" version.Vcs = "ac23eb45" version.Timestamp = "today" version.Hostname = "foo.bar.com" // Initialize as if a bundle plugin is running with 2 bundles bp := pluginBundle.New(&pluginBundle.Config{Bundles: map[string]*pluginBundle.Source{ "b1": {Service: "s1", Resource: "bundle.tar.gz"}, "b2": {Service: "s2", Resource: "bundle.tar.gz"}, }}, f.server.manager) f.server.manager.Register(pluginBundle.Name, bp) req := newReqV1(http.MethodPost, "/data?provenance", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) var result types.DataResponseV1 if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } if result.Provenance == nil { t.Fatalf("Expected non-nil provenance: %v", result.Provenance) } expectedProvenance := &types.ProvenanceV1{ Version: version.Version, Vcs: version.Vcs, Timestamp: version.Timestamp, Hostname: version.Hostname, } if !reflect.DeepEqual(result.Provenance, expectedProvenance) { t.Errorf("Unexpected provenance data: \n\n%+v\n\nExpected:\n%+v\n\n", result.Provenance, expectedProvenance) } // Update bundle revision for a single bundle and make the request again f.server.revisions["b1"] = "r1" req = newReqV1(http.MethodPost, "/data?provenance", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) result = types.DataResponseV1{} if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } if result.Provenance == nil { t.Fatalf("Expected non-nil provenance: %v", result.Provenance) } expectedProvenance.Bundles = map[string]types.ProvenanceBundleV1{ "b1": {Revision: "r1"}, } if !reflect.DeepEqual(result.Provenance, expectedProvenance) { t.Errorf("Unexpected provenance data: \n\n%+v\n\nExpected:\n%+v\n\n", result.Provenance, expectedProvenance) } // Update both and check again f.server.revisions["b1"] = "r2" f.server.revisions["b2"] = "r1" req = newReqV1(http.MethodPost, "/data?provenance", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) result = types.DataResponseV1{} if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } if result.Provenance == nil { t.Fatalf("Expected non-nil provenance: %v", result.Provenance) } expectedProvenance.Bundles = map[string]types.ProvenanceBundleV1{ "b1": {Revision: "r2"}, "b2": {Revision: "r1"}, } if !reflect.DeepEqual(result.Provenance, expectedProvenance) { t.Errorf("Unexpected provenance data: \n\n%+v\n\nExpected:\n%+v\n\n", result.Provenance, expectedProvenance) } } func TestDataMetricsEval(t *testing.T) { f := newFixture(t) // Make a request to evaluate `data` testDataMetrics(t, f, "/data?metrics", []string{ "counter_server_query_cache_hit", "timer_rego_input_parse_ns", "timer_rego_query_parse_ns", "timer_rego_query_compile_ns", "timer_rego_query_eval_ns", "timer_server_handler_ns", }) // Repeat previous request, expect to have hit the query cache // so fewer timers should have been reported. testDataMetrics(t, f, "/data?metrics", []string{ "counter_server_query_cache_hit", "timer_rego_input_parse_ns", "timer_rego_query_eval_ns", "timer_server_handler_ns", }) // Make a request to evaluate `data` and use partial evaluation, // this should not hit the same query cache result as the previous // request. testDataMetrics(t, f, "/data?metrics&partial", []string{ "counter_server_query_cache_hit", "timer_rego_input_parse_ns", "timer_rego_module_compile_ns", "timer_rego_query_parse_ns", "timer_rego_query_compile_ns", "timer_rego_query_eval_ns", "timer_rego_partial_eval_ns", "timer_server_handler_ns", }) // Repeat previous partial eval request, this time it should // be cached testDataMetrics(t, f, "/data?metrics&partial", []string{ "counter_server_query_cache_hit", "timer_rego_input_parse_ns", "timer_rego_query_eval_ns", "timer_server_handler_ns", }) } func testDataMetrics(t *testing.T, f *fixture, url string, expected []string) { t.Helper() f.reset() req := newReqV1(http.MethodPost, url, "") f.server.Handler.ServeHTTP(f.recorder, req) var result types.DataResponseV1 if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } for _, key := range expected { v, ok := result.Metrics[key] if !ok { t.Errorf("Missing expected metric: %s", key) } else if v == nil { t.Errorf("Expected non-nil value for metric: %s", key) } } if len(expected) != len(result.Metrics) { t.Errorf("Expected %d metrics, got %d\n\n\tValues: %+v", len(expected), len(result.Metrics), result.Metrics) } } func TestV1Pretty(t *testing.T) { f := newFixture(t) f.v1(http.MethodPatch, "/data/x", `[{"op": "add", "path":"/", "value": [1,2,3,4]}]`, 204, "") req := newReqV1(http.MethodGet, "/data/x?pretty=true", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) lines := strings.Split(f.recorder.Body.String(), "\n") if len(lines) != 9 { t.Errorf("Expected 8 lines in output but got %d:\n%v", len(lines), lines) } req = newReqV1(http.MethodGet, "/query?q=data.x[i]&pretty=true", "") f.reset() f.server.Handler.ServeHTTP(f.recorder, req) lines = strings.Split(f.recorder.Body.String(), "\n") if len(lines) != 17 { t.Errorf("Expected 16 lines of output but got %d:\n%v", len(lines), lines) } } func TestIndexGetEscaped(t *testing.T) { f := newFixture(t) get, err := http.NewRequest(http.MethodGet, `/?q=`, strings.NewReader("")) if err != nil { panic(err) } f.server.Handler.ServeHTTP(f.recorder, get) if f.recorder.Code != 200 { t.Errorf("Expected success but got: %v", f.recorder) return } page := f.recorder.Body.String() exp := "</textarea><script>alert(1)</script>" if !strings.Contains(page, exp) { t.Fatalf("Expected page to contain escaped URL parameter but got: %v", page) } } func TestIndexGet(t *testing.T) { f := newFixture(t) get, err := http.NewRequest(http.MethodGet, `/?q=foo = 1&input=`, strings.NewReader("")) if err != nil { panic(err) } f.server.Handler.ServeHTTP(f.recorder, get) if f.recorder.Code != 200 { t.Errorf("Expected success but got: %v", f.recorder) return } page := f.recorder.Body.String() if !strings.Contains(page, "Query result") { t.Errorf("Expected page to contain 'Query result' but got: %v", page) return } } func TestIndexGetCompileError(t *testing.T) { f := newFixture(t) // "foo" is not bound get, err := http.NewRequest(http.MethodGet, `/?q=foo`, strings.NewReader("")) if err != nil { panic(err) } f.server.Handler.ServeHTTP(f.recorder, get) if f.recorder.Code != 200 { t.Errorf("Expected success but got: %v", f.recorder) return } page := f.recorder.Body.String() if !strings.Contains(page, "foo is unsafe") { t.Errorf("Expected page to contain 'foo is unsafe' but got: %v", page) return } } func TestPoliciesPutV1(t *testing.T) { f := newFixture(t) req := newReqV1(http.MethodPut, "/policies/1", testMod) f.server.Handler.ServeHTTP(f.recorder, req) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } var response map[string]interface{} if err := json.NewDecoder(f.recorder.Body).Decode(&response); err != nil { t.Fatalf("Unexpected error while unmarshalling response: %v", err) } if len(response) != 0 { t.Fatalf("Expected empty wrapper object") } } func TestPoliciesPutV1Empty(t *testing.T) { f := newFixture(t) req := newReqV1(http.MethodPut, "/policies/1", "") f.server.Handler.ServeHTTP(f.recorder, req) if f.recorder.Code != 400 { t.Fatalf("Expected bad request but got %v", f.recorder) } } func TestPoliciesPutV1ParseError(t *testing.T) { f := newFixture(t) req := newReqV1(http.MethodPut, "/policies/test", ` package a.b.c p ;- true `) f.server.Handler.ServeHTTP(f.recorder, req) if f.recorder.Code != 400 { t.Fatalf("Expected bad request but got %v", f.recorder) } response := map[string]interface{}{} if err := util.NewJSONDecoder(f.recorder.Body).Decode(&response); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } if !reflect.DeepEqual(response["code"], types.CodeInvalidParameter) { t.Fatalf("Expected code %v but got: %v", types.CodeInvalidParameter, response) } v := ast.MustInterfaceToValue(response) name, err := v.Find(ast.MustParseRef("_.errors[0].location.file")[1:]) if err != nil { t.Fatalf("Expecfted to find name in errors but: %v", err) } if name.Compare(ast.String("test")) != 0 { t.Fatalf("Expected name ot equal test but got: %v", name) } req = newReqV1(http.MethodPut, "/policies/test", ``) f.reset() f.server.Handler.ServeHTTP(f.recorder, req) if f.recorder.Code != 400 { t.Fatalf("Expected bad request but got %v", f.recorder) } req = newReqV1(http.MethodPut, "/policies/test", ` package a.b.c p = true`) f.reset() f.server.Handler.ServeHTTP(f.recorder, req) if f.recorder.Code != 200 { t.Fatalf("Expected ok but got %v", f.recorder) } } func TestPoliciesPutV1CompileError(t *testing.T) { f := newFixture(t) req := newReqV1(http.MethodPut, "/policies/test", `package a.b.c p[x] { q[x] } q[x] { p[x] }`, ) f.server.Handler.ServeHTTP(f.recorder, req) if f.recorder.Code != 400 { t.Fatalf("Expected bad request but got %v", f.recorder) } response := map[string]interface{}{} if err := util.NewJSONDecoder(f.recorder.Body).Decode(&response); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } if !reflect.DeepEqual(response["code"], types.CodeInvalidParameter) { t.Fatalf("Expected code %v but got: %v", types.CodeInvalidParameter, response) } v := ast.MustInterfaceToValue(response) name, err := v.Find(ast.MustParseRef("_.errors[0].location.file")[1:]) if err != nil { t.Fatalf("Expecfted to find name in errors but: %v", err) } if name.Compare(ast.String("test")) != 0 { t.Fatalf("Expected name ot equal test but got: %v", name) } } func TestPoliciesPutV1Noop(t *testing.T) { f := newFixture(t) f.v1("PUT", "/policies/test?metrics", `package foo`, 200, "") f.reset() f.v1("PUT", "/policies/test?metrics", `package foo`, 200, "") var resp types.PolicyPutResponseV1 if err := json.NewDecoder(f.recorder.Body).Decode(&resp); err != nil { t.Fatal(err) } exp := []string{"timer_server_read_bytes_ns"} // Sort the metric keys and compare to expected value. We're assuming the // server skips parsing if the bytes are equal. result := []string{} for k := range resp.Metrics { result = append(result, k) } sort.Strings(result) if !reflect.DeepEqual(exp, result) { t.Fatalf("Expected %v but got %v", exp, result) } f.reset() // Ensure subsequent update with changed policy parses the body. f.v1("PUT", "/policies/test?metrics", "package foo\np = 1", 200, "") var resp2 types.PolicyPutResponseV1 if err := json.NewDecoder(f.recorder.Body).Decode(&resp2); err != nil { t.Fatal(err) } if _, ok := resp2.Metrics["timer_rego_module_parse_ns"]; !ok { t.Fatalf("Expected parse module metric in response but got %v", resp2) } } func TestPoliciesListV1(t *testing.T) { f := newFixture(t) putPolicy(t, f, testMod) expected := []types.PolicyV1{ newPolicy("1", testMod), } assertListPolicy(t, f, expected) } func TestPoliciesListV1AfterPartialEval(t *testing.T) { f := newFixture(t) putPolicy(t, f, testMod) expected := []types.PolicyV1{ newPolicy("1", testMod), } assertListPolicy(t, f, expected) eval := newReqV1("POST", "/data?partial", "{}") f.server.Handler.ServeHTTP(f.recorder, eval) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } f.reset() eval = newReqV1("POST", "/data/a/b?partial", "{}") f.server.Handler.ServeHTTP(f.recorder, eval) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } f.reset() // Doesn't matter what the results of eval w/ partial were // We do expect that the partially evaluated policy is _not_ in the listed policies assertListPolicy(t, f, expected) } func putPolicy(t *testing.T, f *fixture, mod string) { t.Helper() put := newReqV1(http.MethodPut, "/policies/1", mod) f.server.Handler.ServeHTTP(f.recorder, put) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } f.reset() } func assertListPolicy(t *testing.T, f *fixture, expected []types.PolicyV1) { t.Helper() list := newReqV1(http.MethodGet, "/policies", "") f.server.Handler.ServeHTTP(f.recorder, list) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } // var policies []*PolicyV1 var response types.PolicyListResponseV1 err := util.NewJSONDecoder(f.recorder.Body).Decode(&response) if err != nil { t.Fatalf("Expected policy list but got error: %v with response body:\n\n%v\n", err, f.recorder) } if len(expected) != len(response.Result) { t.Fatalf("Expected %d policies but got: %v", len(expected), response.Result) } for i := range expected { if !expected[i].Equal(response.Result[i]) { t.Fatalf("Expected policies to be equal. Expected:\n\n%v\n\nGot:\n\n%+v\n", expected, response.Result) } } f.reset() } func TestPoliciesGetV1(t *testing.T) { f := newFixture(t) put := newReqV1(http.MethodPut, "/policies/1", testMod) f.server.Handler.ServeHTTP(f.recorder, put) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } f.reset() get := newReqV1(http.MethodGet, "/policies/1", "") f.server.Handler.ServeHTTP(f.recorder, get) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } var response types.PolicyGetResponseV1 if err := util.NewJSONDecoder(f.recorder.Body).Decode(&response); err != nil { t.Fatalf("Unexpected error: %v", err) } expected := newPolicy("1", testMod) if !expected.Equal(response.Result) { t.Errorf("Expected policies to be equal. Expected:\n\n%v\n\nGot:\n\n%v\n", expected, response.Result) } } func TestPoliciesDeleteV1(t *testing.T) { f := newFixture(t) put := newReqV1(http.MethodPut, "/policies/1", testMod) f.server.Handler.ServeHTTP(f.recorder, put) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } f.reset() del := newReqV1(http.MethodDelete, "/policies/1", "") f.server.Handler.ServeHTTP(f.recorder, del) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } var response map[string]interface{} if err := json.NewDecoder(f.recorder.Body).Decode(&response); err != nil { t.Fatalf("Unexpected unmarshal error: %v", err) } if len(response) > 0 { t.Fatalf("Expected empty response but got: %v", response) } f.reset() get := newReqV1(http.MethodGet, "/policies/1", "") f.server.Handler.ServeHTTP(f.recorder, get) if f.recorder.Code != 404 { t.Fatalf("Expected not found but got %v", f.recorder) } } func TestPoliciesPathSlashes(t *testing.T) { f := newFixture(t) if err := f.v1(http.MethodPut, "/policies/a/b/c.rego", testMod, 200, ""); err != nil { t.Fatalf("Unexpected error: %v", err) } if err := f.v1(http.MethodGet, "/policies/a/b/c.rego", testMod, 200, ""); err != nil { t.Fatalf("Unexpected error: %v", err) } } func TestQueryPostBasic(t *testing.T) { f := newFixture(t) f.server, _ = New(). WithAddresses([]string{":8182"}). WithStore(f.server.store). WithManager(f.server.manager). 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 TestDecisionIDs(t *testing.T) { f := newFixture(t) ids := []string{} ctr := 0 f.server = f.server.WithDecisionLoggerWithErr(func(_ context.Context, info *Info) error { ids = append(ids, info.DecisionID) return nil }).WithDecisionIDFactory(func() string { ctr++ return fmt.Sprint(ctr) }) if err := f.v1("GET", "/data/undefined", "", 200, `{"decision_id": "1"}`); err != nil { t.Fatal(err) } if err := f.v1("POST", "/data/undefined", "", 200, `{"decision_id": "2"}`); err != nil { t.Fatal(err) } if err := f.v1("GET", "/data", "", 200, `{"decision_id": "3", "result": {}}`); err != nil { t.Fatal(err) } if err := f.v1("POST", "/data", "", 200, `{"decision_id": "4", "result": {}}`); err != nil { t.Fatal(err) } exp := []string{"1", "2", "3", "4"} if !reflect.DeepEqual(ids, exp) { t.Fatalf("Expected %v but got %v", exp, ids) } } func TestDecisionLogging(t *testing.T) { f := newFixture(t) decisions := []*Info{} var nextID int f.server = f.server.WithDecisionIDFactory(func() string { nextID++ return fmt.Sprint(nextID) }).WithDecisionLoggerWithErr(func(_ context.Context, info *Info) error { if info.Path == "fail_closed/decision_logger_err" { return fmt.Errorf("some error") } decisions = append(decisions, info) return nil }) reqs := []struct { raw *http.Request v0 bool method string path string body string code int response string }{ { method: "PUT", path: "/policies/test", body: "package system\nmain=true", response: "{}", }, { method: "POST", path: "/data", response: `{"result": {}, "decision_id": "1"}`, }, { method: "GET", path: "/data", response: `{"result": {}, "decision_id": "2"}`, }, { method: "POST", path: "/data/nonexistent", body: `{"input": {"foo": 1}}`, response: `{"decision_id": "3"}`, }, { method: "POST", v0: true, path: "/data", response: `{}`, }, { raw: newReqUnversioned("POST", "/", ""), response: "true", }, { method: "GET", path: "/query?q=data=x", response: `{"result": [{"x": {}}]}`, }, { method: "POST", path: "/query", body: `{"query": "data=x"}`, response: `{"result": [{"x": {}}]}`, }, { method: "PUT", path: "/policies/test2", body: `package foo p { {k: v | k = ["a", "a"][_]; v = [1, 2][_]} }`, response: `{}`, }, { method: "PUT", path: "/policies/test", body: "package system\nmain { data.foo.p }", response: `{}`, }, { method: "POST", path: "/data", code: 500, }, { method: "GET", path: "/data", code: 500, }, { raw: newReqUnversioned("POST", "/", ""), code: 500, }, { method: "POST", path: "/data/fail_closed/decision_logger_err", code: 500, }, { method: "POST", v0: true, path: "/data/test", code: 404, response: `{ "code": "undefined_document", "message": "document missing or undefined: data.test" }`, }, } for _, r := range reqs { code := r.code if code == 0 { code = http.StatusOK } if r.raw != nil { if err := f.executeRequest(r.raw, code, r.response); err != nil { t.Fatal(err) } } else if r.v0 { if err := f.v0(r.method, r.path, r.body, code, r.response); err != nil { t.Fatal(err) } } else { if err := f.v1(r.method, r.path, r.body, code, r.response); err != nil { t.Fatal(err) } } } exp := []struct { input string path string query string wantErr bool }{ {path: ""}, {path: ""}, {path: "nonexistent", input: `{"foo": 1}`}, {path: ""}, {path: "system/main"}, {query: "data = x"}, {query: "data = x"}, {path: "", wantErr: true}, {path: "", wantErr: true}, {path: "system/main", wantErr: true}, {path: `test`, wantErr: true}, } if len(decisions) != len(exp) { t.Fatalf("Expected exactly %d decisions but got: %d", len(exp), len(decisions)) } for i, d := range decisions { if d.DecisionID == "" { t.Fatalf("Expected decision ID on decision %d but got: %v", i, d) } if d.Metrics.Timer(metrics.ServerHandler).Value() == 0 { t.Fatalf("Expected server handler timer to be started on decision %d but got %v", i, d) } if exp[i].path != d.Path || exp[i].query != d.Query { t.Fatalf("Unexpected path or query on %d, want: %v but got: %v", i, exp[i], d) } if exp[i].wantErr && d.Error == nil || !exp[i].wantErr && d.Error != nil { t.Fatalf("Unexpected error on %d, wantErr: %v, got: %v", i, exp[i].wantErr, d) } if exp[i].input != "" { input := util.MustUnmarshalJSON([]byte(exp[i].input)) if d.Input == nil || !reflect.DeepEqual(input, *d.Input) { t.Fatalf("Unexpected input on %d, want: %v, but got: %v", i, exp[i], d) } } } } func TestDecisionLogErrorMessage(t *testing.T) { f := newFixture(t) f.server.WithDecisionLoggerWithErr(func(context.Context, *Info) error { return fmt.Errorf("xxx") }) if err := f.v1(http.MethodPost, "/data", "", 500, `{ "code": "internal_error", "message": "decision_logs: xxx" }`); err != nil { t.Fatal(err) } } func TestQueryV1(t *testing.T) { f := newFixture(t) get := newReqV1(http.MethodGet, `/query?q=a=[1,2,3]%3Ba[i]=x`, "") f.server.Handler.ServeHTTP(f.recorder, get) if f.recorder.Code != 200 { t.Fatalf("Expected success but got %v", f.recorder) } var expected types.QueryResponseV1 err := util.UnmarshalJSON([]byte(`{ "result": [{"a":[1,2,3],"i":0,"x":1},{"a":[1,2,3],"i":1,"x":2},{"a":[1,2,3],"i":2,"x":3}] }`), &expected) if err != nil { panic(err) } var result types.QueryResponseV1 err = util.UnmarshalJSON(f.recorder.Body.Bytes(), &result) if err != nil { t.Fatalf("Unexpected error while unmarshalling result: %v", err) } if !reflect.DeepEqual(result, expected) { t.Fatalf("Expected %v but got: %v", expected, result) } } func TestBadQueryV1(t *testing.T) { f := newFixture(t) expectedErr := `{ "code": "invalid_parameter", "message": "error(s) occurred while parsing query", "errors": [ { "code": "rego_parse_error", "message": "illegal token", "location": { "file": "", "row": 1, "col": 1 }, "details": { "line": "^ -i", "idx": 0 } } ] }` if err := f.v1(http.MethodGet, `/query?q=^ -i`, "", 400, expectedErr); err != nil { recvErr := f.recorder.Body.String() t.Fatalf(`Expected %v but got: %v`, expectedErr, recvErr) } } func TestQueryV1UnsafeBuiltin(t *testing.T) { f := newFixture(t) query := `/query?q=http.send({"method": "get", "url": "foo.com"}, x)` expected := `{ "code": "invalid_parameter", "message": "error(s) occurred while compiling query", "errors": [ { "code": "rego_type_error", "message": "unsafe built-in function calls in expression: http.send", "location": { "file": "", "row": 1, "col": 1 } } ] }` if err := f.v1(http.MethodGet, query, "", 400, expected); err != nil { t.Fatalf(`Expected %v but got: %v`, expected, f.recorder.Body.String()) } } func TestUnversionedPost(t *testing.T) { f := newFixture(t) post := func() *http.Request { return newReqUnversioned(http.MethodPost, "/", ` { "foo": { "bar": [1,2,3] } }`) } f.server.Handler.ServeHTTP(f.recorder, post()) if f.recorder.Code != 404 { t.Fatalf("Expected not found before policy added but got %v", f.recorder) } module := ` package system.main agg = x { sum(input.foo.bar, x) } ` if err := f.v1("PUT", "/policies/test", module, 200, ""); err != nil { t.Fatal(err) } f.reset() f.server.Handler.ServeHTTP(f.recorder, post()) expected := `{"agg":6}` if f.recorder.Code != 200 || f.recorder.Body.String() != expected { t.Fatalf(`Expected HTTP 200 / %v but got: %v`, expected, f.recorder) } } func TestQueryV1Explain(t *testing.T) { f := newFixture(t) get := newReqV1(http.MethodGet, `/query?q=a=[1,2,3]%3Ba[i]=x&explain=full`, "") f.server.Handler.ServeHTTP(f.recorder, get) if f.recorder.Code != 200 { t.Fatalf("Expected 200 but got: %v", f.recorder) } var result types.QueryResponseV1 if err := util.NewJSONDecoder(f.recorder.Body).Decode(&result); err != nil { t.Fatalf("Unexpected JSON decode error: %v", err) } explain := mustUnmarshalTrace(result.Explanation) if len(explain) != 13 { t.Fatalf("Expected exactly 10 trace events for full query but got %d", len(explain)) } } func TestAuthorization(t *testing.T) { ctx := context.Background() store := inmem.New() m, err := plugins.New([]byte{}, "test", store) if err != nil { panic(err) } if err := m.Start(ctx); err != nil { panic(err) } txn := storage.NewTransactionOrDie(ctx, store, storage.WriteParams) authzPolicy := `package system.authz import input.identity default allow = false allow { identity = "bob" } ` if err := store.UpsertPolicy(ctx, txn, "test", []byte(authzPolicy)); err != nil { panic(err) } if err := store.Commit(ctx, txn); err != nil { panic(err) } server, err := New(). WithAddresses([]string{":8182"}). WithStore(store). WithManager(m). WithAuthorization(AuthorizationBasic). Init(ctx) if err != nil { panic(err) } // Test that bob can do stuff. req1, err := http.NewRequest(http.MethodGet, "http://localhost:8182/health", nil) if err != nil { panic(err) } req1 = identifier.SetIdentity(req1, "bob") validateAuthorizedRequest(t, server, req1, http.StatusOK) // Test that alice can't do stuff. req2, err := http.NewRequest(http.MethodGet, "http://localhost:8182/health", nil) if err != nil { panic(err) } req2 = identifier.SetIdentity(req2, "alice") validateAuthorizedRequest(t, server, req2, http.StatusUnauthorized) // Reverse the policy. update := identifier.SetIdentity(newReqV1(http.MethodPut, "/policies/test", ` package system.authz import input.identity default allow = false allow { identity = "alice" } `), "bob") recorder := httptest.NewRecorder() server.Handler.ServeHTTP(recorder, update) if recorder.Code != http.StatusOK { t.Fatalf("Expected policy update to succeed but got: %v", recorder) } // Try alice again. server.Handler.ServeHTTP(recorder, req2) validateAuthorizedRequest(t, server, req2, http.StatusOK) // Try bob again. server.Handler.ServeHTTP(recorder, req1) validateAuthorizedRequest(t, server, req1, http.StatusUnauthorized) } func validateAuthorizedRequest(t *testing.T, s *Server, req *http.Request, exp int) { t.Helper() r := httptest.NewRecorder() // First check the main router s.Handler.ServeHTTP(r, req) if r.Code != exp { t.Fatalf("(Default Handler) Expected %v but got: %v", exp, r) } r = httptest.NewRecorder() // Ensure that auth happens for the diagnostic handler as well s.DiagnosticHandler.ServeHTTP(r, req) if r.Code != exp { t.Fatalf("(Diagnostic Handler) Expected %v but got: %v", exp, r) } } func TestServerReloadTrigger(t *testing.T) { f := newFixture(t) store := f.server.store ctx := context.Background() txn := storage.NewTransactionOrDie(ctx, store, storage.WriteParams) if err := store.UpsertPolicy(ctx, txn, "test", []byte("package test\np = 1")); err != nil { panic(err) } if err := f.v1(http.MethodGet, "/data/test", "", 200, `{}`); err != nil { t.Fatalf("Unexpected error from server: %v", err) } if err := store.Commit(ctx, txn); err != nil { panic(err) } if err := f.v1(http.MethodGet, "/data/test", "", 200, `{"result": {"p": 1}}`); err != nil { t.Fatalf("Unexpected error from server: %v", err) } } func TestServerClearsCompilerConflictCheck(t *testing.T) { f := newFixture(t) store := f.server.store ctx := context.Background() // Make a new transaction params := storage.WriteParams params.Context = storage.NewContext() txn := storage.NewTransactionOrDie(ctx, store, params) // Fresh compiler we will swap on the manager c := ast.NewCompiler() // Add the policy we want to use c.Compile(map[string]*ast.Module{"test": ast.MustParseModule("package test\np=1")}) if len(c.Errors) > 0 { t.Fatalf("Unexpected compile errors: %v", c.Errors) } // Add in a "bad" conflict check c = c.WithPathConflictsCheck(func(_ []string) (bool, error) { t.Fatal("Conflict check should not have been called") return false, nil }) // Set the compiler on the transaction context and commit to trigger listeners plugins.SetCompilerOnContext(params.Context, c) if err := store.UpsertPolicy(ctx, txn, "test", []byte("package test\np = 1")); err != nil { panic(err) } if err := store.Commit(ctx, txn); err != nil { panic(err) } // internal helpers should now give the new compiler back if f.server.getCompiler() != c { t.Fatalf("Expected to get the updated compiler") } // If we request for partial evaluation it will end up using the compiler set from the manager. Ensure it // is using a correct conflict checker. if err := f.v1(http.MethodGet, "/data/test?partial", "", 200, `{"result": {"p": 1}}`); err != nil { t.Fatalf("Unexpected error from server: %v", err) } } type queryBindingErrStore struct { storage.WritesNotSupported storage.PolicyNotSupported } func (s *queryBindingErrStore) Read(ctx context.Context, txn storage.Transaction, path storage.Path) (interface{}, error) { return nil, fmt.Errorf("expected error") } func (*queryBindingErrStore) ListPolicies(ctx context.Context, txn storage.Transaction) ([]string, error) { return nil, nil } func (queryBindingErrStore) NewTransaction(ctx context.Context, params ...storage.TransactionParams) (storage.Transaction, error) { return nil, nil } func (queryBindingErrStore) Commit(ctx context.Context, txn storage.Transaction) error { return nil } func (queryBindingErrStore) Abort(ctx context.Context, txn storage.Transaction) { } func (queryBindingErrStore) Register(context.Context, storage.Transaction, storage.TriggerConfig) (storage.TriggerHandle, error) { return nil, nil } func (queryBindingErrStore) Unregister(context.Context, storage.Transaction, string) { } func TestQueryBindingIterationError(t *testing.T) { ctx := context.Background() mock := &queryBindingErrStore{} m, err := plugins.New([]byte{}, "test", mock) if err != nil { panic(err) } server, err := New().WithStore(mock).WithManager(m).WithAddresses([]string{":8182"}).Init(ctx) if err != nil { panic(err) } recorder := httptest.NewRecorder() f := &fixture{ server: server, recorder: recorder, t: t, } get := newReqV1(http.MethodGet, `/query?q=a=data.foo.bar`, "") f.server.Handler.ServeHTTP(f.recorder, get) if f.recorder.Code != 500 { t.Fatalf("Expected 500 error due to unknown storage error but got: %v", f.recorder) } var resultErr types.ErrorV1 if jsonErr := json.NewDecoder(f.recorder.Body).Decode(&resultErr); jsonErr != nil { t.Fatal(jsonErr) } if resultErr.Code != types.CodeInternal || resultErr.Message != "expected error" { t.Fatal("unexpected response:", resultErr) } } const ( testMod = `package a.b.c import data.x.y as z import data.p q[x] { p[x]; not r[x] } r[x] { z[x] = 4 }` ) type fixture struct { server *Server recorder *httptest.ResponseRecorder t *testing.T } func newFixture(t *testing.T, opts ...func(*Server)) *fixture { ctx := context.Background() store := inmem.New() m, err := plugins.New([]byte{}, "test", store) if err != nil { panic(err) } if err := m.Start(ctx); err != nil { panic(err) } server := New(). WithAddresses([]string{"localhost:8182"}). WithStore(store). WithManager(m) for _, opt := range opts { opt(server) } server, err = server.Init(ctx) if err != nil { panic(err) } recorder := httptest.NewRecorder() return &fixture{ server: server, recorder: recorder, t: t, } } func (f *fixture) loadResponse() interface{} { var v interface{} err := util.NewJSONDecoder(f.recorder.Body).Decode(&v) if err != nil { panic(err) } return v } func (f *fixture) v1TestRequests(trs []tr) error { for i, tr := range trs { if err := f.v1(tr.method, tr.path, tr.body, tr.code, tr.resp); err != nil { return errors.Wrapf(err, "error on test request #%d", i+1) } } return nil } func (f *fixture) v1(method string, path string, body string, code int, resp string) error { // All v1 API's should 404 for the diagnostic handler if err := f.executeDiagnosticRequest(newReqV1(method, path, body), 404, ""); err != nil { return err } return f.executeRequest(newReqV1(method, path, body), code, resp) } func (f *fixture) v0(method string, path string, body string, code int, resp string) error { // All v0 API's should 404 for the diagnostic handler if err := f.executeDiagnosticRequest(newReqV0(method, path, body), 404, ""); err != nil { return err } return f.executeRequest(newReqV0(method, path, body), code, resp) } func (f *fixture) executeRequestForHandler(h http.Handler, req *http.Request, code int, resp string) error { f.reset() h.ServeHTTP(f.recorder, req) if f.recorder.Code != code { return fmt.Errorf("Expected code %v from %v %v but got: %+v", code, req.Method, req.URL, f.recorder) } if resp != "" { var result interface{} if err := util.UnmarshalJSON([]byte(f.recorder.Body.String()), &result); err != nil { return fmt.Errorf("Expected JSON response from %v %v but got: %v", req.Method, req.URL, f.recorder) } var expected interface{} if err := util.UnmarshalJSON([]byte(resp), &expected); err != nil { panic(err) } if !reflect.DeepEqual(result, expected) { a, err := json.MarshalIndent(expected, "", " ") if err != nil { panic(err) } b, err := json.MarshalIndent(result, "", " ") if err != nil { panic(err) } return fmt.Errorf("Expected JSON response from %v %v to equal:\n\n%s\n\nGot:\n\n%s", req.Method, req.URL, a, b) } } return nil } func (f *fixture) executeRequest(req *http.Request, code int, resp string) error { return f.executeRequestForHandler(f.server.Handler, req, code, resp) } func (f *fixture) executeDiagnosticRequest(req *http.Request, code int, resp string) error { return f.executeRequestForHandler(f.server.DiagnosticHandler, req, code, resp) } func (f *fixture) reset() { f.recorder = httptest.NewRecorder() } func executeRequests(t *testing.T, reqs []tr) { t.Helper() f := newFixture(t) for i, req := range reqs { if err := f.v1(req.method, req.path, req.body, req.code, req.resp); err != nil { t.Errorf("Unexpected response on request %d: %v", i+1, err) } } } // Runs through an array of test cases against the v0 REST API tree func executeRequestsv0(t *testing.T, reqs []tr) { t.Helper() f := newFixture(t) for i, req := range reqs { if err := f.v0(req.method, req.path, req.body, req.code, req.resp); err != nil { t.Errorf("Unexpected response on request %d: %v", i+1, err) } } } func validateDiagnosticRequest(t *testing.T, f *fixture, req *http.Request, code int, resp string) { t.Helper() // diagnostic requests need to be available on both the normal handler and diagnostic handler if err := f.executeRequest(req, code, resp); err != nil { t.Errorf("Unexpected error for request %v: %s", req, err) } if err := f.executeDiagnosticRequest(req, code, resp); err != nil { t.Errorf("Unexpected error for request %v: %s", req, err) } } func newPolicy(id, s string) types.PolicyV1 { compiler := ast.NewCompiler() parsed := ast.MustParseModule(s) if compiler.Compile(map[string]*ast.Module{"": parsed}); compiler.Failed() { panic(compiler.Errors) } mod := compiler.Modules[""] return types.PolicyV1{ID: id, AST: mod, Raw: s} } func newReqV1(method string, path string, body string) *http.Request { return newReq(1, method, path, body) } func newReqV0(method string, path string, body string) *http.Request { return newReq(0, method, path, body) } func newReq(version int, method, path, body string) *http.Request { return newReqUnversioned(method, fmt.Sprintf("/v%d", version)+path, body) } func newReqUnversioned(method, path, body string) *http.Request { req, err := http.NewRequest(method, path, strings.NewReader(body)) if err != nil { panic(err) } return req } func mustUnmarshalTrace(t types.TraceV1) (trace types.TraceV1Raw) { if err := json.Unmarshal(t, &trace); err != nil { panic("not reached") } return trace } func TestShutdown(t *testing.T) { f := newFixture(t, func(s *Server) { s.WithDiagnosticAddresses([]string{":8443"}) }) loops, err := f.server.Listeners() if err != nil { t.Errorf("unexpected error: %s", err.Error()) } errc := make(chan error) for _, loop := range loops { go func(serverLoop func() error) { errc <- serverLoop() }(loop) } ctx, cancel := context.WithTimeout(context.Background(), time.Duration(5)*time.Second) defer cancel() err = f.server.Shutdown(ctx) if err != nil { t.Errorf("unexpected error shutting down server: %s", err.Error()) } } func TestShutdownError(t *testing.T) { f := newFixture(t, func(s *Server) { s.WithDiagnosticAddresses([]string{":8443"}) }) errMsg := "failed to shutdown" // Add a mock httpListener to the server m := &mockHTTPListener{ shutdownHook: func() error { return errors.New(errMsg) }, } f.server.httpListeners = []httpListener{m} ctx, cancel := context.WithTimeout(context.Background(), time.Duration(5)*time.Second) defer cancel() err := f.server.Shutdown(ctx) if err == nil { t.Error("expected an error shutting down server but err==nil") } else if !strings.Contains(err.Error(), errMsg) { t.Errorf("unexpected error shutting down server: %s", err.Error()) } } func TestShutdownMultipleErrors(t *testing.T) { f := newFixture(t, func(s *Server) { s.WithDiagnosticAddresses([]string{":8443"}) }) shutdownErrs := []error{errors.New("err1"), nil, errors.New("err3")} // Add mock httpListeners to the server for _, err := range shutdownErrs { m := &mockHTTPListener{} if err != nil { retVal := errors.New(err.Error()) m.shutdownHook = func() error { return retVal } } f.server.httpListeners = append(f.server.httpListeners, m) } ctx, cancel := context.WithTimeout(context.Background(), time.Duration(5)*time.Second) defer cancel() err := f.server.Shutdown(ctx) if err == nil { t.Fatal("expected an error shutting down server but err==nil") } for _, expectedErr := range shutdownErrs { if expectedErr != nil && !strings.Contains(err.Error(), expectedErr.Error()) { t.Errorf("expected error message to contain '%s', full message: '%s'", expectedErr.Error(), err.Error()) } } } func TestAddrsNoListeners(t *testing.T) { s := New() a := s.Addrs() if len(a) != 0 { t.Errorf("expected an empty list of addresses, got: %+v", a) } } func TestAddrsWithEmptyListenAddr(t *testing.T) { s := New() s.httpListeners = []httpListener{&mockHTTPListener{}} a := s.Addrs() if len(a) != 0 { t.Errorf("expected an empty list of addresses, got: %+v", a) } } func TestAddrsWithListenAddr(t *testing.T) { s := New() s.httpListeners = []httpListener{&mockHTTPListener{addrs: ":8181"}} a := s.Addrs() if len(a) != 1 || a[0] != ":8181" { t.Errorf("expected only an ':8181' address, got: %+v", a) } } func TestAddrsWithMixedListenerAddr(t *testing.T) { s := New() addrs := []string{":8181", "", "unix:///var/tmp/foo.sock"} expected := []string{":8181", "unix:///var/tmp/foo.sock"} s.httpListeners = []httpListener{} for _, addr := range addrs { s.httpListeners = append(s.httpListeners, &mockHTTPListener{addrs: addr, t: defaultListenerType}) } a := s.Addrs() if len(a) != 2 { t.Errorf("expected 2 addresses, got: %+v", a) } for _, expectedAddr := range expected { found := false for _, actualAddr := range a { if expectedAddr == actualAddr { found = true break } } if !found { t.Errorf("expected %q in address list, got: %+v", expectedAddr, a) } } } func TestDiagnosticAddrsNoListeners(t *testing.T) { s := New() a := s.DiagnosticAddrs() if len(a) != 0 { t.Errorf("expected an empty list of addresses, got: %+v", a) } } func TestDiagnosticAddrsWithEmptyListenAddr(t *testing.T) { s := New() s.httpListeners = []httpListener{&mockHTTPListener{t: diagnosticListenerType}} a := s.DiagnosticAddrs() if len(a) != 0 { t.Errorf("expected an empty list of addresses, got: %+v", a) } } func TestDiagnosticAddrsWithListenAddr(t *testing.T) { s := New() s.httpListeners = []httpListener{&mockHTTPListener{addrs: ":8181", t: diagnosticListenerType}} a := s.DiagnosticAddrs() if len(a) != 1 || a[0] != ":8181" { t.Errorf("expected only an ':8181' address, got: %+v", a) } } func TestDiagnosticAddrsWithMixedListenerAddr(t *testing.T) { s := New() addrs := []string{":8181", "", "unix:///var/tmp/foo.sock"} expected := []string{":8181", "unix:///var/tmp/foo.sock"} s.httpListeners = []httpListener{} for _, addr := range addrs { s.httpListeners = append(s.httpListeners, &mockHTTPListener{addrs: addr, t: diagnosticListenerType}) } a := s.DiagnosticAddrs() if len(a) != 2 { t.Errorf("expected 2 addresses, got: %+v", a) } for _, expectedAddr := range expected { found := false for _, actualAddr := range a { if expectedAddr == actualAddr { found = true break } } if !found { t.Errorf("expected %q in address list, got: %+v", expectedAddr, a) } } } func TestMixedAddrTypes(t *testing.T) { s := New() s.httpListeners = []httpListener{} addrs := map[string]struct{}{"localhost:8181": {}, "localhost:1234": {}, "unix:///var/tmp/foo.sock": {}} for addr := range addrs { s.httpListeners = append(s.httpListeners, &mockHTTPListener{addrs: addr, t: defaultListenerType}) } diagAddrs := map[string]struct{}{":8181": {}, "https://127.0.0.1": {}} for addr := range diagAddrs { s.httpListeners = append(s.httpListeners, &mockHTTPListener{addrs: addr, t: diagnosticListenerType}) } actualAddrs := s.Addrs() if len(actualAddrs) != len(addrs) { t.Errorf("expected %d addresses, got: %+v", len(addrs), actualAddrs) } for _, addr := range actualAddrs { if _, ok := addrs[addr]; !ok { t.Errorf("Unexpected address %v", addr) } } actualDiagAddrs := s.DiagnosticAddrs() if len(actualDiagAddrs) != len(diagAddrs) { t.Errorf("expected %d addresses, got: %+v", len(diagAddrs), actualDiagAddrs) } for _, addr := range actualDiagAddrs { if _, ok := diagAddrs[addr]; !ok { t.Errorf("Unexpected diagnostic address %v", addr) } } } func TestDiagnosticRoutes(t *testing.T) { cases := []struct { path string should404 bool }{ {"/health", false}, {"/metrics", false}, {"/debug/pprof/", true}, {"/v0/data", true}, {"/v0/data/foo", true}, {"/v1/data/", true}, {"/v1/data/foo", true}, {"/v1/policies", true}, {"/v1/policies/foo", true}, {"/v1/query", true}, {"/v1/compile", true}, {"/", true}, } f := newFixture(t, func(s *Server) { s.WithPprofEnabled(true) s.WithMetrics(new(mockMetricsProvider)) }) for _, tc := range cases { t.Run(tc.path, func(t *testing.T) { req, err := http.NewRequest("GET", tc.path, nil) if err != nil { t.Fatalf("Unexpected error: %s", err) } code := http.StatusOK if tc.should404 { code = http.StatusNotFound } f.reset() f.server.DiagnosticHandler.ServeHTTP(f.recorder, req) if f.recorder.Code != code { t.Errorf("Expected code %v from %v %v but got: %+v", code, req.Method, req.URL, f.recorder) } }) } } type mockHTTPHandler struct{} func (m *mockHTTPHandler) ServeHTTP(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) } type mockMetricsProvider struct{} func (m *mockMetricsProvider) RegisterEndpoints(registrar func(path, method string, handler http.Handler)) { registrar("/metrics", "GET", new(mockHTTPHandler)) } func (m *mockMetricsProvider) InstrumentHandler(handler http.Handler, label string) http.Handler { return handler } type listenerHook func() error type mockHTTPListener struct { shutdownHook listenerHook addrs string t httpListenerType } var _ httpListener = (*mockHTTPListener)(nil) func (m mockHTTPListener) Addr() string { return m.addrs } func (m mockHTTPListener) ListenAndServe() error { return errors.New("not implemented") } func (m mockHTTPListener) ListenAndServeTLS(certFile, keyFile string) error { return errors.New("not implemented") } func (m mockHTTPListener) Shutdown(ctx context.Context) error { var err error if m.shutdownHook != nil { err = m.shutdownHook() } return err } func (m mockHTTPListener) Type() httpListenerType { return m.t }