topdown: fix "a", "a" in {"a"} not returning true (#8747)

It's mostly useless, but aren't we all.

Also added benchmarks to make sure I didn't mess anything up. And one or
two tiny but unrelated fixes.

---------

Signed-off-by: Anders Eknert <anders.eknert@apple.com>
Signed-off-by: Stephan Renatus <stephan.renatus@gmail.com>
Co-authored-by: Stephan Renatus <stephan.renatus@gmail.com>
This commit is contained in:
Anders Eknert
2026-06-26 10:04:30 +02:00
committed by GitHub
parent 8b59ff6e48
commit 2b18f03b2b
10 changed files with 188 additions and 75 deletions
+1 -3
View File
@@ -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)
}
+4 -8
View File
@@ -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)
}
+4 -3
View File
@@ -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
+3 -6
View File
@@ -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)
}
+51 -38
View File
@@ -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
+12 -17
View File
@@ -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))
}
+101
View File
@@ -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)
}
}
})
}
+1
View File
@@ -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);
+3
View File
@@ -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
+8
View File
@@ -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)