diff --git a/cmd/parse.go b/cmd/parse.go index 521c48e050..7b2e69ee75 100644 --- a/cmd/parse.go +++ b/cmd/parse.go @@ -7,7 +7,6 @@ package cmd import ( "encoding/json" "errors" - "fmt" "io" "os" "strings" @@ -108,8 +107,7 @@ func parse(args []string, params *parseParams, stdout io.Writer, stderr io.Write _ = pr.JSON(stderr, pr.Output{Errors: pr.NewOutputErrors(err)}) return 1 } - - _, _ = fmt.Fprint(stdout, string(bs)+"\n") + _, _ = stdout.Write(append(bs, '\n')) default: ast.Pretty(stdout, result.Parsed) } diff --git a/v1/ast/compare.go b/v1/ast/compare.go index 1bf990a6ed..fe21d837d5 100644 --- a/v1/ast/compare.go +++ b/v1/ast/compare.go @@ -342,20 +342,16 @@ func TermValueEqual(a, b *Term) bool { func ValueEqual(a, b Value) bool { switch v := a.(type) { - case Null: - return v.Equal(b) - case Boolean: - return v.Equal(b) + case Null, Boolean, String, Var: + return v == b case Number: return v.Equal(b) - case String: - return v.Equal(b) - case Var: - return v.Equal(b) case Ref: return v.Equal(b) case *Array: return v.Equal(b) + case *Not: + return v.Equal(b) case *TemplateString: return v.Equal(b) } diff --git a/v1/ast/compile.go b/v1/ast/compile.go index 8aa60e0f44..f1453dfcc1 100644 --- a/v1/ast/compile.go +++ b/v1/ast/compile.go @@ -1701,16 +1701,17 @@ func (parser *schemaParser) parseSchemaWithPropertyKey(schema any, propertyKey s // Handle referenced schemas, returns directly when a $ref is found if subSchema.RefSchema != nil { - if existing, ok := parser.definitionCache[subSchema.Ref.String()]; ok { + subSchemaStr := subSchema.Ref.String() + if existing, ok := parser.definitionCache[subSchemaStr]; ok { if existing.processing { if existing.rec == nil { - existing.rec = types.NewRecursive(subSchema.Ref.String(), nil) + existing.rec = types.NewRecursive(subSchemaStr, nil) } return existing.rec, nil } return existing.typ, nil } - return parser.parseSchemaWithPropertyKey(subSchema.RefSchema, subSchema.Ref.String()) + return parser.parseSchemaWithPropertyKey(subSchema.RefSchema, subSchemaStr) } // Cache this $ref definition and finalize it via defer when parsing diff --git a/v1/ast/term.go b/v1/ast/term.go index 178bc09e38..5303b0e9a2 100644 --- a/v1/ast/term.go +++ b/v1/ast/term.go @@ -380,15 +380,12 @@ func (term *Term) Copy() *Term { // Equal returns true if this term equals the other term. Equality is // defined for each kind of term, and does not compare the Location. func (term *Term) Equal(other *Term) bool { - if term == nil && other != nil { - return false - } - if term != nil && other == nil { - return false - } if term == other { return true } + if term == nil || other == nil { + return false + } return ValueEqual(term.Value, other.Value) } diff --git a/v1/test/cases/testdata/v1/aggregates/test-membership.yaml b/v1/test/cases/testdata/v1/aggregates/test-membership.yaml index e689b57f2f..c5033eba23 100644 --- a/v1/test/cases/testdata/v1/aggregates/test-membership.yaml +++ b/v1/test/cases/testdata/v1/aggregates/test-membership.yaml @@ -9,7 +9,20 @@ cases: p if { 1 in {1} } - data: { } + data: {} + want_result: + - x: true + - note: aggregates/member check set with with both key and value + query: data.test.p = x + modules: + - | + package test + + p if { + # a set "key" is also its value + "foo", "foo" in {"foo", "bar"} + } + data: {} want_result: - x: true - note: aggregates/member simple, array @@ -21,7 +34,7 @@ cases: p if { 1 in [1] } - data: { } + data: {} want_result: - x: true - note: aggregates/member simple, object @@ -33,7 +46,7 @@ cases: p if { 1 in {"foo": 1} } - data: { } + data: {} want_result: - x: true - note: aggregates/member object with key @@ -45,7 +58,7 @@ cases: p if { "foo", 1 in {"foo": 1} } - data: { } + data: {} want_result: - x: true - note: aggregates/member array with index @@ -57,7 +70,7 @@ cases: p if { 1, "two" in ["one", "two", "three"] } - data: { } + data: {} want_result: - x: true - note: aggregates/member array with index, nested @@ -69,7 +82,7 @@ cases: p if { 1, 2 in [2] in [false, true] } - data: { } + data: {} want_result: - x: true - note: aggregates/member array with index, nested, associativity without parens @@ -81,7 +94,7 @@ cases: p if { (0, 2 in [2]) in [true] } - data: { } + data: {} want_result: - x: true - note: aggregates/member object with key, nested, associativity without parens @@ -93,7 +106,7 @@ cases: p if { ("foo", 2 in {"foo": 2}) in [true] } - data: { } + data: {} want_result: - x: true - note: aggregates/member object with key, nested @@ -105,7 +118,7 @@ cases: p if { "foo", (2 in {"bar": 2}) in {"foo": true} } - data: { } + data: {} want_result: - x: true - note: aggregates/member simple false, set @@ -117,7 +130,7 @@ cases: p := x if { x := 1 in {2} } - data: { } + data: {} want_result: - x: false - note: aggregates/member simple false, array @@ -129,7 +142,7 @@ cases: p := x if { x := 1 in [2] } - data: { } + data: {} want_result: - x: false - note: aggregates/member simple false, object @@ -141,7 +154,7 @@ cases: p := x if { x := 1 in {"foo": 2} } - data: { } + data: {} want_result: - x: false - note: aggregates/member chained @@ -153,7 +166,7 @@ cases: p if { {1, 2} in [{1, 2}] in [true] } - data: { } + data: {} want_result: - x: true - note: aggregates/member with vars @@ -167,7 +180,7 @@ cases: xs := ["foo", "bar"] x in xs } - data: { } + data: {} want_result: - x: true - note: aggregates/member with not @@ -179,7 +192,7 @@ cases: p if { not "foo" in ["fox"] } - data: { } + data: {} want_result: - x: true - note: aggregates/member operator precedence with other infix operator (+) @@ -191,7 +204,7 @@ cases: p if { (1 + 1) in [2] } - data: { } + data: {} want_result: - x: true - note: aggregates/member operator precedence in list (set) @@ -204,7 +217,7 @@ cases: x := {1, 1 in [2]} x == {1, false} } - data: { } + data: {} want_result: - x: true - note: aggregates/member operator precedence in list with parens (set) @@ -217,7 +230,7 @@ cases: x := {(1, 1 in [2])} x == {false} } - data: { } + data: {} want_result: - x: true - note: aggregates/member operator precedence in list (array) @@ -230,7 +243,7 @@ cases: x := [1, 1 in [2]] x == [1, false] } - data: { } + data: {} want_result: - x: true - note: aggregates/member operator precedence in list with parens (array) @@ -243,7 +256,7 @@ cases: x := [(1, 1 in [2])] x == [false] } - data: { } + data: {} want_result: - x: true - note: aggregates/member operator precedence in list (fun args) @@ -257,7 +270,7 @@ cases: } f(_, _) := true - data: { } + data: {} want_result: - x: true - note: aggregates/member operator precedence in list with parens (fun args) @@ -271,7 +284,7 @@ cases: } f(_) := true - data: { } + data: {} want_result: - x: true - note: aggregates/member composite containee @@ -283,7 +296,7 @@ cases: p if { {"foo": {"baz": 2000}} in [{"foo": {"baz": 2000}}] } - data: { } + data: {} want_result: - x: true - note: aggregates/member non-collection string @@ -295,7 +308,7 @@ cases: p := x if { x := 1 in "foo" } - data: { } + data: {} want_result: - x: false - note: aggregates/member non-collection number @@ -307,7 +320,7 @@ cases: p := x if { x := "foo" in 1 } - data: { } + data: {} want_result: - x: false - note: aggregates/member with key in non-collection (number) @@ -319,7 +332,7 @@ cases: p := x if { x := (1, "foo" in 1) } - data: { } + data: {} want_result: - x: false - note: aggregates/member+some simple, array @@ -331,7 +344,7 @@ cases: p contains x if { some x in [1, 2, 3] } - data: { } + data: {} want_result: - x: - 1 @@ -346,7 +359,7 @@ cases: p if { some "foo" in ["foo"] } - data: { } + data: {} want_result: - x: true - note: aggregates/member+some containee is call @@ -358,7 +371,7 @@ cases: p if { some numbers.range(1, 1) in [[1]] } - data: { } + data: {} want_result: - x: true - note: aggregates/member+some non-ground composite containee @@ -370,7 +383,7 @@ cases: p := x if { some {"foo": x} in [{"foo": 100}, {"what": "ever"}] } - data: { } + data: {} want_result: - x: 100 - note: aggregates/member+some non-ground composite containee, multiple bindings @@ -382,7 +395,7 @@ cases: p := x if { some {"foo": x, "what": y} in [{"foo": 100, "what": "ever"}] } - data: { } + data: {} want_result: - x: 100 - note: aggregates/member+some ground composite containee @@ -394,7 +407,7 @@ cases: p if { some {"foo": 100} in [{"foo": 100}] } - data: { } + data: {} want_result: - x: true - note: aggregates/member+some ground composite containee (false) @@ -406,7 +419,7 @@ cases: p if { some {"foo": 0} in [{"foo": 100}] } - data: { } + data: {} want_result: [] - note: aggregates/member+some+key non-ground value query: data.test.p = x @@ -417,7 +430,7 @@ cases: p := x if { some "foo", x in {"foo": 100, "what": "ever"} } - data: { } + data: {} want_result: - x: 100 - note: aggregates/member+some+key non-ground key @@ -429,7 +442,7 @@ cases: p := x if { some x, "ever" in {"foo": 100, "what": "ever"} } - data: { } + data: {} want_result: - x: what - note: aggregates/member+some+key non-ground key+value @@ -441,7 +454,7 @@ cases: p[k] := v if { some k, v in {"foo": 100, "what": "ever"} } - data: { } + data: {} want_result: - x: foo: 100 @@ -455,7 +468,7 @@ cases: p := x if { some {"foo": x}, "ever" in {{"foo": 100}: "ever"} } - data: { } + data: {} want_result: - x: 100 - note: aggregates/member+some+ref @@ -522,7 +535,7 @@ cases: p contains [k, v] if { some k, v in input with input.foo as "bar" } - data: { } + data: {} want_result: - x: - - foo diff --git a/v1/topdown/aggregates.go b/v1/topdown/aggregates.go index 356e7b38b0..95f8d708aa 100644 --- a/v1/topdown/aggregates.go +++ b/v1/topdown/aggregates.go @@ -260,34 +260,29 @@ func builtinAny(_ BuiltinContext, operands []*ast.Term, iter func(*ast.Term) err } func builtinMember(_ BuiltinContext, operands []*ast.Term, iter func(*ast.Term) error) error { - containee := operands[0] switch c := operands[1].Value.(type) { case ast.Set: - return iter(ast.InternedTerm(c.Contains(containee))) + return iter(ast.InternedTerm(c.Contains(operands[0]))) case *ast.Array: - for i := range c.Len() { - if c.Elem(i).Value.Compare(containee.Value) == 0 { - return iter(ast.InternedTerm(true)) - } - } - return iter(ast.InternedTerm(false)) + return iter(ast.InternedTerm(c.Until(operands[0].Equal))) case ast.Object: return iter(ast.InternedTerm(c.Until(func(_, v *ast.Term) bool { - return v.Value.Compare(containee.Value) == 0 + return operands[0].Equal(v) }))) } return iter(ast.InternedTerm(false)) } func builtinMemberWithKey(_ BuiltinContext, operands []*ast.Term, iter func(*ast.Term) error) error { - key, val := operands[0], operands[1] - switch c := operands[2].Value.(type) { - case interface{ Get(*ast.Term) *ast.Term }: - ret := false - if act := c.Get(key); act != nil { - ret = act.Value.Compare(val.Value) == 0 - } - return iter(ast.InternedTerm(ret)) + type getter interface { + Get(*ast.Term) *ast.Term + } + col, key, val := operands[2], operands[0], operands[1] + switch c := col.Value.(type) { + case ast.Set: + return iter(ast.InternedTerm(c.Contains(key) && key.Equal(val))) + case getter: + return iter(ast.InternedTerm(val.Equal(c.Get(key)))) } return iter(ast.InternedTerm(false)) } diff --git a/v1/topdown/aggregates_bench_test.go b/v1/topdown/aggregates_bench_test.go index c32a8b808e..23d353b6c7 100644 --- a/v1/topdown/aggregates_bench_test.go +++ b/v1/topdown/aggregates_bench_test.go @@ -2,6 +2,7 @@ package topdown import ( "fmt" + "slices" "testing" "github.com/open-policy-agent/opa/v1/ast" @@ -117,3 +118,103 @@ func BenchmarkSumFloatSet(b *testing.B) { } } } + +// BenchmarkMember/set-16 81601833 14.58 ns/op 0 B/op 0 allocs/op +// BenchmarkMember/array-16 579356 1988 ns/op 0 B/op 0 allocs/op +// BenchmarkMember/object-16 90483236 13.14 ns/op 0 B/op 0 allocs/op +func BenchmarkMember(b *testing.B) { + key := ast.InternedTerm(99) + val := key + exp := eqIter(ast.InternedTerm(true)) + + b.Run("set", func(b *testing.B) { + bcx := BuiltinContext{} + set := ast.SetTerm(slices.Collect(ast.InternedIntRange(1, 100))...) + ops := []*ast.Term{key, set} + + for b.Loop() { + if err := builtinMember(bcx, ops, exp); err != nil { + b.Fatalf("unexpected error: %v", err) + } + } + }) + + b.Run("array", func(b *testing.B) { + bcx := BuiltinContext{} + arr := ast.ArrayTerm(slices.Collect(ast.InternedIntRange(1, 100))...) + ops := []*ast.Term{key, arr} + + for b.Loop() { + if err := builtinMember(bcx, ops, exp); err != nil { + b.Fatalf("unexpected error: %v", err) + } + } + }) + + b.Run("object", func(b *testing.B) { + bcx := BuiltinContext{} + obj := ast.NewObject() + for i := 1; i <= 100; i++ { + obj.Insert(ast.InternedTerm(i), ast.InternedTerm(i)) + } + ops := []*ast.Term{key, val, ast.NewTerm(obj)} + + for b.Loop() { + if err := builtinMemberWithKey(bcx, ops, exp); err != nil { + b.Fatalf("unexpected error: %v", err) + } + } + }) +} + +// Was +// BenchmarkMemberWithKey/set-16 error +// BenchmarkMemberWithKey/array-16 58668218 20.41 ns/op 0 B/op 0 allocs/op +// BenchmarkMemberWithKey/object-16 61149481 20.33 ns/op 0 B/op 0 allocs/op +// +// Now +// BenchmarkMemberWithKey/set-16 79370059 15.04 ns/op 0 B/op 0 allocs/op +// BenchmarkMemberWithKey/array-16 97242430 12.00 ns/op 0 B/op 0 allocs/op +// BenchmarkMemberWithKey/object-16 96117901 12.63 ns/op 0 B/op 0 allocs/op +func BenchmarkMemberWithKey(b *testing.B) { + key, val := ast.InternedTerm(99), ast.InternedTerm(99) + exp := eqIter(ast.InternedTerm(true)) + + b.Run("set", func(b *testing.B) { + bcx := BuiltinContext{} + ops := []*ast.Term{key, key, ast.SetTerm(slices.Collect(ast.InternedIntRange(1, 100))...)} + + for b.Loop() { + if err := builtinMemberWithKey(bcx, ops, exp); err != nil { + b.Fatalf("unexpected error: %v", err) + } + } + }) + + b.Run("array", func(b *testing.B) { + bcx := BuiltinContext{} + key := ast.InternedTerm(98) + ops := []*ast.Term{key, val, ast.ArrayTerm(slices.Collect(ast.InternedIntRange(1, 100))...)} + + for b.Loop() { + if err := builtinMemberWithKey(bcx, ops, exp); err != nil { + b.Fatalf("unexpected error: %v", err) + } + } + }) + + b.Run("object", func(b *testing.B) { + bcx := BuiltinContext{} + obj := ast.NewObject() + for i := 1; i <= 100; i++ { + obj.Insert(ast.InternedTerm(i), ast.InternedTerm(i)) + } + ops := []*ast.Term{key, val, ast.NewTerm(obj)} + + for b.Loop() { + if err := builtinMemberWithKey(bcx, ops, exp); err != nil { + b.Fatalf("unexpected error: %v", err) + } + } + }) +} diff --git a/wasm/src/aggregates.c b/wasm/src/aggregates.c index 0b811edba8..df10e31542 100644 --- a/wasm/src/aggregates.c +++ b/wasm/src/aggregates.c @@ -363,6 +363,7 @@ opa_value *builtin_member3(opa_value *key, opa_value *val, opa_value *collection { case OPA_ARRAY: case OPA_OBJECT: + case OPA_SET: return opa_boolean(opa_value_compare(val, opa_value_get(collection, key)) == 0); } return opa_boolean(false); diff --git a/wasm/src/aggregates.h b/wasm/src/aggregates.h index 38660256be..c4b19af7db 100644 --- a/wasm/src/aggregates.h +++ b/wasm/src/aggregates.h @@ -12,4 +12,7 @@ opa_value *opa_agg_sort(opa_value *v); opa_value *opa_agg_all(opa_value *v); opa_value *opa_agg_any(opa_value *v); +opa_value *builtin_member(opa_value *v, opa_value *collection); +opa_value *builtin_member3(opa_value *key, opa_value *val, opa_value *collection); + #endif diff --git a/wasm/tests/test.c b/wasm/tests/test.c index e313a9326f..221a85b442 100644 --- a/wasm/tests/test.c +++ b/wasm/tests/test.c @@ -2062,6 +2062,14 @@ void test_aggregates(void) test("any/set trues", opa_value_compare(opa_agg_any(&set_trues->hdr), opa_boolean(true)) == 0); test("any/set mixed", opa_value_compare(opa_agg_any(&set_mixed->hdr), opa_boolean(true)) == 0); test("any/set falses", opa_value_compare(opa_agg_any(&set_falses->hdr), opa_boolean(false)) == 0); + + opa_set_t *set_foo_bar = opa_cast_set(opa_set()); + opa_set_add(set_foo_bar, opa_string_terminated("foo")); + opa_set_add(set_foo_bar, opa_string_terminated("bar")); + + test("member3/set key==val found", opa_value_compare(builtin_member3(opa_string_terminated("foo"), opa_string_terminated("foo"), &set_foo_bar->hdr), opa_boolean(true)) == 0); + test("member3/set key==val not found", opa_value_compare(builtin_member3(opa_string_terminated("baz"), opa_string_terminated("baz"), &set_foo_bar->hdr), opa_boolean(false)) == 0); + test("member3/set key!=val", opa_value_compare(builtin_member3(opa_string_terminated("foo"), opa_string_terminated("bar"), &set_foo_bar->hdr), opa_boolean(false)) == 0); } WASM_EXPORT(test_base64)