From 12dad370d5b42e1e7d783d3b3382180841090eca Mon Sep 17 00:00:00 2001 From: Torin Sandall Date: Wed, 11 May 2016 09:08:56 -0700 Subject: [PATCH] Add basic REST API support to server mode - REST APIs * CRUDL on policy modules * Ad-hoc queries * Query and patch base documents * Query virtual documents - Add PolicyStore to manage policy definition/module CRUDL operations. * Supports persistence of policy definitons. * Serve REST API CRUDL operations. * Manage install/uninstall of rules into data store. * Manage persistence of policy definitions. - Misc. refactoring * Move storage creation into runtime Init. * Make AST types JSON serializable. Tweaked ast.Import to use Term instead of Value for the path. --- .gitignore | 8 +- Makefile | 7 +- ast/compile.go | 4 +- ast/parser_test.go | 8 +- ast/policy.go | 70 ++++- ast/policy_test.go | 95 ++++++- ast/rego.peg | 4 +- ast/term.go | 161 +++++++++++- ast/term_test.go | 41 ++- cmd/run.go | 36 ++- cmd/version.go | 9 +- eval/index.go | 14 +- eval/policystore.go | 321 +++++++++++++++++++++++ eval/policystore_test.go | 305 ++++++++++++++++++++++ eval/storage.go | 131 ++-------- eval/storage_test.go | 58 ----- eval/topdown.go | 21 +- eval/topdown_test.go | 40 ++- eval/tracer_test.go | 13 +- runtime/repl.go | 21 +- runtime/runtime.go | 199 +++++++++++++-- runtime/runtime_test.go | 125 +++++++++ runtime/server.go | 538 +++++++++++++++++++++++++++++++++++++++ runtime/server_test.go | 360 ++++++++++++++++++++++++++ 24 files changed, 2305 insertions(+), 284 deletions(-) create mode 100644 eval/policystore.go create mode 100644 eval/policystore_test.go create mode 100644 runtime/runtime_test.go create mode 100644 runtime/server.go create mode 100644 runtime/server_test.go diff --git a/.gitignore b/.gitignore index cd33ccefa2..9e3f5b8090 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,11 @@ +# development environment .DS_Store .vscode + +# build artifacts +coverage opa ast/parser.go -coverage + +# runtime artifacts +policies diff --git a/Makefile b/Makefile index ad6f9ed650..06ca9d5d9d 100644 --- a/Makefile +++ b/Makefile @@ -37,14 +37,17 @@ generate: build: generate $(GO) build -o opa $(LDFLAGS) +install: generate + $(GO) install $(LDFLAGS) + test: generate $(GO) test -v $(PACKAGES) COVER_PACKAGES=$(PACKAGES) $(COVER_PACKAGES): @mkdir -p coverage/$(shell dirname $@) - go test -covermode=count -coverprofile=coverage/$(shell dirname $@)/coverage.out $@ - go tool cover -html=coverage/$(shell dirname $@)/coverage.out || true + $(GO) test -covermode=count -coverprofile=coverage/$(shell dirname $@)/coverage.out $@ + $(GO) tool cover -html=coverage/$(shell dirname $@)/coverage.out || true cover: $(COVER_PACKAGES) diff --git a/ast/compile.go b/ast/compile.go index 9ce0aed353..fcc6803fbe 100644 --- a/ast/compile.go +++ b/ast/compile.go @@ -196,7 +196,7 @@ func (c *Compiler) setGlobals() { // Populate globals with imports within this module. for _, i := range m.Imports { if len(i.Alias) > 0 { - switch p := i.Path.(type) { + switch p := i.Path.Value.(type) { case Ref: globals[i.Alias] = p case Var: @@ -205,7 +205,7 @@ func (c *Compiler) setGlobals() { c.err("unexpected %T: %v", p, i) } } else { - switch p := i.Path.(type) { + switch p := i.Path.Value.(type) { case Ref: switch v := p[len(p)-1].Value.(type) { case String: diff --git a/ast/parser_test.go b/ast/parser_test.go index 6dfe37e250..cbb24a844e 100644 --- a/ast/parser_test.go +++ b/ast/parser_test.go @@ -200,12 +200,12 @@ func TestPackage(t *testing.T) { } func TestImport(t *testing.T) { - assertParseImport(t, "single", "import foo", &Import{Path: VarTerm("foo").Value}) - ref := RefTerm(VarTerm("foo"), StringTerm("bar"), StringTerm("baz")).Value + assertParseImport(t, "single", "import foo", &Import{Path: VarTerm("foo")}) + ref := RefTerm(VarTerm("foo"), StringTerm("bar"), StringTerm("baz")) assertParseImport(t, "multiple", "import foo.bar.baz", &Import{Path: ref}) - assertParseImport(t, "single alias", "import foo as bar", &Import{Path: VarTerm("foo").Value, Alias: Var("bar")}) + assertParseImport(t, "single alias", "import foo as bar", &Import{Path: VarTerm("foo"), Alias: Var("bar")}) assertParseImport(t, "multiple alias", "import foo.bar.baz as qux", &Import{Path: ref, Alias: Var("qux")}) - ref2 := RefTerm(VarTerm("foo"), StringTerm("bar"), StringTerm("white space")).Value + ref2 := RefTerm(VarTerm("foo"), StringTerm("bar"), StringTerm("white space")) assertParseImport(t, "white space", "import foo.bar[\"white space\"]", &Import{Path: ref2}) assertParseError(t, "non-ground ref", "import foo[x]") } diff --git a/ast/policy.go b/ast/policy.go index 4c58fac2e1..84b939e7db 100644 --- a/ast/policy.go +++ b/ast/policy.go @@ -4,8 +4,11 @@ package ast -import "fmt" -import "strings" +import ( + "encoding/json" + "fmt" + "strings" +) // DefaultRootDocument is the default root document. // All package directives inside source files are implicitly @@ -31,25 +34,25 @@ type ( // Package represents the namespace of the documents produced // by rules inside the module. Package struct { - Location *Location + Location *Location `json:"-"` Path Ref } // Import represents a dependency on a document outside of the policy // namespace. Imports are optional. Import struct { - Location *Location - Path Value - Alias Var + Location *Location `json:"-"` + Path *Term + Alias Var `json:",omitempty"` } // Rule represents a rule as defined in the language. Rules define the // content of documents that represent policy decisions. Rule struct { - Location *Location + Location *Location `json:"-"` Name Var - Key *Term - Value *Term + Key *Term `json:",omitempty"` + Value *Term `json:",omitempty"` Body Body } @@ -58,8 +61,8 @@ type ( // Expr represents a single expression contained inside the body of a rule. Expr struct { - Location *Location - Negated bool + Location *Location `json:"-"` + Negated bool `json:",omitempty"` Terms interface{} } ) @@ -264,6 +267,51 @@ func (expr *Expr) String() string { return strings.Join(buf, " ") } +// UnmarshalJSON parses the byte array and stores the result in expr. +func (expr *Expr) UnmarshalJSON(bs []byte) error { + v := map[string]interface{}{} + if err := json.Unmarshal(bs, &v); err != nil { + return err + } + + n, ok := v["Negated"] + if !ok { + expr.Negated = false + } else { + b, ok := n.(bool) + if !ok { + return unmarshalError(n, "bool") + } + expr.Negated = b + } + + switch ts := v["Terms"].(type) { + case map[string]interface{}: + v, err := unmarshalValue(ts) + if err != nil { + return err + } + expr.Terms = &Term{Value: v} + case []interface{}: + buf := []*Term{} + for _, v := range ts { + e, ok := v.(map[string]interface{}) + if !ok { + return unmarshalError(v, "map[string]interface{}") + } + v, err := unmarshalValue(e) + if err != nil { + return err + } + buf = append(buf, &Term{Value: v}) + } + expr.Terms = buf + default: + return unmarshalError(v["Terms"], "Term or []Term") + } + return nil +} + // NewBuiltinExpr creates a new Expr object with the supplied terms. // The builtin operator must be the first term. func NewBuiltinExpr(terms ...*Term) *Expr { diff --git a/ast/policy_test.go b/ast/policy_test.go index 3879a6c0f3..d936f23040 100644 --- a/ast/policy_test.go +++ b/ast/policy_test.go @@ -4,7 +4,38 @@ package ast -import "testing" +import ( + "encoding/json" + "reflect" + "testing" +) + +func TestModuleJSONRoundTrip(t *testing.T) { + mod := MustParseModule(` + package a.b.c + import data.x.y as z + import data.u.i + p = [1,2,{"foo":3}] :- r[x] = 1, not q[x] + r[y] = v :- i[1] = y, v = i[2] + q[x] :- a=[true,false,null,{"x":[1,2,3]}], a[i] = x + `) + + bs, err := json.Marshal(mod) + if err != nil { + panic(err) + } + + roundtrip := &Module{} + + err = json.Unmarshal(bs, roundtrip) + if err != nil { + panic(err) + } + + if !roundtrip.Equal(mod) { + t.Errorf("Expected roundtripped module to be equal to original:\nExpected:\n\n%v\n\nGot:\n\n%v\n", mod, roundtrip) + } +} func TestPackageEquals(t *testing.T) { pkg1 := &Package{Path: RefTerm(VarTerm("foo"), StringTerm("bar"), StringTerm("baz")).Value.(Ref)} @@ -26,12 +57,12 @@ func TestPackageString(t *testing.T) { } func TestImportEquals(t *testing.T) { - imp1 := &Import{Path: Var("foo"), Alias: Var("bar")} - imp11 := &Import{Path: Var("foo"), Alias: Var("bar")} - imp2 := &Import{Path: Var("foo")} - imp3 := &Import{Path: RefTerm(VarTerm("bar"), VarTerm("baz"), VarTerm("qux")).Value, Alias: Var("corge")} - imp33 := &Import{Path: RefTerm(VarTerm("bar"), VarTerm("baz"), VarTerm("qux")).Value, Alias: Var("corge")} - imp4 := &Import{Path: RefTerm(VarTerm("bar"), VarTerm("baz"), VarTerm("qux")).Value} + imp1 := &Import{Path: VarTerm("foo"), Alias: Var("bar")} + imp11 := &Import{Path: VarTerm("foo"), Alias: Var("bar")} + imp2 := &Import{Path: VarTerm("foo")} + imp3 := &Import{Path: RefTerm(VarTerm("bar"), VarTerm("baz"), VarTerm("qux")), Alias: Var("corge")} + imp33 := &Import{Path: RefTerm(VarTerm("bar"), VarTerm("baz"), VarTerm("qux")), Alias: Var("corge")} + imp4 := &Import{Path: RefTerm(VarTerm("bar"), VarTerm("baz"), VarTerm("qux"))} assertImportsEqual(t, imp1, imp1) assertImportsEqual(t, imp1, imp11) assertImportsEqual(t, imp3, imp3) @@ -47,10 +78,10 @@ func TestImportEquals(t *testing.T) { } func TestImportString(t *testing.T) { - imp1 := &Import{Path: Var("foo"), Alias: Var("bar")} - imp2 := &Import{Path: Var("foo")} - imp3 := &Import{Path: RefTerm(VarTerm("bar"), StringTerm("baz"), StringTerm("qux")).Value, Alias: Var("corge")} - imp4 := &Import{Path: RefTerm(VarTerm("bar"), StringTerm("baz"), StringTerm("qux")).Value} + imp1 := &Import{Path: VarTerm("foo"), Alias: Var("bar")} + imp2 := &Import{Path: VarTerm("foo")} + imp3 := &Import{Path: RefTerm(VarTerm("bar"), StringTerm("baz"), StringTerm("qux")), Alias: Var("corge")} + imp4 := &Import{Path: RefTerm(VarTerm("bar"), StringTerm("baz"), StringTerm("qux"))} assertImportToString(t, imp1, "import foo as bar") assertImportToString(t, imp2, "import foo") assertImportToString(t, imp3, "import bar.baz.qux as corge") @@ -126,6 +157,48 @@ func TextExprString(t *testing.T) { assertExprString(t, expr4, "ne({foo: [1, a.b]}, false)") } +func TestExprBadJSON(t *testing.T) { + + assert := func(js string, exp error) { + expr := Expr{} + err := json.Unmarshal([]byte(js), &expr) + if !reflect.DeepEqual(exp, err) { + t.Errorf("Expected %v but got: %v", exp, err) + } + } + + js := ` + { + "Negated": 100, + "Terms": { + "Value": "foo", + "Type": "string" + } + } + ` + + exp := unmarshalError(100.0, "bool") + assert(js, exp) + + js = ` + { + "Terms": [ + "foo" + ] + } + ` + exp = unmarshalError("foo", "map[string]interface{}") + assert(js, exp) + + js = ` + { + "Terms": "bad value" + } + ` + exp = unmarshalError("bad value", "Term or []Term") + assert(js, exp) +} + func TestRuleHeadEquals(t *testing.T) { assertRulesEqual(t, &Rule{}, &Rule{}) diff --git a/ast/rego.peg b/ast/rego.peg index 8b432307de..8bba07917b 100644 --- a/ast/rego.peg +++ b/ast/rego.peg @@ -74,8 +74,8 @@ Package <- "package" ws val:(Ref / Var) { Import <- "import" ws path:(Ref / Var) alias:(ws "as" ws Var)? { imp := &Import{} imp.Location = currentLocation(c) - imp.Path = path.(*Term).Value - switch p := imp.Path.(type) { + imp.Path = path.(*Term) + switch p := imp.Path.Value.(type) { case Ref: if !p.IsGround() { return nil, fmt.Errorf("import cannot contain variables in tail: %v", p) diff --git a/ast/term.go b/ast/term.go index e7c719e3a6..6fb07406f9 100644 --- a/ast/term.go +++ b/ast/term.go @@ -4,10 +4,14 @@ package ast -import "fmt" -import "regexp" -import "strconv" -import "strings" +import ( + "encoding/json" + "fmt" + "regexp" + "strconv" + "strings" +) + import "hash/fnv" // Location records a position in source code @@ -48,7 +52,7 @@ type Value interface { // Term is an argument to a function. type Term struct { Value Value // the value of the Term as represented in Go - Location *Location // the location of the Term in the source + Location *Location `json:"-"` // the location of the Term in the source } // Equal returns true if this term equals the other term. Equality is @@ -71,10 +75,55 @@ func (term *Term) IsGround() bool { return term.Value.IsGround() } +// MarshalJSON returns the JSON encoding of the term. +// Specialized marshalling logic is required to include a type hint +// for Value. +func (term *Term) MarshalJSON() ([]byte, error) { + var typ string + switch term.Value.(type) { + case Null: + typ = "null" + case Boolean: + typ = "boolean" + case Number: + typ = "number" + case String: + typ = "string" + case Ref: + typ = "ref" + case Var: + typ = "var" + case Array: + typ = "array" + case Object: + typ = "object" + } + d := map[string]interface{}{ + "Type": typ, + "Value": term.Value, + } + return json.Marshal(d) +} + func (term *Term) String() string { return term.Value.String() } +// UnmarshalJSON parses the byte array and stores the result in term. +// Specialized unmarshalling is required to handle Value. +func (term *Term) UnmarshalJSON(bs []byte) error { + v := map[string]interface{}{} + if err := json.Unmarshal(bs, &v); err != nil { + return err + } + val, err := unmarshalValue(v) + if err != nil { + return err + } + term.Value = val + return nil +} + // Null represents the null value defined by JSON. type Null struct{} @@ -590,3 +639,105 @@ func termSliceIsGround(a []*Term) bool { } return true } + +func unmarshalError(v interface{}, e string) error { + return fmt.Errorf("ast: cannot unmarshal %T into Go value of type %v", v, e) +} + +func unmarshalTermSlice(d map[string]interface{}) ([]*Term, error) { + s, ok := d["Value"].([]interface{}) + if !ok { + return nil, unmarshalError(d["Value"], "[]interface{}") + } + buf := []*Term{} + for _, i := range s { + m, ok := i.(map[string]interface{}) + if !ok { + return nil, unmarshalError(i, "map[string]interface{}") + } + v, err := unmarshalValue(m) + if err != nil { + return nil, err + } + buf = append(buf, &Term{Value: v}) + } + return buf, nil +} + +func unmarshalValue(d map[string]interface{}) (Value, error) { + switch d["Type"] { + case "null": + return Null{}, nil + case "boolean": + b, ok := d["Value"].(bool) + if !ok { + return nil, unmarshalError(d["Value"], "bool") + } + return Boolean(b), nil + case "number": + f, ok := d["Value"].(float64) + if !ok { + return nil, unmarshalError(d["Value"], "float64") + } + return Number(f), nil + case "string": + s, ok := d["Value"].(string) + if !ok { + return nil, unmarshalError(d["Value"], "string") + } + return String(s), nil + case "ref": + s, err := unmarshalTermSlice(d) + if err != nil { + return nil, err + } + return Ref(s), nil + case "var": + s, ok := d["Value"].(string) + if !ok { + return nil, unmarshalError(d["Value"], "ast.Var") + } + return Var(s), nil + case "array": + s, err := unmarshalTermSlice(d) + if err != nil { + return nil, err + } + return Array(s), nil + case "object": + buf := Object{} + s, ok := d["Value"].([]interface{}) + if !ok { + return nil, unmarshalError(d["Value"], "[]interface{}") + } + for _, i := range s { + p, ok := i.([]interface{}) + if !ok { + return nil, unmarshalError(i, "[]interface{}") + } + if len(p) != 2 { + return nil, unmarshalError(p, "[2]interface{}") + } + km, ok := p[0].(map[string]interface{}) + if !ok { + return nil, unmarshalError(p[0], "map[string]interface{}") + } + k, err := unmarshalValue(km) + if err != nil { + return nil, err + } + vm, ok := p[1].(map[string]interface{}) + if !ok { + return nil, unmarshalError(p[1], "map[string]interface{}") + } + v, err := unmarshalValue(vm) + if err != nil { + return nil, err + } + buf = append(buf, [2]*Term{&Term{Value: k}, &Term{Value: v}}) + } + return buf, nil + default: + return nil, fmt.Errorf("ast: cannot unmarshal Term with Type %v", d["Type"]) + } +} diff --git a/ast/term_test.go b/ast/term_test.go index e6b6a4641d..a6a697d51f 100644 --- a/ast/term_test.go +++ b/ast/term_test.go @@ -5,6 +5,7 @@ package ast import ( + "encoding/json" "fmt" "reflect" "sort" @@ -106,7 +107,43 @@ func TestQuery(t *testing.T) { } } -func TestEqualTerms(t *testing.T) { +func TestTermBadJSON(t *testing.T) { + + assert := func(js string, exp error) { + term := Term{} + err := json.Unmarshal([]byte(js), &term) + if !reflect.DeepEqual(exp, err) { + t.Errorf("Expected %v but got: %v", exp, err) + } + } + + castTests := []struct { + input string + val interface{} + expected string + }{ + {`{"Value": null, "Type": "boolean"}`, nil, "bool"}, + {`{"Value": false, "Type": "number"}`, false, "float64"}, + {`{"Value": 100, "Type": "string"}`, 100.0, "string"}, + {`{"Value": "hello", "Type": "number"}`, "hello", "float64"}, + {`{"Value": 100, "Type": "var"}`, 100.0, "ast.Var"}, + {`{"Value": "abc", "Type": "ref"}`, "abc", "[]interface{}"}, + {`{"Value": ["abc"], "Type": "ref"}`, "abc", "map[string]interface{}"}, + {`{"Value": "abc", "Type": "array"}`, "abc", "[]interface{}"}, + {`{"Value": ["abc"], "Type": "array"}`, "abc", "map[string]interface{}"}, + {`{"Value": "abc", "Type": "object"}`, "abc", "[]interface{}"}, + {`{"Value": ["abc"], "Type": "object"}`, "abc", "[]interface{}"}, + {`{"Value": [["abc"]], "Type": "object"}`, []interface{}{}, "[2]interface{}"}, + {`{"Value": [["abc", "abc"]], "Type": "object"}`, "abc", "map[string]interface{}"}, + {`{"Value": [[{"Value": "abc", "Type": "string"}, "abc"]], "Type": "object"}`, "abc", "map[string]interface{}"}, + } + + for _, tc := range castTests { + assert(tc.input, unmarshalError(tc.val, tc.expected)) + } +} + +func TestTermEqual(t *testing.T) { assertTermEqual(t, NullTerm(), NullTerm()) assertTermEqual(t, BooleanTerm(true), BooleanTerm(true)) assertTermEqual(t, NumberTerm(5), NumberTerm(5)) @@ -155,7 +192,7 @@ func TestHash(t *testing.T) { } } -func TestTermsToString(t *testing.T) { +func TestTermString(t *testing.T) { assertToString(t, Null{}, "null") assertToString(t, Boolean(true), "true") assertToString(t, Boolean(false), "false") diff --git a/cmd/run.go b/cmd/run.go index bc13bca4fb..00fd52addd 100644 --- a/cmd/run.go +++ b/cmd/run.go @@ -7,6 +7,7 @@ package cmd import ( "os" "path" + "path/filepath" "github.com/open-policy-agent/opa/runtime" "github.com/spf13/cobra" @@ -15,6 +16,12 @@ import ( // default filename for the interactive shell's history var defaultHistoryFile = ".opa_history" +// default policy definition storage directory +var defaultPolicyDir = "policies" + +// default listening address for the server +var defaultAddr = ":8181" + func init() { params := &runtime.Params{} @@ -24,6 +31,18 @@ func init() { Short: "Start OPA in interative or server mode", Long: `Start an instance of the Open Policy Agent (OPA). +To run the interactive shell: + + $ opa run + +To run the server without saving policies: + + $ opa run -s + +To run the server and persist policies to a local directory: + + $ opa run -s -p ./policies/ + The 'run' command starts an instance of the OPA runtime. The OPA runtime can be started as an interactive shell or a server. @@ -33,6 +52,11 @@ a server, users can access OPA's APIs via HTTP. The runtime can be initialized with one or more files that represent base documents (e.g., example.json) or policies (e.g., example.rego). + +If the --policy-dir option is specified any files inside the directory +will be considered policy definitions and will be loaded on startup. API +calls to create new policies save the definition file to this direcory. +In addition, API calls to delete policies will remove the definition file. `, Run: func(cmd *cobra.Command, args []string) { params.Paths = args @@ -43,10 +67,12 @@ base documents (e.g., example.json) or policies (e.g., example.rego). runCommand.Flags().BoolVarP(¶ms.Server, "server", "s", false, "start the runtime in server mode") runCommand.Flags().StringVarP(¶ms.HistoryPath, "history", "H", historyPath(), "set path of history file") + runCommand.Flags().StringVarP(¶ms.PolicyDir, "policy-dir", "p", "", "set directory to store policy definitions") + runCommand.Flags().StringVarP(¶ms.Addr, "addr", "a", defaultAddr, "set listening address of the server") usageTemplate := `Usage: {{.UseLine}} [flags] [files] - + Flags: {{.LocalFlags.FlagUsages | trimRightSpace}} ` @@ -63,3 +89,11 @@ func historyPath() string { } return path.Join(home, defaultHistoryFile) } + +func policyDir() string { + cwd, err := os.Getwd() + if err != nil { + return defaultPolicyDir + } + return filepath.Join(cwd, defaultPolicyDir) +} diff --git a/cmd/version.go b/cmd/version.go index 956157d7d9..ce1eb00191 100644 --- a/cmd/version.go +++ b/cmd/version.go @@ -4,9 +4,12 @@ package cmd -import "fmt" -import "github.com/spf13/cobra" -import "github.com/open-policy-agent/opa/version" +import ( + "fmt" + + "github.com/open-policy-agent/opa/version" + "github.com/spf13/cobra" +) var versionCommand = &cobra.Command{ Use: "version", diff --git a/eval/index.go b/eval/index.go index 5cd92410bd..4160105bd0 100644 --- a/eval/index.go +++ b/eval/index.go @@ -293,11 +293,8 @@ func iterStorage(store *Storage, ref ast.Ref, path ast.Ref, bindings *Bindings, if len(ref) == 0 { node, err := lookup(store, path) if err != nil { - switch err := err.(type) { - case *StorageError: - if err.Code == StorageNotFoundErr { - return nil - } + if IsStorageNotFound(err) { + return nil } return err } @@ -318,11 +315,8 @@ func iterStorage(store *Storage, ref ast.Ref, path ast.Ref, bindings *Bindings, node, err := lookup(store, path) if err != nil { - switch err := err.(type) { - case *StorageError: - if err.Code == StorageNotFoundErr { - return nil - } + if IsStorageNotFound(err) { + return nil } return err } diff --git a/eval/policystore.go b/eval/policystore.go new file mode 100644 index 0000000000..bebcfd8a16 --- /dev/null +++ b/eval/policystore.go @@ -0,0 +1,321 @@ +// 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 eval + +import ( + "fmt" + "io/ioutil" + "os" + "path/filepath" + "strings" + + "github.com/open-policy-agent/opa/ast" + "github.com/pkg/errors" +) + +// PolicyStore provides a storage abstraction for policy definitions and modules. +// +type PolicyStore struct { + dataStore *Storage + policyDir string + raw map[string][]byte + modules map[string]*ast.Module +} + +// LoadPolicies is the default callback function that will be used when +// opening the policy store. +func LoadPolicies(bufs map[string][]byte) (map[string]*ast.Module, error) { + + parsed := map[string]*ast.Module{} + + for id, bs := range bufs { + mod, err := ast.ParseModule(string(bs)) + if err != nil { + return nil, err + } + parsed[id] = mod + } + + c := ast.NewCompiler() + if c.Compile(parsed); c.Failed() { + return nil, c.Errors[0] + } + + return c.Modules, nil +} + +// NewPolicyStore returns an empty PolicyStore. +func NewPolicyStore(store *Storage, policyDir string) *PolicyStore { + return &PolicyStore{ + dataStore: store, + policyDir: policyDir, + raw: map[string][]byte{}, + modules: map[string]*ast.Module{}, + } +} + +// List returns all of the modules. +func (p *PolicyStore) List() map[string]*ast.Module { + cpy := map[string]*ast.Module{} + for k, v := range p.modules { + cpy[k] = v + } + return cpy +} + +// Open initializes the policy store. +// +// This should be called on startup to load policies from persistent storage. +// The callback function "f" will be invoked with the buffers representing the +// persisted policies. The callback should return the compiled version of the +// policies so that they can be installed into the data store. +// +func (p *PolicyStore) Open(f func(map[string][]byte) (map[string]*ast.Module, error)) error { + + if len(p.policyDir) == 0 { + return nil + } + + info, err := ioutil.ReadDir(p.policyDir) + if err != nil { + return err + } + + raw := map[string][]byte{} + + for _, i := range info { + + f := i.Name() + bs, err := ioutil.ReadFile(filepath.Join(p.policyDir, f)) + + if err != nil { + return err + } + + id := p.getID(f) + raw[id] = bs + } + + mods, err := f(raw) + if err != nil { + return err + } + + for id, mod := range mods { + if err := p.Add(id, mod, raw[id], false); err != nil { + return err + } + } + + return nil +} + +// Add inserts the policy module into the store. If an existing policy module exists with the same ID, +// it is overwritten. If persist is false, then the policy will not be persisted. +func (p *PolicyStore) Add(id string, mod *ast.Module, raw []byte, persist bool) error { + + if persist && len(p.policyDir) == 0 { + return fmt.Errorf("cannot persist without --policy-dir set") + } + + old := p.modules[id] + + if old != nil { + if err := p.uninstallModule(old); err != nil { + return errors.Wrapf(err, "failed to uninstall old version of module: %v", id) + } + } + + if err := p.installModule(mod); err != nil { + return errors.Wrapf(err, "failed to install module but old version of module was uninstalled: %v", id) + } + + p.raw[id] = raw + p.modules[id] = mod + + if persist { + filename := p.getFilename(id) + if err := ioutil.WriteFile(filename, raw, 0644); err != nil { + return errors.Wrapf(err, "failed to persist definition but new version was installed: %v", id) + } + } + + return nil +} + +// Remove removes the policy module for id. +func (p *PolicyStore) Remove(id string) error { + + mod, err := p.Get(id) + if err != nil { + return err + } + + if err := p.uninstallModule(mod); err != nil { + return errors.Wrapf(err, "failed to uninstall module: %v", id) + } + + filename := p.getFilename(id) + + if strings.HasPrefix(filename, p.policyDir) { + if err := os.Remove(filename); err != nil { + if !os.IsNotExist(err) { + return errors.Wrapf(err, "failed to delete persisted definition but module was uninstalled: %v", id) + } + } + } + + delete(p.raw, id) + delete(p.modules, id) + + return nil +} + +// Get returns the policy module for id. +func (p *PolicyStore) Get(id string) (*ast.Module, error) { + mod, ok := p.modules[id] + if !ok { + return nil, notFoundErrorf("module not found: %v", id) + } + return mod, nil +} + +// GetRaw returns the raw content of the module for id. +func (p *PolicyStore) GetRaw(id string) ([]byte, error) { + bs, ok := p.raw[id] + if !ok { + return nil, notFoundErrorf("definition not found: %v", id) + } + return bs, nil +} + +func (p *PolicyStore) getFilename(id string) string { + return filepath.Join(p.policyDir, id) +} + +func (p *PolicyStore) getID(f string) string { + return filepath.Base(f) +} + +func (p *PolicyStore) installModule(mod *ast.Module) error { + + installed := map[*ast.Rule][]interface{}{} + + for _, r := range mod.Rules { + fqn := append(ast.Ref{}, mod.Package.Path...) + fqn = append(fqn, &ast.Term{Value: ast.String(r.Name)}) + path, _ := fqn.Underlying() + path = path[1:] + if err := p.installRule(path, r); err != nil { + for r, path := range installed { + if err := p.uninstallRule(path, r); err != nil { + return err + } + } + return err + } + installed[r] = path + } + + return nil +} + +func (p *PolicyStore) installRule(path []interface{}, rule *ast.Rule) error { + + err := p.dataStore.MakePath(path[:len(path)-1]) + if err != nil { + return errors.Wrapf(err, "unable to make path for rule set") + } + + node, err := p.dataStore.Get(path) + if err != nil { + switch err := err.(type) { + case *StorageError: + if err.Code == StorageNotFoundErr { + rules := []*ast.Rule{rule} + if err := p.dataStore.Patch(StorageAdd, path, rules); err != nil { + return errors.Wrapf(err, "unable to add new rule set") + } + return nil + } + } + return err + } + + rs, ok := node.([]*ast.Rule) + if !ok { + return fmt.Errorf("unable to add rule to base document") + } + + for i := range rs { + if rs[i].Equal(rule) { + return nil + } + } + + rs = append(rs, rule) + + if err := p.dataStore.Patch(StorageReplace, path, rs); err != nil { + return errors.Wrapf(err, "unable to add rule to existing rule set") + } + + return nil +} + +func (p *PolicyStore) uninstallModule(mod *ast.Module) error { + uninstalled := map[*ast.Rule][]interface{}{} + for _, r := range mod.Rules { + fqn := append(ast.Ref{}, mod.Package.Path...) + fqn = append(fqn, &ast.Term{Value: ast.String(r.Name)}) + path, _ := fqn.Underlying() + path = path[1:] + if err := p.uninstallRule(path, r); err != nil { + for r, path := range uninstalled { + if err := p.installRule(path, r); err != nil { + return err + } + } + return err + } + uninstalled[r] = path + } + return nil +} + +// uninstallRule removes the rule located at the path. If the path is not found +// or the rule does not exist in the ruleset, this function returns nil (no error). +func (p *PolicyStore) uninstallRule(path []interface{}, rule *ast.Rule) error { + + node, err := p.dataStore.Get(path) + + if IsStorageNotFound(err) { + return nil + } + + rs, ok := node.([]*ast.Rule) + if !ok { + return fmt.Errorf("unable to remove rule: path refers to base document") + } + + found := false + + for i := range rs { + if rs[i].Equal(rule) { + rs = append(rs[:i], rs[i+1:]...) + found = true + break + } + } + + if !found { + return nil + } + + if len(rs) == 0 { + return p.dataStore.Patch(StorageRemove, path, nil) + } + + return p.dataStore.Patch(StorageReplace, path, rs) +} diff --git a/eval/policystore_test.go b/eval/policystore_test.go new file mode 100644 index 0000000000..b627dadf1d --- /dev/null +++ b/eval/policystore_test.go @@ -0,0 +1,305 @@ +// 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 eval + +import ( + "io/ioutil" + "os" + "path/filepath" + "testing" + + "github.com/open-policy-agent/opa/ast" +) + +func TestPolicyStoreDefaultOpen(t *testing.T) { + + dir, err := ioutil.TempDir("", "policyDir") + if err != nil { + panic(err) + } + + defer os.RemoveAll(dir) + + filename := filepath.Join(dir, "testMod1") + + err = ioutil.WriteFile(filename, []byte(testMod1), 0644) + if err != nil { + panic(err) + } + + dataStore := NewStorage() + policyStore := NewPolicyStore(dataStore, dir) + + err = policyStore.Open(LoadPolicies) + if err != nil { + t.Errorf("Unexpected error on Open(): %v", err) + return + } + + c := ast.NewCompiler() + mod := ast.MustParseModule(testMod1) + if c.Compile(map[string]*ast.Module{"testMod1": mod}); c.Failed() { + panic(c.FlattenErrors()) + } + + stored, err := policyStore.Get("testMod1") + if err != nil { + t.Errorf("Unexpected error on Get(): %v", err) + return + } + + if !c.Modules["testMod1"].Equal(stored) { + t.Errorf("Expected %v from policy store but got: %v", c.Modules["testMod1"], stored) + return + } +} + +func TestPolicyStoreAdd(t *testing.T) { + + f := newFixture() + defer f.cleanup() + + mod1 := f.compile1(testMod1) + mod2 := f.compile1(testMod2) + + err := f.policyStore.Add("testMod1", mod1, []byte(testMod1), true) + if err != nil { + t.Errorf("Unexpected error on Add(): %v", err) + return + } + + err = f.policyStore.Add("testMod2", mod2, []byte(testMod2), true) + if err != nil { + t.Errorf("Unexpected error on Add(): %v", err) + return + } + + r, err := f.policyStore.Get("testMod1") + if err != nil { + t.Errorf("Unexpected error on Get(): %v", err) + return + } + + if !mod1.Equal(r) { + t.Errorf("Expected %v for Get() but got: %v", mod1, r) + return + } + + raw, err := f.policyStore.GetRaw("testMod1") + if err != nil { + t.Errorf("Unexpected error on GetRaw(): %v", err) + return + } + + if string(raw) != testMod1 { + t.Errorf("Expected %v for GetRaw() but got: %v", testMod1, raw) + } + + mods := f.policyStore.List() + + if len(mods) != 2 { + t.Errorf("Expected a single module from List() but got: %v", mods) + return + } + + if !mods["testMod1"].Equal(mod1) { + t.Errorf("Expected List() result to equal %v but got %v", mod1, mods["testMod1"]) + return + } +} + +func TestPolicyStoreAddIdempotent(t *testing.T) { + + f := newFixture() + defer f.cleanup() + + mod1 := f.compile1(testMod1) + + err := f.policyStore.Add("testMod1", mod1, []byte(testMod1), true) + if err != nil { + t.Errorf("Unexpected error on Add(): %v", err) + return + } + + err = f.policyStore.Add("testMod1", mod1, []byte(testMod1), true) + if err != nil { + t.Errorf("Unexpected error on Add(): %v", err) + return + } + + node, err := f.dataStore.Get(path("a.b.p")) + if err != nil { + t.Errorf("Unexpected error on Get(): %v", err) + return + } + + rules := node.([]*ast.Rule) + if len(rules) != 1 { + t.Errorf("Expected ruleset to exactly one rule: %v", rules) + return + } + + if !rules[0].Equal(mod1.Rules[0]) { + t.Errorf("Expected rule to be %v but got: %v", mod1, rules[0]) + return + } +} + +func TestPolicyStoreRemove(t *testing.T) { + + f := newFixture() + defer f.cleanup() + + mod1 := f.compile1(testMod1) + mod2 := f.compile1(testMod2) + + err := f.policyStore.Add("testMod1", mod1, []byte(testMod1), true) + if err != nil { + t.Errorf("Unexpected error on Add(): %v", err) + return + } + + err = f.policyStore.Add("testMod2", mod2, []byte(testMod2), true) + if err != nil { + t.Errorf("Unexpected error on Add(): %v", err) + return + } + + if err := f.policyStore.Remove("testMod1"); err != nil { + t.Errorf("Unexpected error on Remove(): %v", err) + return + } + + mods := f.policyStore.List() + + if len(mods) != 1 { + t.Errorf("Expected one module to remain after Remove(): %v", mods) + return + } + + if _, err := f.policyStore.Get("testMod2"); err != nil { + t.Errorf("Expected testMod2 to remain after Remove(): %v", mods) + return + } + + _, err = os.Stat(f.policyStore.getFilename("testMod1")) + if !os.IsNotExist(err) { + info, err := ioutil.ReadDir(f.policyStore.policyDir) + if err != nil { + panic(err) + } + files := []string{} + for _, i := range info { + files = append(files, i.Name()) + } + t.Errorf("Expected testMod1 to be removed from disk but %v contains: %v", f.policyStore.policyDir, files) + return + } +} + +func TestPolicyStoreUpdate(t *testing.T) { + f := newFixture() + defer f.cleanup() + + mod1 := f.compile1(testMod1) + mod2 := f.compile1(testMod2) + + err := f.policyStore.Add("testMod1", mod1, []byte(testMod1), true) + if err != nil { + t.Errorf("Unexpected error on Add(): %v", err) + return + } + + err = f.policyStore.Add("testMod1", mod2, []byte(testMod2), true) + if err != nil { + t.Errorf("Unexpected error on Add(): %v", err) + return + } + + node, err := f.dataStore.Get(path("a.b.p")) + if err != nil { + t.Errorf("Unexpected error on Get(): %v", err) + return + } + + rules := node.([]*ast.Rule) + if len(rules) != 1 { + t.Errorf("Expected exactly one rule but got: %v", rules) + return + } + + if !rules[0].Equal(mod2.Rules[0]) { + t.Errorf("Expected rule to equal %v but got: %v", mod2.Rules[0], rules[0]) + return + } + + node, err = f.dataStore.Get(path("a.b.q")) + if !IsStorageNotFound(err) { + t.Errorf("Expected storage not found error but got: %v (err: %v)", node, err) + return + } +} + +const ( + testMod1 = ` + package a.b + + p = true :- true + q = true :- true + ` + + testMod2 = ` + package a.b + + p = true :- false + ` +) + +type fixture struct { + policyStore *PolicyStore + dataStore *Storage +} + +func newFixture() *fixture { + + dir, err := ioutil.TempDir("", "policyDir") + if err != nil { + panic(err) + } + + dataStore := NewStorage() + policyStore := NewPolicyStore(dataStore, dir) + err = policyStore.Open(func(map[string][]byte) (map[string]*ast.Module, error) { + return nil, nil + }) + if err != nil { + panic(err) + } + + f := &fixture{ + policyStore: policyStore, + dataStore: dataStore, + } + + return f +} + +func (f *fixture) cleanup() { + os.RemoveAll(f.policyStore.policyDir) +} + +func (f *fixture) compile1(m string) *ast.Module { + + mods := f.policyStore.List() + mod := ast.MustParseModule(m) + mods[""] = mod + + c := ast.NewCompiler() + if c.Compile(mods); c.Failed() { + panic(c.FlattenErrors()) + } + + return c.Modules[""] +} diff --git a/eval/storage.go b/eval/storage.go index 9ab7809f40..a75ba58c7e 100644 --- a/eval/storage.go +++ b/eval/storage.go @@ -5,12 +5,9 @@ package eval import ( - "encoding/json" "fmt" - "os" "github.com/open-policy-agent/opa/ast" - "github.com/pkg/errors" ) // StorageErrorCode represents the collection of error types that can be @@ -36,6 +33,15 @@ func (err *StorageError) Error() string { return fmt.Sprintf("storage error (code: %d): %v", err.Code, err.Message) } +// IsStorageNotFound returns true if this error is a StorageNotFoundErr +func IsStorageNotFound(err error) bool { + switch err := err.(type) { + case *StorageError: + return err.Code == StorageNotFoundErr + } + return false +} + var doesNotExistMsg = "document does not exist" var outOfRangeMsg = "array index out of range" var nonEmptyMsg = "path must be non-empty" @@ -62,6 +68,11 @@ func notFoundError(path []interface{}, f string, a ...interface{}) *StorageError if len(f) > 0 { msg += ", " + fmt.Sprintf(f, a...) } + return notFoundErrorf(msg) +} + +func notFoundErrorf(f string, a ...interface{}) *StorageError { + msg := fmt.Sprintf(f, a...) return &StorageError{ Code: StorageNotFoundErr, Message: msg, @@ -74,110 +85,18 @@ type Storage struct { data map[string]interface{} } -// NewEmptyStorage is a helper for creating a new, empty Storage. -func NewEmptyStorage() *Storage { +// NewStorage is a helper for creating a new, empty Storage. +func NewStorage() *Storage { return &Storage{ Indices: NewIndices(), data: map[string]interface{}{}, } } -// NewStorage is a helper for creating a new Storage containing -// the given base documents and rules. -func NewStorage(docs []map[string]interface{}, mods map[string]*ast.Module) (*Storage, error) { - - store := NewEmptyStorage() - - for _, d := range docs { - // TODO(tsandall): recursive merge instead of replace? - for k, v := range d { - if err := store.Patch(StorageAdd, []interface{}{k}, v); err != nil { - return nil, err - } - } - } - - for _, m := range mods { - - for _, r := range m.Rules { - - fqn := append(ast.Ref{}, m.Package.Path...) - fqn = append(fqn, &ast.Term{Value: ast.String(r.Name)}) - - path, _ := fqn.Underlying() - path = path[1:] - - err := store.MakePath(path[:len(path)-1]) - if err != nil { - return nil, errors.Wrapf(err, "unable to make path for rule set") - } - - node, err := store.Get(path) - if err != nil { - switch err := err.(type) { - case *StorageError: - if err.Code == StorageNotFoundErr { - rules := []*ast.Rule{r} - if err := store.Patch(StorageAdd, path, rules); err != nil { - return nil, errors.Wrapf(err, "unable to add new rule set") - } - continue - } - } - return nil, err - } - - rs, ok := node.([]*ast.Rule) - if !ok { - return nil, fmt.Errorf("unable to add rule to base document") - } - - rs = append(rs, r) - - if err := store.Patch(StorageReplace, path, rs); err != nil { - return nil, errors.Wrapf(err, "unable to add rule to existing rule set") - } - } - } - - return store, nil -} - -// NewStorageFromFiles is a helper for creating a new Storage containing -// documents stored in files and/or policy modules. -func NewStorageFromFiles(files []string) (*Storage, error) { - - modules := map[string]*ast.Module{} - docs := []map[string]interface{}{} - - for _, file := range files { - m, astErr := ast.ParseModuleFile(file) - if astErr == nil { - modules[file] = m - continue - } - d, jsonErr := parseJSONObjectFile(file) - if jsonErr == nil { - docs = append(docs, d) - continue - } - // TODO(tsandall): add heuristic to determine whether this supposed - // to be a policy module or a JSON file. Format appropriate error. - return nil, fmt.Errorf("parse error: %v: %v: %v", file, astErr, jsonErr) - } - - c := ast.NewCompiler() - c.Compile(modules) - if c.Failed() { - return nil, fmt.Errorf(c.FlattenErrors()) - } - - return NewStorage(docs, c.Modules) -} - // NewStorageFromJSONObject returns Storage by converting from map[string]interface{} +// This is mostly for test purposes. func NewStorageFromJSONObject(data map[string]interface{}) *Storage { - store := NewEmptyStorage() + store := NewStorage() for k, v := range data { if err := store.Patch(StorageAdd, []interface{}{k}, v); err != nil { panic(err) @@ -661,17 +580,3 @@ func checkArrayIndex(path []interface{}, node []interface{}, v interface{}) (int } return i, nil } - -func parseJSONObjectFile(file string) (map[string]interface{}, error) { - f, err := os.Open(file) - if err != nil { - return nil, err - } - defer f.Close() - reader := json.NewDecoder(f) - var data map[string]interface{} - if err := reader.Decode(&data); err != nil { - return nil, err - } - return data, nil -} diff --git a/eval/storage_test.go b/eval/storage_test.go index 6498528a5f..1d9b6e5af1 100644 --- a/eval/storage_test.go +++ b/eval/storage_test.go @@ -7,70 +7,12 @@ package eval import ( "encoding/json" "fmt" - "io/ioutil" - "os" "reflect" "testing" "github.com/open-policy-agent/opa/ast" ) -func TestLoadFromFiles(t *testing.T) { - tmp1, err := ioutil.TempFile("", "docFile") - if err != nil { - panic(err) - } - defer os.Remove(tmp1.Name()) - doc1 := `{"foo": "bar", "a": {"b": {"d": [1]}}}` - if _, err := tmp1.Write([]byte(doc1)); err != nil { - panic(err) - } - if err := tmp1.Close(); err != nil { - panic(err) - } - - tmp2, err := ioutil.TempFile("", "policyFile") - if err != nil { - panic(err) - } - defer os.Remove(tmp2.Name()) - mod1 := ` - package a.b.c - import data.foo - p = true :- foo = "bar" - p = true :- 1 = 2 - ` - if _, err := tmp2.Write([]byte(mod1)); err != nil { - panic(err) - } - if err := tmp2.Close(); err != nil { - panic(err) - } - - store, err := NewStorageFromFiles([]string{tmp1.Name(), tmp2.Name()}) - if err != nil { - t.Errorf("Unexpected error: %v", err) - return - } - - r, err := store.Get(path("foo")) - if Compare(r, "bar") != 0 || err != nil { - t.Errorf("Expected %v but got %v (err: %v)", "bar", r, err) - return - } - - r, err = store.Get(path("a.b.c.p")) - rules, ok := r.([]*ast.Rule) - if !ok { - t.Errorf("Expected rules but got: %v", r) - return - } - if !rules[0].Name.Equal(ast.Var("p")) { - t.Errorf("Expected rule p but got: %v", rules[0]) - return - } -} - func TestStorageGet(t *testing.T) { data := loadSmallTestData() diff --git a/eval/topdown.go b/eval/topdown.go index c5b925e625..d7fd26b907 100644 --- a/eval/topdown.go +++ b/eval/topdown.go @@ -660,11 +660,8 @@ func evalRefRecEnumColl(ctx *TopDownContext, path, tail ast.Ref, iter TopDownIte node, err := lookup(ctx.Store, path) if err != nil { - switch err := err.(type) { - case *StorageError: - if err.Code == StorageNotFoundErr { - return nil - } + if IsStorageNotFound(err) { + return nil } return err } @@ -715,11 +712,8 @@ func evalRefRecGround(ctx *TopDownContext, path, tail ast.Ref, iter TopDownItera path = append(path, tail[0]) node, err := lookupRule(ctx.Store, path) if err != nil { - switch err := err.(type) { - case *StorageError: - if err.Code == StorageNotFoundErr { - return nil - } + if IsStorageNotFound(err) { + return nil } return err } @@ -1221,11 +1215,8 @@ func lookup(store *Storage, ref ast.Ref) (interface{}, error) { func lookupExists(store *Storage, ref ast.Ref) (bool, error) { _, err := lookup(store, ref) if err != nil { - switch err := err.(type) { - case *StorageError: - if err.Code == StorageNotFoundErr { - return false, nil - } + if IsStorageNotFound(err) { + return false, nil } return false, err } diff --git a/eval/topdown_test.go b/eval/topdown_test.go index 7de14e4e5c..627c7d59b7 100644 --- a/eval/topdown_test.go +++ b/eval/topdown_test.go @@ -525,9 +525,15 @@ func TestTopDownEmbeddedVirtualDoc(t *testing.T) { q[x] :- g[j][k] = x`}) data := loadSmallTestData() - store, err := NewStorage([]map[string]interface{}{data}, mods) - if err != nil { - panic(err) + + store := NewStorageFromJSONObject(data) + policyStore := NewPolicyStore(store, "") + + for id, mod := range mods { + err := policyStore.Add(id, mod, []byte(""), false) + if err != nil { + panic(err) + } } assertTopDown(t, store, 0, "deep embedded vdoc", []string{"b", "c", "d", "p"}, "[1, 2, 4]") @@ -583,9 +589,14 @@ func TestExample(t *testing.T) { mods := compileModules([]string{vd}) - store, err := NewStorage([]map[string]interface{}{doc}, mods) - if err != nil { - panic(err) + store := NewStorageFromJSONObject(doc) + policyStore := NewPolicyStore(store, "") + + for id, mod := range mods { + err := policyStore.Add(id, mod, []byte(""), false) + if err != nil { + panic(err) + } } assertTopDown(t, store, 0, "public servers", []string{"opa", "example", "public_servers"}, ` @@ -629,7 +640,7 @@ func compileRules(imports []string, input []string) map[string]*ast.Module { is := []*ast.Import{} for _, i := range imports { is = append(is, &ast.Import{ - Path: ast.MustParseRef(i), + Path: ast.MustParseTerm(i), }) } @@ -757,10 +768,19 @@ func runTopDownTestCase(t *testing.T, data map[string]interface{}, i int, note s for k := range data { imports = append(imports, "data."+k) } - store, err := NewStorage([]map[string]interface{}{data}, compileRules(imports, rules)) - if err != nil { - panic(err) + + mods := compileRules(imports, rules) + + store := NewStorageFromJSONObject(data) + policyStore := NewPolicyStore(store, "") + + for id, mod := range mods { + err := policyStore.Add(id, mod, []byte(""), false) + if err != nil { + panic(err) + } } + assertTopDown(t, store, i, note, []string{"p"}, expected) } diff --git a/eval/tracer_test.go b/eval/tracer_test.go index 435069c73d..ee2cdb2804 100644 --- a/eval/tracer_test.go +++ b/eval/tracer_test.go @@ -21,14 +21,21 @@ func (t *mockTracer) Trace(ctx *TopDownContext, f string, a ...interface{}) { func TestTracer(t *testing.T) { + data := loadSmallTestData() + mods := compileRules([]string{"data.a"}, []string{ "p[x] :- q[x] = y", "q[i] = j :- a[i] = j", }) - store, err := NewStorage([]map[string]interface{}{loadSmallTestData()}, mods) - if err != nil { - panic(err) + store := NewStorageFromJSONObject(data) + policyStore := NewPolicyStore(store, "") + + for id, mod := range mods { + err := policyStore.Add(id, mod, []byte(""), false) + if err != nil { + panic(err) + } } tracer := &mockTracer{[]string{}} diff --git a/runtime/repl.go b/runtime/repl.go index 5b9dac0764..20a4c457bc 100644 --- a/runtime/repl.go +++ b/runtime/repl.go @@ -8,7 +8,6 @@ import ( "encoding/json" "fmt" "io" - "math/rand" "os" "sort" "strings" @@ -33,6 +32,7 @@ type Repl struct { InitPrompt string BufferPrompt string Buffer []string + nextID int } // NewRepl creates a new Repl. @@ -157,8 +157,10 @@ func (r *Repl) cmdTrace() bool { } func (r *Repl) compileBody(body ast.Body) (ast.Body, error) { + name := fmt.Sprintf("repl%d", r.nextID) + r.nextID++ rule := &ast.Rule{ - Name: ast.Var(randString(32)), + Name: ast.Var(name), Body: body, } // TODO(tsandall): refactor to use current implicit module @@ -170,11 +172,11 @@ func (r *Repl) compileBody(body ast.Body) (ast.Body, error) { Rules: []*ast.Rule{rule}, } c := ast.NewCompiler() - c.Compile(map[string]*ast.Module{"tmp": m}) + c.Compile(map[string]*ast.Module{name: m}) if len(c.Errors) > 0 { return nil, fmt.Errorf(c.FlattenErrors()) } - return c.Modules["tmp"].Rules[0].Body, nil + return c.Modules[name].Rules[0].Body, nil } func (r *Repl) compileRule(rule *ast.Rule) (*ast.Rule, error) { @@ -455,14 +457,3 @@ func buildHeader(fields map[string]struct{}, term *ast.Term) { } } } - -// randString returns a random string of letters. -// http://stackoverflow.com/a/31832326 -func randString(length int) string { - letters := []rune("abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ") - s := make([]rune, length) - for i := range s { - s[i] = letters[rand.Intn(len(letters))] - } - return string(s) -} diff --git a/runtime/runtime.go b/runtime/runtime.go index e378f69848..fb66d2f0a5 100644 --- a/runtime/runtime.go +++ b/runtime/runtime.go @@ -5,50 +5,217 @@ package runtime import ( + "encoding/json" "fmt" + "io/ioutil" "os" + "path/filepath" + "github.com/open-policy-agent/opa/ast" "github.com/open-policy-agent/opa/eval" + "github.com/pkg/errors" ) // Params stores the configuration for an OPA instance. type Params struct { - Server bool - Paths []string + + // Addr is the listening address that the OPA server will bind to. + Addr string + + // Server flag controls whether the OPA instance will start a server. + // By default, the OPA instance acts as an interactive shell. + Server bool + + // Paths contains filenames of base documents and policy modules to + // load on startup. + Paths []string + + // HistoryPath is the filename to store the interactive shell user + // input history. HistoryPath string + + // PolicyDir is the filename of the directory to persist policy + // definitions in. Policy definitions stored in this directory + // are automatically loaded on startup. + PolicyDir string } // Runtime represents a single OPA instance. type Runtime struct { - Store *eval.Storage + Store *eval.Storage + PolicyStore *eval.PolicyStore +} + +// Init initializes the OPA instance. +func (rt *Runtime) Init(params *Params) error { + + if len(params.PolicyDir) > 0 { + if err := os.MkdirAll(params.PolicyDir, 0755); err != nil { + return errors.Wrap(err, "unable to make --policy-dir") + } + } + + parsed, err := parseInputs(params.Paths) + if err != nil { + return errors.Wrapf(err, "parse error") + } + + // Open data store and load base documents. + dataStore := eval.NewStorage() + + for _, doc := range parsed.docs { + for k, v := range doc { + if err := dataStore.Patch(eval.StorageAdd, []interface{}{k}, v); err != nil { + return errors.Wrap(err, "unable to open data store") + } + } + } + + // Open policy store and load existing policies. + policyStore := eval.NewPolicyStore(dataStore, params.PolicyDir) + if err := policyStore.Open(eval.LoadPolicies); err != nil { + return errors.Wrap(err, "unable to open policy store") + } + + // Load policies provided via input. + if err := compileAndStoreInputs(parsed.modules, policyStore); err != nil { + return errors.Wrapf(err, "compile error") + } + + rt.PolicyStore = policyStore + rt.Store = dataStore + + return nil } // Start is the entry point of an OPA instance. func (rt *Runtime) Start(params *Params) { - store, err := eval.NewStorageFromFiles(params.Paths) - - if err != nil { - fmt.Println("failed to open storage:", err) + if err := rt.Init(params); err != nil { + fmt.Println("error initializing runtime:", err) os.Exit(1) } - rt.Store = store - if !params.Server { - rt.runRepl(params) + rt.startRepl(params) } else { - rt.runServer(params) + rt.startServer(params) } } -func (rt *Runtime) runServer(params *Params) { - fmt.Println("not implemented: server mode") - os.Exit(1) +func (rt *Runtime) startServer(params *Params) { + persist := len(params.PolicyDir) > 0 + server := NewServer(rt, params.Addr, persist) + server.Loop() } -func (rt *Runtime) runRepl(params *Params) { - +func (rt *Runtime) startRepl(params *Params) { repl := NewRepl(rt, params.HistoryPath, os.Stdout) repl.Loop() } + +func compileAndStoreInputs(parsed map[string]*parsedModule, policyStore *eval.PolicyStore) error { + + mods := policyStore.List() + for _, p := range parsed { + mods[p.id] = p.mod + } + + c := ast.NewCompiler() + if c.Compile(mods); c.Failed() { + // TODO(tsandall): add another call on compiler to flatten into error type + return c.Errors[0] + } + + for id := range parsed { + mod := c.Modules[id] + if err := policyStore.Add(id, mod, parsed[id].raw, false); err != nil { + return err + } + } + + return nil +} + +type parsedModule struct { + id string + mod *ast.Module + raw []byte +} + +type parsedInput struct { + docs []map[string]interface{} + modules map[string]*parsedModule +} + +func parseInputs(paths []string) (*parsedInput, error) { + + parsedDocs := []map[string]interface{}{} + parsedModules := map[string]*parsedModule{} + + for _, file := range paths { + + info, err := os.Stat(file) + if err != nil { + return nil, err + } + + if info.IsDir() { + continue + } + + bs, err := ioutil.ReadFile(file) + + if err != nil { + return nil, err + } + + m, astErr := ast.ParseModuleFile(file) + + if astErr == nil { + parsedModules[file] = &parsedModule{ + id: file, + mod: m, + raw: bs, + } + continue + } + + d, jsonErr := parseJSONObjectFile(file) + + if jsonErr == nil { + parsedDocs = append(parsedDocs, d) + continue + } + + switch filepath.Ext(file) { + case ".json": + return nil, jsonErr + case ".rego": + return nil, astErr + default: + return nil, fmt.Errorf("unrecognizable file: %v", file) + } + } + + r := &parsedInput{ + docs: parsedDocs, + modules: parsedModules, + } + + return r, nil +} + +func parseJSONObjectFile(file string) (map[string]interface{}, error) { + f, err := os.Open(file) + if err != nil { + return nil, err + } + defer f.Close() + reader := json.NewDecoder(f) + var data map[string]interface{} + if err := reader.Decode(&data); err != nil { + return nil, err + } + return data, nil +} diff --git a/runtime/runtime_test.go b/runtime/runtime_test.go new file mode 100644 index 0000000000..2fed36a4a4 --- /dev/null +++ b/runtime/runtime_test.go @@ -0,0 +1,125 @@ +// 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 runtime + +import ( + "fmt" + "io/ioutil" + "os" + "path/filepath" + "testing" + + "github.com/open-policy-agent/opa/ast" + "github.com/open-policy-agent/opa/eval" +) + +func TestInit(t *testing.T) { + tmp1, err := ioutil.TempFile("", "docFile") + if err != nil { + panic(err) + } + defer os.Remove(tmp1.Name()) + doc1 := `{"foo": "bar", "a": {"b": {"d": [1]}}}` + if _, err := tmp1.Write([]byte(doc1)); err != nil { + panic(err) + } + if err := tmp1.Close(); err != nil { + panic(err) + } + + tmp2, err := ioutil.TempFile("", "policyFile") + if err != nil { + panic(err) + } + defer os.Remove(tmp2.Name()) + mod1 := ` + package a.b.c + import data.foo + p = true :- foo = "bar" + p = true :- 1 = 2 + ` + if _, err := tmp2.Write([]byte(mod1)); err != nil { + panic(err) + } + if err := tmp2.Close(); err != nil { + panic(err) + } + + tmp3, err := ioutil.TempDir("", "policyDir") + if err != nil { + panic(err) + } + + defer os.RemoveAll(tmp3) + + tmp4 := filepath.Join(tmp3, "existingPolicy") + + err = ioutil.WriteFile(tmp4, []byte(` + package a.b.c + q = true :- p + `), 0644) + if err != nil { + panic(err) + } + + rt := Runtime{} + + err = rt.Init(&Params{ + Paths: []string{tmp1.Name(), tmp2.Name()}, + PolicyDir: tmp3, + }) + + if err != nil { + t.Errorf("Unexpected error: %v", err) + return + } + + node, err := rt.Store.Get(path("foo")) + if eval.Compare(node, "bar") != 0 || err != nil { + t.Errorf("Expected %v but got %v (err: %v)", "bar", node, err) + return + } + + node, err = rt.Store.Get(path("a.b.c.p")) + rules, ok := node.([]*ast.Rule) + if !ok { + t.Errorf("Expected rules but got: %v", node) + return + } + if !rules[0].Name.Equal(ast.Var("p")) { + t.Errorf("Expected rule p but got: %v", rules[0]) + return + } + + node, err = rt.Store.Get(path("a.b.c.q")) + rules, ok = node.([]*ast.Rule) + if !ok { + t.Errorf("Expected rules but got: %v", node) + return + } + if !rules[0].Name.Equal(ast.Var("q")) { + t.Errorf("Expected rule q but got: %v", rules[0]) + return + } +} + +func path(input interface{}) []interface{} { + switch input := input.(type) { + case []interface{}: + return input + case string: + switch v := ast.MustParseTerm(input).Value.(type) { + case ast.Var: + return []interface{}{string(v)} + case ast.Ref: + path, err := v.Underlying() + if err != nil { + panic(err) + } + return path + } + } + panic(fmt.Sprintf("illegal value: %v", input)) +} diff --git a/runtime/server.go b/runtime/server.go new file mode 100644 index 0000000000..1b36686d40 --- /dev/null +++ b/runtime/server.go @@ -0,0 +1,538 @@ +// 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 runtime + +import ( + "encoding/json" + "fmt" + "io" + "io/ioutil" + "net/http" + "strings" + "sync" + "time" + + "github.com/gorilla/mux" + "github.com/open-policy-agent/opa/ast" + "github.com/open-policy-agent/opa/eval" + "github.com/open-policy-agent/opa/version" + "github.com/pkg/errors" +) + +// apiErrorV1 models an error response sent to the client. +type apiErrorV1 struct { + Code int + Message string +} + +func (err *apiErrorV1) Bytes() []byte { + if bs, err := json.MarshalIndent(err, "", " "); err == nil { + return bs + } + return nil +} + +// patchV1 models a single patch operation against a document. +type patchV1 struct { + Op string `json:"op"` + Path string `json:"path"` + Value interface{} `json:"value"` +} + +// policyV1 models a policy module in OPA. +type policyV1 struct { + ID string + Module *ast.Module +} + +func (p *policyV1) Equal(other *policyV1) bool { + return p.ID == other.ID && p.Module.Equal(other.Module) +} + +// resultSetV1 models the result of an ad-hoc query. +type resultSetV1 []map[string]interface{} + +// Server contains runtime state specific to the server-mode persona, e.g., +// HTTP router. +// +// Notes: +// +// - In the future, the HTTP routing could be factored into a separate module +// relying on the server (which would remain RPC-based). For now, it's simpler +// to keep the HTTP routing and backend implementation in one place. +type Server struct { + Addr string + Persist bool + Runtime *Runtime + Router *mux.Router + + mtx sync.RWMutex +} + +// NewServer returns a new Server. +func NewServer(rt *Runtime, addr string, persist bool) *Server { + + s := &Server{ + Addr: addr, + Persist: persist, + Runtime: rt, + Router: mux.NewRouter(), + } + + s.registerHandlerV1("/data/{path:.+}", "GET", s.v1DataGet) + s.registerHandlerV1("/data/{path:.+}", "PATCH", s.v1DataPatch) + s.registerHandlerV1("/policies", "GET", s.v1PoliciesList) + s.registerHandlerV1("/policies/{id}", "DELETE", s.v1PoliciesDelete) + s.registerHandlerV1("/policies/{id}", "GET", s.v1PoliciesGet) + s.registerHandlerV1("/policies/{id}/raw", "GET", s.v1PoliciesRawGet) + s.registerHandlerV1("/policies/{id}", "PUT", s.v1PoliciesPut) + s.registerHandlerV1("/query", "GET", s.v1QueryGet) + + s.Router.HandleFunc("/", s.indexGet).Methods("GET") + + return s +} + +// Loop starts the server. This function does not return. +func (s *Server) Loop() { + http.ListenAndServe(s.Addr, s.Router) +} + +func (s *Server) execQuery(qStr string) (resultSetV1, error) { + + query, err := ast.ParseBody(qStr) + + if err != nil { + return nil, err + } + + path := ast.Ref{ast.DefaultRootDocument} + + rule := &ast.Rule{ + Body: query, + } + + mod := &ast.Module{ + Package: &ast.Package{ + Path: path, + }, + Rules: []*ast.Rule{rule}, + } + + c := ast.NewCompiler() + + if c.Compile(map[string]*ast.Module{"": mod}); c.Failed() { + return nil, c.Errors[0] + } + + compiled := c.Modules[""].Rules[0].Body + + ctx := eval.NewTopDownContext(compiled, s.Runtime.Store) + + results := resultSetV1{} + + s.mtx.RLock() + defer s.mtx.RUnlock() + + err = eval.TopDown(ctx, func(ctx *eval.TopDownContext) error { + result := map[string]interface{}{} + var err error + ctx.Bindings.Iter(func(k, v ast.Value) bool { + kv, ok := k.(ast.Var) + if !ok { + return false + } + vv, e := eval.ValueToInterface(v, ctx) + if err != nil { + err = e + return true + } + result[string(kv)] = vv + return false + }) + if err != nil { + return err + } + if len(result) > 0 { + results = append(results, result) + } + return nil + }) + + return results, err +} + +func (s *Server) indexGet(w http.ResponseWriter, r *http.Request) { + + renderHeader(w) + renderBanner(w) + renderVersion(w) + + values := r.URL.Query() + qStrs := values["q"] + + renderQueryForm(w, qStrs) + + if len(qStrs) > 0 { + qStr := qStrs[len(qStrs)-1] + t0 := time.Now() + results, err := s.execQuery(qStr) + dt := time.Since(t0) + renderQueryResult(w, results, err, dt) + } + + renderFooter(w) +} + +func (s *Server) registerHandlerV1(path string, method string, h func(http.ResponseWriter, *http.Request)) { + s.Router.HandleFunc("/v1"+path, h).Methods(method) +} + +func (s *Server) v1DataGet(w http.ResponseWriter, r *http.Request) { + vars := mux.Vars(r) + path := strings.Split(vars["path"], "/") + params := &eval.TopDownQueryParams{ + Store: s.Runtime.Store, + Path: path, + } + + s.mtx.RLock() + defer s.mtx.RUnlock() + + result, err := eval.TopDownQuery(params) + + if err != nil { + handleErrorAuto(w, err) + return + } + + handleResponseJSON(w, 200, result) +} + +func (s *Server) v1DataPatch(w http.ResponseWriter, r *http.Request) { + vars := mux.Vars(r) + root := strings.Split(vars["path"], "/") + + ops := []patchV1{} + if err := json.NewDecoder(r.Body).Decode(&ops); err != nil { + handleError(w, 400, err) + return + } + + path := []interface{}{} + for _, x := range root { + path = append(path, x) + } + + s.mtx.Lock() + defer s.mtx.Unlock() + + for i := range ops { + + var op eval.StorageOp + + // TODO this could be refactored for failure handling + switch ops[i].Op { + case "add": + op = eval.StorageAdd + case "remove": + op = eval.StorageRemove + case "replace": + op = eval.StorageReplace + default: + handleErrorf(w, 400, "bad patch operation: %v", ops[i].Op) + return + } + + parts := strings.Split(ops[i].Path[1:], "/") + + for _, x := range parts { + if x == "" { + continue + } + path = append(path, x) + } + + if err := s.Runtime.Store.Patch(op, path, ops[i].Value); err != nil { + handleErrorAuto(w, err) + return + } + + path = path[:len(root)] + } + + handleResponse(w, 204, nil) +} + +func (s *Server) v1PoliciesDelete(w http.ResponseWriter, r *http.Request) { + vars := mux.Vars(r) + id := vars["id"] + + _, err := s.Runtime.PolicyStore.Get(id) + if err != nil { + handleErrorAuto(w, err) + return + } + + s.mtx.Lock() + defer s.mtx.Unlock() + + mods := s.Runtime.PolicyStore.List() + delete(mods, id) + + c := ast.NewCompiler() + if c.Compile(mods); c.Failed() { + handleErrorf(w, 400, c.FlattenErrors()) + return + } + + if err := s.Runtime.PolicyStore.Remove(id); err != nil { + handleErrorAuto(w, err) + return + } + + handleResponse(w, 204, nil) +} + +func (s *Server) v1PoliciesGet(w http.ResponseWriter, r *http.Request) { + vars := mux.Vars(r) + id := vars["id"] + + s.mtx.RLock() + defer s.mtx.RUnlock() + + mod, err := s.Runtime.PolicyStore.Get(id) + if err != nil { + handleErrorAuto(w, err) + return + } + + policy := &policyV1{ + ID: id, + Module: mod, + } + + handleResponseJSONPretty(w, 200, policy) +} + +func (s *Server) v1PoliciesRawGet(w http.ResponseWriter, r *http.Request) { + vars := mux.Vars(r) + id := vars["id"] + + s.mtx.RLock() + defer s.mtx.RUnlock() + + bs, err := s.Runtime.PolicyStore.GetRaw(id) + + if err != nil { + handleErrorAuto(w, err) + return + } + + handleResponse(w, 200, bs) +} + +func (s *Server) v1PoliciesList(w http.ResponseWriter, r *http.Request) { + + policies := []*policyV1{} + + s.mtx.RLock() + defer s.mtx.RUnlock() + + for id, mod := range s.Runtime.PolicyStore.List() { + policy := &policyV1{ + ID: id, + Module: mod, + } + policies = append(policies, policy) + } + + handleResponseJSONPretty(w, 200, policies) +} + +func (s *Server) v1PoliciesPut(w http.ResponseWriter, r *http.Request) { + + vars := mux.Vars(r) + id := vars["id"] + + buf, err := ioutil.ReadAll(r.Body) + if err != nil { + handleError(w, 500, err) + return + } + + mod, err := ast.ParseModule(string(buf)) + if err != nil { + handleError(w, 400, err) + return + } + if mod == nil { + handleErrorf(w, 400, "refusing to add empty module") + return + } + + s.mtx.Lock() + defer s.mtx.Unlock() + + mods := s.Runtime.PolicyStore.List() + mods[id] = mod + + c := ast.NewCompiler() + + if c.Compile(mods); c.Failed() { + handleErrorf(w, 400, c.FlattenErrors()) + return + } + + mod = c.Modules[id] + + if err := s.Runtime.PolicyStore.Add(id, mod, buf, s.Persist); err != nil { + handleErrorAuto(w, err) + return + } + + policy := &policyV1{ + ID: id, + Module: mod, + } + + handleResponseJSONPretty(w, 200, policy) +} + +func (s *Server) v1QueryGet(w http.ResponseWriter, r *http.Request) { + values := r.URL.Query() + qStrs := values["q"] + if len(qStrs) == 0 { + handleErrorf(w, 400, "missing query parameter 'q'") + return + } + + qStr := qStrs[len(qStrs)-1] + results, err := s.execQuery(qStr) + + if err != nil { + handleErrorAuto(w, err) + return + } + + handleResponseJSON(w, 200, results) +} + +func handleError(w http.ResponseWriter, code int, err error) { + handleErrorf(w, code, err.Error()) +} + +func handleErrorAuto(w http.ResponseWriter, err error) { + var prev error + for curr := err; curr != prev; { + if eval.IsStorageNotFound(curr) { + handleError(w, 404, err) + return + } + prev = curr + curr = errors.Cause(prev) + } + handleError(w, 500, err) +} + +func handleErrorf(w http.ResponseWriter, code int, f string, a ...interface{}) { + headers := w.Header() + headers.Add("Content-Type", "application/json") + e := &apiErrorV1{Code: code, Message: fmt.Sprintf(f, a...)} + w.WriteHeader(code) + w.Write(e.Bytes()) +} + +func handleResponse(w http.ResponseWriter, code int, bs []byte) { + w.WriteHeader(code) + if code == 204 { + return + } + w.Write(bs) +} + +func handleResponseJSON(w http.ResponseWriter, code int, v interface{}) { + bs, err := json.Marshal(v) + if err != nil { + handleErrorAuto(w, err) + return + } + headers := w.Header() + headers.Add("Content-Type", "application/json") + handleResponse(w, code, bs) +} + +func handleResponseJSONPretty(w http.ResponseWriter, code int, v interface{}) { + bs, err := json.MarshalIndent(v, "", " ") + if err != nil { + handleErrorAuto(w, err) + return + } + headers := w.Header() + headers.Add("Content-Type", "application/json") + handleResponse(w, code, bs) +} + +func renderBanner(w http.ResponseWriter) { + fmt.Fprintln(w, `
+ ________      ________    ________
+|\   __  \    |\   __  \  |\   __  \
+\ \  \|\  \   \ \  \|\  \ \ \  \|\  \
+ \ \  \\\  \   \ \   ____\ \ \   __  \
+  \ \  \\\  \   \ \  \___|  \ \  \ \  \
+   \ \_______\   \ \__\      \ \__\ \__\
+    \|_______|    \|__|       \|__|\|__|
+	
`) + fmt.Fprintln(w, "Open Policy Agent - An open source project to policy enable any application.
") + fmt.Fprintln(w, "
") +} + +func renderFooter(w http.ResponseWriter) { + fmt.Fprintln(w, "") + fmt.Fprintln(w, "") +} + +func renderHeader(w http.ResponseWriter) { + fmt.Fprintln(w, "") + fmt.Fprintln(w, "") +} + +func renderQueryForm(w http.ResponseWriter, qStrs []string) { + + input := "" + + if len(qStrs) > 0 { + input = qStrs[len(qStrs)-1] + } + + fmt.Fprintf(w, ` +
+ Query:
+
+ +
`, input) +} + +func renderQueryResult(w io.Writer, results resultSetV1, err error, d time.Duration) { + + buf, err2 := json.MarshalIndent(results, "", " ") + + if err != nil { + fmt.Fprintf(w, "Query error (took %v):
%v
", d, err) + } else if err2 != nil { + fmt.Fprintf(w, "JSON marshal error:
%v
", err2) + } else { + fmt.Fprintf(w, "Query results (took %v):
", d) + fmt.Fprintf(w, "
%s
", string(buf)) + } +} + +func renderVersion(w http.ResponseWriter) { + fmt.Fprintln(w, "Version: "+version.Version+"
") + fmt.Fprintln(w, "Build Commit: "+version.Vcs+"
") + fmt.Fprintln(w, "Build Timestamp: "+version.Timestamp+"
") + fmt.Fprintln(w, "Build Hostname: "+version.Hostname+"
") + fmt.Fprintln(w, "
") +} diff --git a/runtime/server_test.go b/runtime/server_test.go new file mode 100644 index 0000000000..53710011f7 --- /dev/null +++ b/runtime/server_test.go @@ -0,0 +1,360 @@ +// 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 runtime + +import ( + "encoding/json" + "io/ioutil" + "net/http" + "net/http/httptest" + "os" + "reflect" + "strings" + "testing" + + "github.com/open-policy-agent/opa/ast" +) + +var policyDir string + +// TestMain creates a temporary direcotry for the server to +// save policies to. The directory name is stored in policyDir +// and is used by the newFixture function. +func TestMain(m *testing.M) { + d, err := ioutil.TempDir("", "server_test") + if err != nil { + panic(err) + } + defer os.RemoveAll(d) + policyDir = d + rc := m.Run() + os.Exit(rc) +} + +func TestDataPatchV1(t *testing.T) { + f := newFixture(t) + patch := newReqV1("PATCH", "/data/x", `[{"op": "add", "path": "/", "value": {"a": 1, "b": 2}}]`) + f.server.Router.ServeHTTP(f.recorder, patch) + + if f.recorder.Code != 204 { + t.Errorf("Expected success/no-content but got %v", f.recorder) + return + } + + get := newReqV1("GET", "/data/x/a", "") + f.reset() + f.server.Router.ServeHTTP(f.recorder, get) + + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + resp := f.loadResponse().(float64) + exp := float64(1) + if resp != exp { + t.Errorf("Expected %v but got: %v", exp, resp) + } +} + +func TestIndexGet(t *testing.T) { + f := newFixture(t) + get, err := http.NewRequest("GET", `/?q=foo = 1`, strings.NewReader("")) + if err != nil { + panic(err) + } + f.server.Router.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 TestPoliciesPutV1(t *testing.T) { + f := newFixture(t) + req := newReqV1("PUT", "/policies/1", testMod) + + f.server.Router.ServeHTTP(f.recorder, req) + + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + policy := f.loadPolicy() + expected := newPolicy("1", testMod) + if !expected.Equal(policy) { + t.Errorf("Expected policies to be equal. Expected:\n\n%v\n\nGot:\n\n%v\n", expected, policy) + } +} + +func TestPoliciesPutV1Empty(t *testing.T) { + f := newFixture(t) + req := newReqV1("PUT", "/policies/1", "") + + f.server.Router.ServeHTTP(f.recorder, req) + + if f.recorder.Code != 400 { + t.Errorf("Expected bad request but got %v", f.recorder) + return + } +} + +func TestPoliciesPutV1ParseError(t *testing.T) { + f := newFixture(t) + req := newReqV1("PUT", "/policies/1", ` + package a.b.c + + p[x] %%^ ;- + `) + + f.server.Router.ServeHTTP(f.recorder, req) + + if f.recorder.Code != 400 { + t.Errorf("Expected bad request but got %v", f.recorder) + return + } +} + +// TODO(tsandall): revisit once safety checks are in place +func testPoliciesPutV1CompileError(t *testing.T) { + f := newFixture(t) + req := newReqV1("PUT", "/policies/1", ` + package a.b.c + p[x] :- q[x] + q[x] :- p[x] + `) + + f.server.Router.ServeHTTP(f.recorder, req) + + if f.recorder.Code != 400 { + t.Errorf("Expected bad request but got %v", f.recorder) + return + } +} + +func TestPoliciesListV1(t *testing.T) { + f := newFixture(t) + put := newReqV1("PUT", "/policies/1", testMod) + f.server.Router.ServeHTTP(f.recorder, put) + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + f.reset() + list := newReqV1("GET", "/policies", "") + + f.server.Router.ServeHTTP(f.recorder, list) + + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + var policies []*policyV1 + err := json.NewDecoder(f.recorder.Body).Decode(&policies) + if err != nil { + t.Errorf("Expected policy list but got error: %v", err) + return + } + + expected := []*policyV1{ + newPolicy("1", testMod), + } + if len(expected) != len(policies) { + t.Errorf("Expected %d policies but got: %v", len(expected), policies) + return + } + for i := range expected { + if !expected[i].Equal(policies[i]) { + t.Errorf("Expected policies to be equal. Expected:\n\n%v\n\nGot:\n\n%v\n", expected[i], policies[i]) + } + } +} + +func TestPoliciesGetV1(t *testing.T) { + f := newFixture(t) + put := newReqV1("PUT", "/policies/1", testMod) + f.server.Router.ServeHTTP(f.recorder, put) + + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + f.reset() + get := newReqV1("GET", "/policies/1", "") + + f.server.Router.ServeHTTP(f.recorder, get) + + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + policy := f.loadPolicy() + expected := newPolicy("1", testMod) + if !expected.Equal(policy) { + t.Errorf("Expected policies to be equal. Expected:\n\n%v\n\nGot:\n\n%v\n", expected, policy) + } +} + +func TestPoliciesGetRawV1(t *testing.T) { + f := newFixture(t) + put := newReqV1("PUT", "/policies/1", testMod) + f.server.Router.ServeHTTP(f.recorder, put) + + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + f.reset() + get := newReqV1("GET", "/policies/1/raw", "") + + f.server.Router.ServeHTTP(f.recorder, get) + + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + raw := f.recorder.Body.String() + if raw != testMod { + t.Errorf("Expected raw string to equal testMod:\n\nExpected:\n\n%v\n\nGot:\n\n%v\n", testMod, raw) + } + +} + +func TestPoliciesDeleteV1(t *testing.T) { + f := newFixture(t) + put := newReqV1("PUT", "/policies/1", testMod) + f.server.Router.ServeHTTP(f.recorder, put) + + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + f.reset() + del := newReqV1("DELETE", "/policies/1", "") + + f.server.Router.ServeHTTP(f.recorder, del) + + if f.recorder.Code != 204 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + f.reset() + get := newReqV1("GET", "/policies/1", "") + f.server.Router.ServeHTTP(f.recorder, get) + if f.recorder.Code != 404 { + t.Errorf("Expected not found but got %v", f.recorder) + return + } +} + +func TestQueryV1(t *testing.T) { + f := newFixture(t) + get := newReqV1("GET", `/query?q=a=[1,2,3],a[i]=x`, "") + f.server.Router.ServeHTTP(f.recorder, get) + + if f.recorder.Code != 200 { + t.Errorf("Expected success but got %v", f.recorder) + return + } + + var expected resultSetV1 + err := json.Unmarshal([]byte(`[{"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 resultSetV1 + err = json.Unmarshal(f.recorder.Body.Bytes(), &result) + if err != nil { + t.Errorf("Unexpected error while unmarshalling result: %v", err) + return + } + + if !reflect.DeepEqual(result, expected) { + t.Errorf("Expected %v but got: %v", expected, result) + } +} + +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 { + runtime *Runtime + server *Server + recorder *httptest.ResponseRecorder + t *testing.T +} + +func newFixture(t *testing.T) *fixture { + runtime := &Runtime{} + runtime.Init(&Params{Server: true, PolicyDir: policyDir}) + server := NewServer(runtime, ":8182", false) + recorder := httptest.NewRecorder() + return &fixture{ + runtime: runtime, + server: server, + recorder: recorder, + t: t, + } +} + +func (f *fixture) loadPolicy() *policyV1 { + policy := &policyV1{} + err := json.NewDecoder(f.recorder.Body).Decode(policy) + if err != nil { + panic(err) + } + return policy +} + +func (f *fixture) loadResponse() interface{} { + var v interface{} + err := json.NewDecoder(f.recorder.Body).Decode(&v) + if err != nil { + panic(err) + } + return v +} + +func (f *fixture) reset() { + f.recorder = httptest.NewRecorder() +} + +func newPolicy(id, s string) *policyV1 { + compiler := ast.NewCompiler() + parsed := ast.MustParseModule(s) + if compiler.Compile(map[string]*ast.Module{"": parsed}); compiler.Failed() { + panic(compiler.FlattenErrors()) + } + mod := compiler.Modules[""] + return &policyV1{ID: id, Module: mod} +} + +func newReqV1(method string, path string, body string) *http.Request { + req, err := http.NewRequest(method, "/v1"+path, strings.NewReader(body)) + if err != nil { + panic(err) + } + return req +}