diff --git a/ast/parser_ext.go b/ast/parser_ext.go index 4980acaf01..4077ab2e20 100644 --- a/ast/parser_ext.go +++ b/ast/parser_ext.go @@ -166,6 +166,7 @@ func ParseStatements(input string) ([]interface{}, error) { return nil, err } stmts := parsed.([]interface{}) + postProcess(stmts) return stmts, err } @@ -266,6 +267,18 @@ func parseModule(stmts []interface{}) (*Module, error) { return mod, nil } +func postProcess(stmts []interface{}) { + mangleWildcards(stmts) +} + +func mangleWildcards(stmts []interface{}) { + + mangler := &wildcardMangler{} + for _, stmt := range stmts { + Walk(mangler, stmt) + } +} + type wildcardMangler struct { c int } diff --git a/ast/parser_test.go b/ast/parser_test.go index cb4bc11c00..5ca0882255 100644 --- a/ast/parser_test.go +++ b/ast/parser_test.go @@ -140,6 +140,47 @@ func TestCompositesWithRefs(t *testing.T) { assertParseOneTerm(t, "ref values", "[{8: a[i].b, f: c[0][\"d\"].e[j]}]", ArrayTerm(ObjectTerm(Item(NumberTerm(8), ref1), Item(VarTerm("f"), ref2)))) } +func TestArrayComprehensions(t *testing.T) { + + input := `[ + {"x": [a[i] | xs = [{"a": ["baz", j]} | q[p], p.a != "bar", j = "foo"], + xs[j].a[k] = "foo"]} + ]` + + expected := ArrayTerm( + ObjectTerm(Item( + StringTerm("x"), + ArrayComprehensionTerm( + RefTerm(VarTerm("a"), VarTerm("i")), + Body{ + NewBuiltinExpr( + VarTerm("="), + VarTerm("xs"), + ArrayComprehensionTerm( + ObjectTerm(Item(StringTerm("a"), ArrayTerm(StringTerm("baz"), VarTerm("j")))), + Body{ + &Expr{ + Terms: RefTerm(VarTerm("q"), VarTerm("p")), + }, + NewBuiltinExpr(VarTerm("!="), RefTerm(VarTerm("p"), StringTerm("a")), StringTerm("bar")), + NewBuiltinExpr(VarTerm("="), VarTerm("j"), StringTerm("foo")), + }, + ), + ), + NewBuiltinExpr( + VarTerm("="), + RefTerm(VarTerm("xs"), VarTerm("j"), StringTerm("a"), VarTerm("k")), + StringTerm("foo"), + ), + }, + ), + )), + ) + + assertParseOneTerm(t, "nested", input, expected) + +} + func TestInfixExpr(t *testing.T) { assertParseOneExpr(t, "scalars 1", "true = false", NewBuiltinExpr(VarTerm("="), BooleanTerm(true), BooleanTerm(false))) assertParseOneExpr(t, "scalars 2", "3.14 = null", NewBuiltinExpr(VarTerm("="), NumberTerm(3.14), NullTerm())) @@ -295,6 +336,10 @@ func TestComments(t *testing.T) { :- m = [1,2, 3], a = m[i] + + r[x] :- x = [ a | # inside comprehension + a = z[i], + b[i].a = a ] ` assertParseModule(t, "module comments", testModule, &Module{ @@ -307,6 +352,7 @@ func TestComments(t *testing.T) { Rules: []*Rule{ MustParseStatement("p[x] = y :- y = \"foo\", x = \"bar\", x != y, q[x]").(*Rule), MustParseStatement("q[a] :- m = [1,2,3], a = m[i]").(*Rule), + MustParseStatement("r[x] :- x = [a | a = z[i], b[i].a = a]").(*Rule), }, }) } @@ -416,6 +462,25 @@ func TestWildcards(t *testing.T) { ), }, }) + + assertParseOneExpr(t, "comprehension", "_ = [x | a = a[_]]", &Expr{ + Terms: []*Term{ + VarTerm("="), + VarTerm("$0"), + ArrayComprehensionTerm( + VarTerm("x"), + Body{ + &Expr{ + Terms: []*Term{ + VarTerm("="), + VarTerm("a"), + RefTerm(VarTerm("a"), VarTerm("$1")), + }, + }, + }, + ), + }, + }) } func assertParse(t *testing.T, msg string, input string, correct func([]interface{})) { @@ -485,7 +550,7 @@ func assertParseOneExpr(t *testing.T, msg string, input string, correct *Expr) { } expr := body[0] if !expr.Equal(correct) { - t.Errorf("Error on test %s: expressions not equal: %v (parsed), %v (correct)", msg, expr, correct) + t.Errorf("Error on test %s: expressions not equal:\n%v (parsed)\n%v (correct)", msg, expr, correct) } }) } diff --git a/ast/policy.go b/ast/policy.go index 75204e20b7..5506a20a82 100644 --- a/ast/policy.go +++ b/ast/policy.go @@ -400,32 +400,7 @@ func (expr *Expr) UnmarshalJSON(bs []byte) error { if err := json.Unmarshal(bs, &v); err != nil { return err } - - if x, ok := v["Negated"]; ok { - if b, ok := x.(bool); ok { - expr.Negated = b - } else { - return fmt.Errorf("ast: unable to unmarshal Negated field with type: %T (expected true or false)", v["Negated"]) - } - } - - 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{}: - terms, err := unmarshalTermSlice(ts) - if err != nil { - return err - } - expr.Terms = terms - default: - return fmt.Errorf(`ast: unable to unmarshal Terms field with type: %T (expected {"Value": ..., "Type": ...} or [{"Value": ..., "Type": ...}, ...])`, v["Terms"]) - } - return nil + return unmarshalExpr(expr, v) } // Vars returns a VarSet containing all of the variables in the expression. diff --git a/ast/policy_test.go b/ast/policy_test.go index 3751b3e3d2..a8661ff83a 100644 --- a/ast/policy_test.go +++ b/ast/policy_test.go @@ -19,6 +19,7 @@ func TestModuleJSONRoundTrip(t *testing.T) { 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 + t = true :- xs = [{"x": a[i].a} | a[i].n = "bob", b[x]] `) bs, err := json.Marshal(mod) diff --git a/ast/rego.peg b/ast/rego.peg index bae81804e9..2d6c778f3f 100644 --- a/ast/rego.peg +++ b/ast/rego.peg @@ -90,7 +90,6 @@ Import <- "import" ws path:(Ref / Var) alias:(ws "as" ws Var)? { return imp, nil } -// TODO(tsandall): update to handle underscore variables Rule <- name:Var key:( _ "[" _ Term _ "]" _ )? value:( _ "=" _ Term )? body:( _ ":-" _ Body) { rule := &Rule{} @@ -130,8 +129,6 @@ Body <- head:Expr tail:( _ "," _ Expr)* { expr := s.([]interface{})[3].(*Expr) buf = append(buf, expr) } - mangler := &wildcardMangler{} - Walk(mangler, buf) return buf, nil } @@ -169,10 +166,18 @@ PrefixExpr <- op:Var "(" _ head:Term? tail:( _ "," _ Term )* _ ")" { return buf, nil } -Term <- val:( Composite / Scalar / Ref / Var ) { +Term <- val:( Comprehension / Composite / Scalar / Ref / Var ) { return val, nil } +Comprehension <- ArrayComprehension + +ArrayComprehension <- "[" _ term:Term _ "|" _ body:Body _ "]" { + ac := ArrayComprehensionTerm(term.(*Term), body.(Body)) + ac.Location = currentLocation(c) + return ac, nil +} + Composite <- Object / Array Scalar <- Number / String / Bool / Null diff --git a/ast/term.go b/ast/term.go index 82ae3ac891..b98841d82b 100644 --- a/ast/term.go +++ b/ast/term.go @@ -34,6 +34,7 @@ func NewLocation(text []byte, file string, row int, col int) *Location { // - Object, Array // - Variables // - References +// - Array Comprehensions // type Value interface { // Equal returns true if this value equals the other value. @@ -70,6 +71,11 @@ func (term *Term) Equal(other *Term) bool { return term.Value.Equal(other.Value) } +// Hash returns the hash code of the Term's value. +func (term *Term) Hash() int { + return term.Value.Hash() +} + // IsGround returns true if this terms' Value is ground. func (term *Term) IsGround() bool { return term.Value.IsGround() @@ -97,6 +103,8 @@ func (term *Term) MarshalJSON() ([]byte, error) { typ = "array" case Object: typ = "object" + case *ArrayComprehension: + typ = "array-comprehension" } d := map[string]interface{}{ "Type": typ, @@ -582,6 +590,48 @@ func (obj Object) queryRec(ref Ref, keys map[Var]Value, iter QueryIterator) erro } } +// ArrayComprehension represents an array comprehension as defined in the language. +type ArrayComprehension struct { + Term *Term + Body Body +} + +// ArrayComprehensionTerm creates a new Term with an ArrayComprehension value. +func ArrayComprehensionTerm(term *Term, body Body) *Term { + return &Term{ + Value: &ArrayComprehension{ + Term: term, + Body: body, + }, + } +} + +// Equal returns true if this array comprehension is syntactically equal to another. +func (ac *ArrayComprehension) Equal(other Value) bool { + if ac == other { + return true + } + o, ok := other.(*ArrayComprehension) + if !ok { + return false + } + return o.Term.Equal(ac.Term) && o.Body.Equal(ac.Body) +} + +// Hash returns the hash code of the Value. +func (ac *ArrayComprehension) Hash() int { + return ac.Term.Hash() + ac.Body.Hash() +} + +// IsGround returns true if the Term and Body are ground. +func (ac *ArrayComprehension) IsGround() bool { + return ac.Term.IsGround() && ac.Body.IsGround() +} + +func (ac *ArrayComprehension) String() string { + return "[" + ac.Term.String() + " | " + ac.Body.String() + "]" +} + func queryRec(v Value, ref Ref, tail Ref, keys map[Var]Value, iter QueryIterator, skipScalar bool) error { if len(tail) == 0 { if err := iter(keys, v); err != nil { @@ -631,7 +681,7 @@ func termSliceIsGround(a []*Term) bool { return true } -// TODO(tsandall): The unmarshalling errors in these functions are not +// NOTE(tsandall): The unmarshalling errors in these functions are not // helpful for callers because they do not identify the source of the // unmarshalling error. Because OPA doesn't accept JSON describing ASTs // from callers, this is acceptable (for now). If that changes in the future, @@ -639,12 +689,64 @@ func termSliceIsGround(a []*Term) bool { // on the happy path and treats all errors the same. If better error // reporting is needed, the error paths will need to be fleshed out. +func unmarshalBody(b []interface{}) (Body, error) { + buf := Body{} + for _, e := range b { + if m, ok := e.(map[string]interface{}); ok { + expr := &Expr{} + if err := unmarshalExpr(expr, m); err == nil { + buf = append(buf, expr) + continue + } + } + goto unmarshal_error + } + return buf, nil +unmarshal_error: + return nil, fmt.Errorf("ast: unable to unmarshal body") +} + +func unmarshalExpr(expr *Expr, v map[string]interface{}) error { + if x, ok := v["Negated"]; ok { + if b, ok := x.(bool); ok { + expr.Negated = b + } else { + return fmt.Errorf("ast: unable to unmarshal Negated field with type: %T (expected true or false)", v["Negated"]) + } + } + switch ts := v["Terms"].(type) { + case map[string]interface{}: + t, err := unmarshalTerm(ts) + if err != nil { + return err + } + expr.Terms = t + case []interface{}: + terms, err := unmarshalTermSlice(ts) + if err != nil { + return err + } + expr.Terms = terms + default: + return fmt.Errorf(`ast: unable to unmarshal Terms field with type: %T (expected {"Value": ..., "Type": ...} or [{"Value": ..., "Type": ...}, ...])`, v["Terms"]) + } + return nil +} + +func unmarshalTerm(m map[string]interface{}) (*Term, error) { + v, err := unmarshalValue(m) + if err != nil { + return nil, err + } + return &Term{Value: v}, nil +} + func unmarshalTermSlice(s []interface{}) ([]*Term, error) { buf := []*Term{} for _, x := range s { if m, ok := x.(map[string]interface{}); ok { - if v, err := unmarshalValue(m); err == nil { - buf = append(buf, &Term{Value: v}) + if t, err := unmarshalTerm(m); err == nil { + buf = append(buf, t) continue } } @@ -704,6 +806,22 @@ func unmarshalValue(d map[string]interface{}) (Value, error) { } return buf, nil } + case "array-comprehension": + if m, ok := v.(map[string]interface{}); ok { + if t, ok := m["Term"].(map[string]interface{}); ok { + if term, err := unmarshalTerm(t); err == nil { + if b, ok := m["Body"].([]interface{}); ok { + if body, err := unmarshalBody(b); err == nil { + buf := &ArrayComprehension{ + Term: term, + Body: body, + } + return buf, nil + } + } + } + } + } } unmarshal_error: return nil, fmt.Errorf("ast: unable to unmarshal term") diff --git a/ast/term_test.go b/ast/term_test.go index e855438c60..68693fd99c 100644 --- a/ast/term_test.go +++ b/ast/term_test.go @@ -141,6 +141,7 @@ func TestTermEqual(t *testing.T) { assertTermEqual(t, ArrayTerm(NumberTerm(1), NumberTerm(2), NumberTerm(3)), ArrayTerm(NumberTerm(1), NumberTerm(2), NumberTerm(3))) assertTermEqual(t, VarTerm("foo"), VarTerm("foo")) assertTermEqual(t, RefTerm(VarTerm("foo"), VarTerm("i"), NumberTerm(2)), RefTerm(VarTerm("foo"), VarTerm("i"), NumberTerm(2))) + assertTermEqual(t, ArrayComprehensionTerm(VarTerm("x"), Body{&Expr{Terms: RefTerm(VarTerm("a"), VarTerm("i"))}}), ArrayComprehensionTerm(VarTerm("x"), Body{&Expr{Terms: RefTerm(VarTerm("a"), VarTerm("i"))}})) assertTermNotEqual(t, NullTerm(), BooleanTerm(true)) assertTermNotEqual(t, BooleanTerm(true), BooleanTerm(false)) assertTermNotEqual(t, NumberTerm(5), NumberTerm(7)) @@ -153,6 +154,7 @@ func TestTermEqual(t *testing.T) { assertTermNotEqual(t, ArrayTerm(NumberTerm(1), NumberTerm(2), NumberTerm(3)), ArrayTerm(NumberTerm(1), NumberTerm(2), NumberTerm(4))) assertTermNotEqual(t, VarTerm("foo"), VarTerm("bar")) assertTermNotEqual(t, RefTerm(VarTerm("foo"), VarTerm("i"), NumberTerm(2)), RefTerm(VarTerm("foo"), StringTerm("i"), NumberTerm(2))) + assertTermNotEqual(t, ArrayComprehensionTerm(VarTerm("x"), Body{&Expr{Terms: RefTerm(VarTerm("a"), VarTerm("j"))}}), ArrayComprehensionTerm(VarTerm("x"), Body{&Expr{Terms: RefTerm(VarTerm("a"), VarTerm("i"))}})) } func TestHash(t *testing.T) { @@ -164,7 +166,8 @@ func TestHash(t *testing.T) { ], "e": { 100: a[i].b - } + }, + "k": [ "foo" | true ] } ` @@ -195,6 +198,7 @@ func TestTermString(t *testing.T) { assertToString(t, ArrayTerm().Value, "[]") assertToString(t, ObjectTerm().Value, "{}") assertToString(t, ArrayTerm(ObjectTerm(Item(VarTerm("foo"), ArrayTerm(RefTerm(VarTerm("bar"), VarTerm("i"))))), StringTerm("foo"), BooleanTerm(true), NullTerm(), NumberTerm(42.1)).Value, "[{foo: [bar[i]]}, \"foo\", true, null, 42.1]") + assertToString(t, ArrayComprehensionTerm(ArrayTerm(VarTerm("x")), Body{&Expr{Terms: RefTerm(VarTerm("a"), VarTerm("i"))}}).Value, "[[x] | a[i]]") } func TestRefUnderlying(t *testing.T) { diff --git a/ast/visit.go b/ast/visit.go index 3251660862..f249348c04 100644 --- a/ast/visit.go +++ b/ast/visit.go @@ -72,5 +72,8 @@ func Walk(v Visitor, x interface{}) { for _, t := range x { Walk(w, t.Value) } + case *ArrayComprehension: + Walk(w, x.Term) + Walk(w, x.Body) } } diff --git a/ast/visit_test.go b/ast/visit_test.go index d295692242..bbdf87cbcb 100644 --- a/ast/visit_test.go +++ b/ast/visit_test.go @@ -20,7 +20,10 @@ func TestVisitor(t *testing.T) { rule := MustParseModule(` package a.b import x.y as z - t[x] = y :- p[x] = {"foo": [y,2,{"bar": 3}]}, not q[x] + t[x] = y :- + p[x] = {"foo": [y,2,{"bar": 3}]}, + not q[x], + y = [ [x,z] | x = "x", z = "z" ] `) vis := &testVis{} Walk(vis, rule) @@ -59,9 +62,25 @@ func TestVisitor(t *testing.T) { ref2 q x + expr3 + = + y + compr + array + x + z + body + expr4 + = + x + "x" + expr5 + = + z + "z" */ - if len(vis.elems) != 33 { - t.Errorf("Expected exactly 33 elements in AST but got %d: %v", len(vis.elems), vis.elems) + if len(vis.elems) != 49 { + t.Errorf("Expected exactly 49 elements in AST but got %d: %v", len(vis.elems), vis.elems) } }