From c416fd964ef326186cf7c45e84ff241f53d0a377 Mon Sep 17 00:00:00 2001 From: Sebastian Spaink Date: Fri, 6 Mar 2026 14:08:34 -0600 Subject: [PATCH] Revert "ast: make rule index track var assignments and `x in {...}` (#8341)" (#8410) This reverts commit 32b97ae08c94ac9cdcb434678190c77678ef89c5. Signed-off-by: Sebastian Spaink --- v1/ast/index.go | 402 ++++++++------------------------ v1/ast/index_debug.go | 219 ------------------ v1/ast/index_test.go | 525 +----------------------------------------- 3 files changed, 103 insertions(+), 1043 deletions(-) delete mode 100644 v1/ast/index_debug.go diff --git a/v1/ast/index.go b/v1/ast/index.go index acddd0776d..663072ed83 100644 --- a/v1/ast/index.go +++ b/v1/ast/index.go @@ -5,6 +5,7 @@ package ast import ( + "fmt" "slices" "sort" "strings" @@ -67,7 +68,6 @@ var ( globMatchRef = GlobMatch.Ref() internalPrintRef = InternalPrint.Ref() internalTestCaseRef = InternalTestCase.Ref() - internalMemberRef = Member.Ref() skipIndexing = NewSet(NewTerm(internalPrintRef), NewTerm(internalTestCaseRef)) ) @@ -87,7 +87,6 @@ func (i *baseDocEqIndex) Build(rules []*Rule) bool { i.kind = rules[0].Head.RuleKind() indices := newrefindices(i.isVirtual) - values := make(map[Var]Value) // build indices for each rule. for idx := range rules { @@ -107,9 +106,8 @@ func (i *baseDocEqIndex) Build(rules []*Rule) bool { } } if !skip { - clear(values) for i := range rule.Body { - indices.Update(rule, rule.Body[i], values) + indices.Update(rule, rule.Body[i]) } } return false @@ -126,46 +124,7 @@ func (i *baseDocEqIndex) Build(rules []*Rule) bool { node := i.root if indices.Indexed(rule) { for _, ref := range indices.Sorted() { - var values []*refindex - for _, ri := range indices.rules[rule] { - if ri.Ref.Equal(ref) { - values = append(values, ri) - } - } - if len(values) == 0 { - node = node.Insert(ref, nil, nil) - } else if len(values) == 1 { - node = node.Insert(ref, values[0].Value, values[0].Mapper) - } else { - var hasVar bool - for i := range values { - if _, isVar := values[i].Value.(Var); isVar { - hasVar = true - break - } - } - - if hasVar { - child := node.Insert(ref, anyValue, values[0].Mapper) - for i := range values { - if values[i].Mapper != nil { - node.next.addMapper(values[i].Mapper) - } - } - node = child - } else { - // When a rule has multiple scalar values (e.g., internal.member_2 with a set), - // each value should have its own child node, and the rule is appended to each. - // This creates separate paths for each value so different rules with overlapping - // values don't interfere with each other. - for _, val := range values { - child := node.Insert(ref, val.Value, val.Mapper) - child.append([...]int{idx, prio}, rule) - } - prio++ - return false - } - } + node = node.Insert(ref, indices.Value(rule, ref), indices.Mapper(rule, ref)) } } // Insert rule into trie with (insertion order, priority order) @@ -333,7 +292,7 @@ var anyValue = Var("__any__") // Update attempts to update the refindices for the given expression in the // given rule. If the expression cannot be indexed the update does not affect // the indices. -func (i *refindices) Update(rule *Rule, expr *Expr, values map[Var]Value) { +func (i *refindices) Update(rule *Rule, expr *Expr) { if len(expr.With) > 0 { // NOTE(tsandall): In the future, we may need to consider expressions @@ -355,44 +314,30 @@ func (i *refindices) Update(rule *Rule, expr *Expr, values map[Var]Value) { // function with a undefined argument, there's no point to recording // "needs to be anything" for function args if ref, ok := ts.Value.(Ref); ok { // "naked ref" - i.updateEq(rule, ref, anyValue, nil) + i.updateEq(rule, ref, anyValue) } } } - equalish := op.Equal(equalityRef) || // unification, no 3-operands version exists + a, b := expr.Operand(0), expr.Operand(1) + switch { + case op.Equal(equalityRef): + i.updateEq(rule, a.Value, b.Value) + + case op.Equal(equalRef) && len(expr.Operands()) == 2: // NOTE(tsandall): if equal() is called with more than two arguments the // output value is being captured in which case the indexer cannot // exclude the rule if the equal() call would return false (because the // false value must still be produced.) - (op.Equal(equalRef) && len(expr.Operands()) == 2) - - a, b := expr.Operand(0), expr.Operand(1) - switch { - case equalish: - if !i.updateEqWildcardRef(rule, a.Value, b.Value, values) { - i.updateEq(rule, a.Value, b.Value, values) - } + i.updateEq(rule, a.Value, b.Value) case op.Equal(globMatchRef) && len(expr.Operands()) == 3: // NOTE(sr): Same as with equal() above -- 4 operands means the output // of `glob.match` is captured and the rule can thus not be excluded. i.updateGlobMatch(rule, expr) - - case op.Equal(internalMemberRef) && len(expr.Operands()) == 2: - // NOTE(sr): Again, 3 operands means captured output (like above). - i.updateMember(rule, expr, values) } } -func (i *refindices) isValidIndexRef(ref Ref) bool { - // NB(sr): the ordering is intentional, cheapest-first - return RootDocumentNames.Contains(ref[0]) && - !ref.IsNested() && - ref.IsGround() && - !i.isVirtual(ref) -} - // Sorted returns a sorted list of references that the indices were built from. // References that appear more frequently in the indexed rules are ordered // before less frequently appearing references. @@ -439,46 +384,16 @@ func (i *refindices) Mapper(rule *Rule, ref Ref) *valueMapper { return nil } -func (i *refindices) updateEq(rule *Rule, a, b Value, constants map[Var]Value) { +func (i *refindices) updateEq(rule *Rule, a, b Value) { args := rule.Head.Args - if !i.eqOperandsToRefAndValue(rule, args, a, b, constants) { - i.eqOperandsToRefAndValue(rule, args, b, a, constants) + if idx, ok := eqOperandsToRefAndValue(i.isVirtual, args, a, b); ok { + i.insert(rule, idx) + return } -} - -func (i *refindices) updateEqWildcardRef(rule *Rule, a, b Value, constants map[Var]Value) bool { - return i.tryIndexWildcardRef(rule, a, b, constants) || - i.tryIndexWildcardRef(rule, b, a, constants) -} - -func (i *refindices) tryIndexWildcardRef(rule *Rule, a, b Value, constants map[Var]Value) bool { - ref, ok := a.(Ref) - if !ok { - return false + if idx, ok := eqOperandsToRefAndValue(i.isVirtual, args, b, a); ok { + i.insert(rule, idx) + return } - - groundPrefix := ref.GroundPrefix() - if len(groundPrefix) != len(ref)-1 || !i.isValidIndexRef(groundPrefix) { - return false - } - - resolvedValue := b - if bvar, ok := b.(Var); ok { - if resolved, ok := constants[bvar]; ok { - resolvedValue = resolved - } - } else if val, ok := indexValue(b); ok { - resolvedValue = val - } else { - return false - } - - if !IsScalar(resolvedValue) { - return false - } - - i.insert(rule, &refindex{Ref: groundPrefix, Value: resolvedValue}) - return true } func (i *refindices) updateGlobMatch(rule *Rule, expr *Expr) { @@ -494,8 +409,21 @@ func (i *refindices) updateGlobMatch(rule *Rule, expr *Expr) { // 3rd operand was a reference that has been rewritten and bound to a // variable earlier in the query OR a function argument variable. match := expr.Operand(2) - if v, ok := match.Value.(Var); ok { - if ref := resolveVarToRef(i.rules[rule], args, v); ref != nil { + if _, ok := match.Value.(Var); ok { + var ref Ref + for _, other := range i.rules[rule] { + if ov, ok := other.Value.(Var); ok && ov.Equal(match.Value) { + ref = other.Ref + } + } + if ref == nil { + for j, arg := range args { + if arg.Equal(match) { + ref = Ref{FunctionArgRootDocument, InternedTerm(j)} + } + } + } + if ref != nil { i.insert(rule, &refindex{ Ref: ref, Value: arr.Value, @@ -514,137 +442,14 @@ func (i *refindices) updateGlobMatch(rule *Rule, expr *Expr) { } } -func (i *refindices) updateMember(rule *Rule, expr *Expr, constants map[Var]Value) { - args := rule.Head.Args - lhs, rhs := expr.Operand(0), expr.Operand(1) - - lvar, ok := lhs.Value.(Var) - if ok { - lref := resolveVarToRef(i.rules[rule], args, lvar) - if lref != nil { - i.updateMemberRefInValue(rule, lref, rhs, constants) // `ref in value` - return - } - } - - // `var0 in var1` case (var0 may be constant, var1 ref) - i.updateMemberValueInRef(rule, args, lhs.Value, rhs, constants) -} - -func (i *refindices) updateMemberValueInRef(rule *Rule, args []*Term, lval Value, rhs *Term, constants map[Var]Value) { - if lvar, ok := lval.(Var); ok { - val, ok := constants[lvar] - if ok { - lval = val - } - } else if !IsScalar(lval) { - return - } - - rref := i.resolveAndValidateRef(rule, args, rhs) - if rref == nil { - return - } - - i.insert(rule, &refindex{Ref: rref, Value: lval}) -} - -func (i *refindices) updateMemberRefInValue(rule *Rule, ref Ref, rhs *Term, constants map[Var]Value) { - rval := rhs.Value - if rvar, ok := rval.(Var); ok { // rhs is var, try to resolve - if resolved, ok := constants[rvar]; ok { - rval = resolved - } - } - - addRef := func(t *Term) error { - i.insert(rule, &refindex{Ref: ref, Value: t.Value}) - return nil - } - - switch rcol := rval.(type) { - case *Array: - _ = rcol.Iter(addRef) - case Set: - _ = rcol.Iter(addRef) - case Object: - _ = rcol.Iter(func(_, v *Term) error { - return addRef(v) - }) - } -} - -func (i *refindices) resolveAndValidateRef(rule *Rule, args []*Term, term *Term) Ref { - var ref Ref - switch v := term.Value.(type) { - case Ref: - ref = v - case Var: - ref = resolveVarToRef(i.rules[rule], args, v) - default: - return nil - } - - if ref == nil || !i.isValidIndexRef(ref) { - return nil - } - - return ref -} - -// resolveVarToRef checks the previously prepared `*refindex` slice for -// occurrences of the var `v`. Since we store `ref = var` expressions for -// "any" lookups (i.e. "return the rule if ref is anything"), we can -// resolve vars to refs in these simple cases: -// -// __local2__ = input.foo -// __local2__ = -// -// This what builtin calls involving refs are rewritten to, so it is used -// for var -> ref lookup when buiding the RI for glob.match or `v in col`. -// -// For convenience, we also resolve function arg vars here. -// -// NB: This also covers explicit var assignments, like `role := input.rule`, -// but it is no help with chains of assignments, like -// -// x := input.role -// y := x -// -// -// as we're not capturing `var = var` expressions in the index. -func resolveVarToRef(ri []*refindex, args []*Term, v Var) Ref { - for _, other := range ri { - if ov, ok := other.Value.(Var); ok && ov.Equal(v) { - return other.Ref - } - } - for j, arg := range args { - if arg.Value.Compare(v) == 0 { - return Ref{FunctionArgRootDocument, InternedTerm(j)} - } - } - - return nil -} - func (i *refindices) insert(rule *Rule, index *refindex) { count, _ := i.frequency.Get(index.Ref) i.frequency.Put(index.Ref, count+1) - _, indexValueIsVar := index.Value.(Var) - for pos, other := range i.rules[rule] { if other.Ref.Equal(index.Ref) { - - if ValueEqual(other.Value, index.Value) { - return - } - _, otherValueIsVar := other.Value.(Var) - if !indexValueIsVar && otherValueIsVar { - i.rules[rule][pos] = index - return - } + i.rules[rule][pos] = index + return } } @@ -690,12 +495,7 @@ func (tr *trieTraversalResult) Add(t *trieNode) { if !ok { tr.ordering = append(tr.ordering, root) } - // Deduplicate: check if a ruleNode with this priority already exists - if !slices.ContainsFunc(nodes, func(existing *ruleNode) bool { - return existing.prio == node.prio - }) { - tr.unordered[root] = append(nodes, node) - } + tr.unordered[root] = append(nodes, node) } if t.multiple { tr.multiple = true @@ -723,6 +523,45 @@ type trieNode struct { multiple bool } +func (node *trieNode) String() string { + var flags []string + flags = append(flags, fmt.Sprintf("self:%p", node)) + if len(node.ref) > 0 { + flags = append(flags, node.ref.String()) + } + if node.next != nil { + flags = append(flags, fmt.Sprintf("next:%p", node.next)) + } + if node.any != nil { + flags = append(flags, fmt.Sprintf("any:%p", node.any)) + } + if node.undefined != nil { + flags = append(flags, fmt.Sprintf("undefined:%p", node.undefined)) + } + if node.array != nil { + flags = append(flags, fmt.Sprintf("array:%p", node.array)) + } + if node.scalars.Len() > 0 { + buf := make([]string, 0, node.scalars.Len()) + node.scalars.Iter(func(key Value, val *trieNode) bool { + buf = append(buf, fmt.Sprintf("scalar(%v):%p", key, val)) + return false + }) + sort.Strings(buf) + flags = append(flags, strings.Join(buf, " ")) + } + if len(node.rules) > 0 { + flags = append(flags, fmt.Sprintf("%d rule(s)", len(node.rules))) + } + if len(node.mappers) > 0 { + flags = append(flags, fmt.Sprintf("%d mapper(s)", len(node.mappers))) + } + if node.value != nil { + flags = append(flags, "value exists") + } + return strings.Join(flags, " ") +} + func (node *trieNode) append(prio [2]int, rule *Rule) { node.rules = append(node.rules, &ruleNode{prio, rule}) @@ -902,19 +741,8 @@ func (node *trieNode) traverse(resolver ValueResolver, tr *trieTraversalResult) func (node *trieNode) traverseValue(resolver ValueResolver, tr *trieTraversalResult, value Value) error { switch value := value.(type) { - case *Array, Set, Object: - if node.array != nil { - if arr, ok := value.(*Array); ok { - return node.array.traverseArray(resolver, tr, arr) - } - return nil - } - - if node.scalars.Len() > 0 { - return node.traverseCollectionMembership(resolver, tr, value) - } - - return nil + case *Array: + return node.array.traverseArray(resolver, tr, value) case Null, Boolean, Number, String: child, ok := node.scalars.Get(value) @@ -927,29 +755,6 @@ func (node *trieNode) traverseValue(resolver ValueResolver, tr *trieTraversalRes return nil } -func (node *trieNode) traverseCollectionMembership(resolver ValueResolver, tr *trieTraversalResult, collection Value) error { - checkMember := func(t *Term) error { - if IsScalar(t.Value) { - child, _ := node.scalars.Get(t.Value) - return child.Traverse(resolver, tr) - } - return nil - } - - switch col := collection.(type) { - case *Array: - return col.Iter(checkMember) - case Set: - return col.Iter(checkMember) - case Object: - return col.Iter(func(_, v *Term) error { - return checkMember(v) - }) - } - - return nil -} - func (node *trieNode) traverseArray(resolver ValueResolver, tr *trieTraversalResult, arr *Array) error { if node == nil { return nil @@ -1012,42 +817,31 @@ func (node *trieNode) traverseUnknown(resolver ValueResolver, tr *trieTraversalR // for the argument number. So for `f(x, y) { x = 10; y = 12 }`, we'll // bind `args[0]` and `args[1]` to this rule when called for (x=10) and // (y=12) respectively. -func (i *refindices) eqOperandsToRefAndValue(rule *Rule, args []*Term, a, b Value, constants map[Var]Value) bool { +func eqOperandsToRefAndValue(isVirtual func(Ref) bool, args []*Term, a, b Value) (*refindex, bool) { switch v := a.(type) { case Var: - // a is a var, but we have not been able to resolve it to a ref, save for later - if IsConstant(b) { - constants[v] = b - } - - bval, ok := indexValue(b) - if !ok { - return false - } - if ref := resolveVarToRef(i.rules[rule], args, v); ref != nil { - i.insert(rule, &refindex{Ref: ref, Value: bval}) - return true - } - - case Ref: - if !i.isValidIndexRef(v) { - return false - } - - if bvar, ok := b.(Var); ok { // cheaper lookup first: constants - if resolved, ok := constants[bvar]; ok { - b = resolved + for i, arg := range args { + if arg.Value.Compare(a) == 0 { + if bval, ok := indexValue(b); ok { + return &refindex{Ref: Ref{FunctionArgRootDocument, InternedTerm(i)}, Value: bval}, true + } } - } else if bval, ok := indexValue(b); ok { - b = bval - } else { - return false } - - i.insert(rule, &refindex{Ref: v, Value: b}) - return true + case Ref: + if !RootDocumentNames.Contains(v[0]) { + return nil, false + } + if isVirtual(v) { + return nil, false + } + if v.IsNested() || !v.IsGround() { + return nil, false + } + if bval, ok := indexValue(b); ok { + return &refindex{Ref: v, Value: bval}, true + } } - return false + return nil, false } func indexValue(b Value) (Value, bool) { diff --git a/v1/ast/index_debug.go b/v1/ast/index_debug.go deleted file mode 100644 index 88d451b175..0000000000 --- a/v1/ast/index_debug.go +++ /dev/null @@ -1,219 +0,0 @@ -// Copyright 2026 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 ast - -import ( - "fmt" - "sort" - "strings" -) - -func (node *trieNode) mermaid() string { - var sb strings.Builder - sb.WriteString("graph TD\n") - nodeCounter := 0 - nodeIDs := make(map[*trieNode]string) - node.mermaidFormat(&sb, &nodeCounter, nodeIDs, "") - return sb.String() -} - -func (node *trieNode) mermaidFormat(sb *strings.Builder, counter *int, nodeIDs map[*trieNode]string, parentID string) { - currentID, exists := nodeIDs[node] - if !exists { - currentID = fmt.Sprintf("n%d", *counter) - *counter++ - nodeIDs[node] = currentID - label := node.mermaidLabel() - fmt.Fprintf(sb, " %s[\"%s\"]\n", currentID, label) - } - - if parentID != "" { - fmt.Fprintf(sb, " %s --> %s\n", parentID, currentID) - } - - if exists { - return - } - - if node.undefined != nil { - if childID, childExists := nodeIDs[node.undefined]; childExists { - fmt.Fprintf(sb, " %s -->|undefined| %s\n", currentID, childID) - } else { - node.undefined.mermaidFormat(sb, counter, nodeIDs, "") - fmt.Fprintf(sb, " %s -->|undefined| %s\n", currentID, nodeIDs[node.undefined]) - } - } - - if node.any != nil { - if childID, childExists := nodeIDs[node.any]; childExists { - fmt.Fprintf(sb, " %s -->|any| %s\n", currentID, childID) - } else { - node.any.mermaidFormat(sb, counter, nodeIDs, "") - fmt.Fprintf(sb, " %s -->|any| %s\n", currentID, nodeIDs[node.any]) - } - } - - if node.scalars.Len() > 0 { - type scalarPair struct { - key Value - node *trieNode - } - pairs := make([]scalarPair, 0, node.scalars.Len()) - node.scalars.Iter(func(key Value, val *trieNode) bool { - pairs = append(pairs, scalarPair{key, val}) - return false - }) - sort.Slice(pairs, func(a, b int) bool { - return pairs[a].key.Compare(pairs[b].key) < 0 - }) - for _, pair := range pairs { - var scalarLabel string - if s, ok := pair.key.(String); ok { - scalarLabel = string(s) - } else { - scalarLabel = pair.key.String() - } - if len(scalarLabel) > 20 { - scalarLabel = scalarLabel[:20] + "..." - } - scalarLabel = mermaidEscape(scalarLabel) - if childID, childExists := nodeIDs[pair.node]; childExists { - fmt.Fprintf(sb, " %s -->|\"%s\"| %s\n", currentID, scalarLabel, childID) - } else { - pair.node.mermaidFormat(sb, counter, nodeIDs, "") - fmt.Fprintf(sb, " %s -->|\"%s\"| %s\n", currentID, scalarLabel, nodeIDs[pair.node]) - } - } - } - - if node.array != nil { - if childID, childExists := nodeIDs[node.array]; childExists { - fmt.Fprintf(sb, " %s -->|array| %s\n", currentID, childID) - } else { - node.array.mermaidFormat(sb, counter, nodeIDs, "") - fmt.Fprintf(sb, " %s -->|array| %s\n", currentID, nodeIDs[node.array]) - } - } - - if node.next != nil { - node.next.mermaidFormat(sb, counter, nodeIDs, currentID) - } -} - -func (node *trieNode) mermaidLabel() string { - var parts []string - - if len(node.ref) > 0 { - parts = append(parts, node.ref.String()) - } - - if len(node.rules) > 0 { - for _, rn := range node.rules { - bodyStr := "" - if rn.rule.Body != nil { - bodyStr = rn.rule.Body.String() - if len(bodyStr) > 50 { - bodyStr = bodyStr[:50] + "..." - } - } - bodyStr = mermaidEscape(bodyStr) - parts = append(parts, bodyStr) - } - } - - if len(node.mappers) > 0 { - parts = append(parts, fmt.Sprintf("%d mapper(s)", len(node.mappers))) - } - if node.multiple { - parts = append(parts, "multiple") - } - - if len(parts) == 0 { - return "ยท" - } - - return strings.Join(parts, "
") -} - -func mermaidEscape(s string) string { - s = strings.ReplaceAll(s, `"`, `"`) - return s -} - -func (node *trieNode) String() string { - var sb strings.Builder - node.format(&sb, 0) - return sb.String() -} - -func (node *trieNode) format(sb *strings.Builder, depth int) { - indent := strings.Repeat(" ", depth) - - if len(node.ref) > 0 { - sb.WriteString(indent) - sb.WriteString(node.ref.String()) - } else if depth == 0 { - sb.WriteString("root") - } - - if len(node.rules) > 0 { - fmt.Fprintf(sb, " [%d rule(s)]", len(node.rules)) - } - if len(node.mappers) > 0 { - fmt.Fprintf(sb, " [%d mapper(s)]", len(node.mappers)) - } - if node.value != nil { - fmt.Fprintf(sb, " value=%v", node.value) - } - if node.multiple { - sb.WriteString(" [multiple]") - } - sb.WriteString("\n") - - if node.undefined != nil { - sb.WriteString(indent) - sb.WriteString(" undefined:\n") - node.undefined.format(sb, depth+2) - } - - if node.any != nil { - sb.WriteString(indent) - sb.WriteString(" any:\n") - node.any.format(sb, depth+2) - } - - if node.scalars.Len() > 0 { - scalars := make([]Value, 0, node.scalars.Len()) - nodes := make([]*trieNode, 0, node.scalars.Len()) - node.scalars.Iter(func(key Value, val *trieNode) bool { - scalars = append(scalars, key) - nodes = append(nodes, val) - return false - }) - sort.Slice(scalars, func(a, b int) bool { - return scalars[a].Compare(scalars[b]) < 0 - }) - for i := range scalars { - sb.WriteString(indent) - fmt.Fprintf(sb, " %v:\n", scalars[i]) - for j := range nodes { - if ValueEqual(scalars[i], scalars[j]) { - nodes[j].format(sb, depth+2) - break - } - } - } - } - - if node.array != nil { - sb.WriteString(indent) - sb.WriteString(" array:\n") - node.array.format(sb, depth+2) - } - - if node.next != nil { - node.next.format(sb, depth) - } -} diff --git a/v1/ast/index_test.go b/v1/ast/index_test.go index 4aadf3b3fc..7f2eca8744 100644 --- a/v1/ast/index_test.go +++ b/v1/ast/index_test.go @@ -164,7 +164,7 @@ func TestBaseDocEqIndexing(t *testing.T) { x = [1,2,3] x[0] = 1 } { - input.x[_] = 1 # can be indexed since 1.14.0 + input.x[_] = 1 } { input.x[input.y] = 1 } { @@ -431,13 +431,10 @@ func TestBaseDocEqIndexing(t *testing.T) { expectedRS: mod.RuleSet(Var("filtering")), }, { - note: "non-indexable rules", - ruleset: "filtering", - input: `{}`, - expectedRS: mod.RuleSet(Var("filtering")).Diff(NewRuleSet( - MustParseRuleWithOpts(`filtering if { input.x = 1 }`, opts), - MustParseRuleWithOpts(`filtering if { input.x[_] = 1 }`, opts), - )), + note: "non-indexable rules", + ruleset: "filtering", + input: `{}`, + expectedRS: mod.RuleSet(Var("filtering")).Diff(NewRuleSet(MustParseRuleWithOpts(`filtering if { input.x = 1 }`, opts))), }, { note: "unknown: all", @@ -742,484 +739,6 @@ func TestBaseDocEqIndexing(t *testing.T) { // input: `{"k": 1, "v": 2}`, // expectedRS: RuleSet([]*Rule{refMod.Rules[3]}), // }, - { - note: "var assignments: var = ref; var = value", - module: module(`package test - p if { - x = input.foo - x = "bar" - }`), - ruleset: "p", - input: `{"foo": "bar"}`, - expectedRS: []string{ - `p if { x = input.foo; x = "bar" }`, - }, - }, - { - note: "var assignments: var = ref; var = value (no match)", - module: module(`package test - p if { - z = input.foo - z = "bar" - }`), - ruleset: "p", - input: `{"foo": "baz"}`, - expectedRS: []string{}, - }, - { - note: "var assignments: var = value; var = ref (reverse order)", - module: module(`package test - p if { - y = "bar" - y = input.foo - }`), - ruleset: "p", - input: `{"foo": "bar"}`, - expectedRS: []string{ - `p if { y = "bar"; y = input.foo }`, - }, - }, - { - note: "var assignments: var = value; var = ref (reverse order, no match)", - module: module(`package test - p if { - y = "bar" - y = input.foo - }`), - ruleset: "p", - input: `{"foo": "baz"}`, - expectedRS: []string{}, - }, - { - note: "var assignments: value = var; ref = var", - module: module(`package test - p if { - "bar" = x - input.foo = x - }`), - ruleset: "p", - input: `{"foo": "bar"}`, - expectedRS: []string{ - `p if { "bar" = x; input.foo = x }`, - }, - }, - { - note: "var assignments: value = var; ref = var (no match)", - module: module(`package test - p if { - "bar" = x - input.foo = x - }`), - ruleset: "p", - input: `{"foo": "baz"}`, - expectedRS: []string{}, - }, - { - note: "internal.member_2: rhs = value (no match)", - module: module(`package test - p if { - __local0__ = input.role - internal.member_2(__local0__, {"admin", "foo"}) - }`), // this is `p if input.role in {"admin", "foo"}` sent through the compiler - ruleset: "p", - input: `{"role": "bar"}`, - expectedRS: []string{}, - }, - { - note: "internal.member_2: rhs = value", - module: module(`package test - p if { - __local0__ = input.role - internal.member_2(__local0__, {"admin", "foo"}) - }`), - ruleset: "p", - input: `{"role": "foo"}`, - expectedRS: []string{ - `p if { __local0__ = input.role; internal.member_2(__local0__, {"admin", "foo"}) }`}, - }, - { - note: "internal.member_2: rhs = value (other match)", - module: module(`package test - p if { - __local0__ = input.role - internal.member_2(__local0__, {"admin", "foo"}) - }`), - ruleset: "p", - input: `{"role": "admin"}`, - expectedRS: []string{ - `p if { __local0__ = input.role; internal.member_2(__local0__, {"admin", "foo"}) }`}, - }, - { - note: "internal.member_2: lhs = value (2 out of 3)", - module: module(`package test - p if { - x = "a" - x in input.foo - } - p if { - x = "b" - x in input.foo - } - p if { - x = "c" - x in input.foo - } - `), - ruleset: "p", - input: `{"foo": {"a", "b"}}`, - expectedRS: []string{ - `p if { x = "a"; x in input.foo }`, - `p if { x = "b"; x in input.foo }`, - }, - }, - { - note: "internal.member_2: rhs = value (2 out of 3)", - module: module(`package test - p if { - __local0__ = input.role - internal.member_2(__local0__, {"a", "b", "c"}) - } - p if { - __local0__ = input.role - internal.member_2(__local0__, {"x", "b", "z"}) - } - p if { - __local0__ = input.role - internal.member_2(__local0__, {"x", "y", "z"}) - }`), - ruleset: "p", - input: `{"role": "b"}`, - expectedRS: []string{ - `p if { __local0__ = input.role; internal.member_2(__local0__, {"a", "b", "c"}) }`, - `p if { __local0__ = input.role; internal.member_2(__local0__, {"x", "b", "z"}) }`, - }, - }, - { - note: "internal.member_2: unknown value triggers traverseUnknown (duplicate prevention)", - module: module(`package test - p if { - __local0__ = input.role - internal.member_2(__local0__, {"admin", "foo"}) - }`), - ruleset: "p", - unknowns: []string{`input.role`}, - expectedRS: []string{ - `p if { __local0__ = input.role; internal.member_2(__local0__, {"admin", "foo"}) }`}, - }, - { - note: "internal.member_2: object values in rhs (match)", - module: module(`package test - p if { - __local0__ = input.role - internal.member_2(__local0__, {"k1": "admin", "k2": "user"}) - }`), - ruleset: "p", - input: `{"role": "admin"}`, - expectedRS: []string{ - `p if { __local0__ = input.role; internal.member_2(__local0__, {"k1": "admin", "k2": "user"}) }`}, - }, - { - note: "internal.member_2: object values in rhs (no match)", - module: module(`package test - p if { - __local0__ = input.role - internal.member_2(__local0__, {"k1": "admin", "k2": "user"}) - }`), - ruleset: "p", - input: `{"role": "guest"}`, - expectedRS: []string{}, - }, - { - note: "internal.member_2: var with array value in rhs (match)", - module: module(`package test - p if { - __local0__ = input.role - __local1__ = ["admin", "user"] - internal.member_2(__local0__, __local1__) - }`), - ruleset: "p", - input: `{"role": "admin"}`, - expectedRS: []string{ - `p if { __local0__ = input.role; __local1__ = ["admin", "user"]; internal.member_2(__local0__, __local1__) }`}, - }, - { - note: "internal.member_2: var with array value in rhs (no match)", - module: module(`package test - p if { - __local0__ = input.role - __local1__ = ["admin", "user"] - internal.member_2(__local0__, __local1__) - }`), - ruleset: "p", - input: `{"role": "guest"}`, - expectedRS: []string{}, - }, - { - note: "internal.member_2: var with empty array in rhs (one other match)", - module: module(`package test - p if { - __local0__ = input.role - __local1__ = [] - internal.member_2(__local0__, __local1__) - } - p if { - __local0__ = input.role - __local1__ = ["admin", "guest"] - internal.member_2(__local0__, __local1__) - }`), - ruleset: "p", - input: `{"role": "guest"}`, - expectedRS: []string{ - `p if { __local0__ = input.role; __local1__ = []; internal.member_2(__local0__, __local1__) }`, - `p if { __local0__ = input.role; __local1__ = ["admin", "guest"]; internal.member_2(__local0__, __local1__) }`, - }, - }, - { - note: "internal.member_2: var with empty array in rhs (no index)", - module: module(`package test - p if { - __local0__ = input.role - __local1__ = [] - internal.member_2(__local0__, __local1__) - }`), - ruleset: "p", - input: `{"role": "guest"}`, - expectedRS: []string{ - `p if { __local0__ = input.role; __local1__ = []; internal.member_2(__local0__, __local1__) }`, - }, - }, - { - note: "internal.member_2: var with empty array in rhs (one other is mismatch)", - module: module(`package test - p if { - __local0__ = input.role - __local1__ = [] - internal.member_2(__local0__, __local1__) - } - p if { - __local0__ = input.role - __local1__ = ["admin", "guest"] - internal.member_2(__local0__, __local1__) - }`), - ruleset: "p", - input: `{"role": "user"}`, - expectedRS: []string{ - `p if { __local0__ = input.role; __local1__ = []; internal.member_2(__local0__, __local1__) }`, - }, - }, - { - note: "internal.member_2: var with set value in rhs (match)", - module: module(`package test - p if { - __local0__ = input.role - __local1__ = {"admin", "user"} - internal.member_2(__local0__, __local1__) - }`), - ruleset: "p", - input: `{"role": "admin"}`, - expectedRS: []string{ - `p if { __local0__ = input.role; __local1__ = {"admin", "user"}; internal.member_2(__local0__, __local1__) }`}, - }, - { - note: "internal.member_2: var with set value in rhs (no match)", - module: module(`package test - p if { - __local0__ = input.role - __local1__ = {"admin", "user"} - internal.member_2(__local0__, __local1__) - }`), - ruleset: "p", - input: `{"role": "guest"}`, - expectedRS: []string{}, - }, - { - note: "functions: member_2 in function, arg not matching", - module: module(`package test - member2_f(a) if a in {"a", "b"}`), - ruleset: "member2_f", - args: []Value{String("c")}, - expectedRS: []string{}, - }, - { - note: "functions: member_2 in function, arg matching", - module: module(`package test - member2_f(a) if a in {"a", "b"}`), - ruleset: "member2_f", - args: []Value{String("a")}, - expectedRS: []string{ - `member2_f(a) = true if { a in {"a", "b"} }`, - }, - }, - { - note: "internal.member_2: reverse - scalar in collection (match)", - module: module(`package test - p if { - __local0__ = input.roles - internal.member_2("admin", __local0__) - }`), - ruleset: "p", - input: `{"roles": ["admin", "user"]}`, - expectedRS: []string{ - `p if { __local0__ = input.roles; internal.member_2("admin", __local0__) }`}, - }, - { - note: "internal.member_2: reverse - scalar in collection (no match)", - module: module(`package test - p if { - __local0__ = input.roles - internal.member_2("guest", __local0__) - }`), - ruleset: "p", - input: `{"roles": ["admin", "user"]}`, - expectedRS: []string{}, - }, - { - note: "internal.member_2: reverse - var in collection (match)", - module: module(`package test - p if { - __local0__ = input.roles - __local1__ = "admin" - internal.member_2(__local1__, __local0__) - }`), - ruleset: "p", - input: `{"roles": ["admin", "user"]}`, - expectedRS: []string{ - `p if { __local0__ = input.roles; __local1__ = "admin"; internal.member_2(__local1__, __local0__) }`}, - }, - { - note: "internal.member_2: reverse - var in collection (no match)", - module: module(`package test - p if { - __local0__ = input.roles - __local1__ = "guest" - internal.member_2(__local1__, __local0__) - }`), - ruleset: "p", - input: `{"roles": ["admin", "user"]}`, - expectedRS: []string{}, - }, - { - note: "internal.member_2: reverse - scalar in set (match)", - module: module(`package test - p if { - __local0__ = input.items - internal.member_2("b", __local0__) - }`), - ruleset: "p", - input: `{"items": {"a", "b", "c"}}`, - expectedRS: []string{ - `p if { __local0__ = input.items; internal.member_2("b", __local0__) }`}, - }, - { - note: "internal.member_2: reverse - scalar in set (no match)", - module: module(`package test - p if { - __local0__ = input.items - internal.member_2("z", __local0__) - }`), - ruleset: "p", - input: `{"items": {"a", "b", "c"}}`, - expectedRS: []string{}, - }, - { - note: "internal.member_2: reverse - number in array (match)", - module: module(`package test - p if { - __local0__ = input.nums - internal.member_2(2, __local0__) - }`), - ruleset: "p", - input: `{"nums": [1, 2, 3]}`, - expectedRS: []string{ - `p if { __local0__ = input.nums; internal.member_2(2, __local0__) }`}, - }, - { - note: "internal.member_2: reverse - number in array (no match)", - module: module(`package test - p if { - __local0__ = input.nums - internal.member_2(99, __local0__) - }`), - ruleset: "p", - input: `{"nums": [1, 2, 3]}`, - expectedRS: []string{}, - }, - { - note: "internal.member_2: reverse - scalar in object values (match)", - module: module(`package test - p if { - __local0__ = input.obj - internal.member_2("bar", __local0__) - }`), - ruleset: "p", - input: `{"obj": {"a": "foo", "b": "bar"}}`, - expectedRS: []string{ - `p if { __local0__ = input.obj; internal.member_2("bar", __local0__) }`}, - }, - { - note: "internal.member_2: reverse - scalar in object values (no match)", - module: module(`package test - p if { - __local0__ = input.obj - internal.member_2("baz", __local0__) - }`), - ruleset: "p", - input: `{"obj": {"a": "foo", "b": "bar"}}`, - expectedRS: []string{}, - }, - { - note: "internal.member_2: reverse - non-scalar not indexed", - module: module(`package test - p if { - __local0__ = input.items - internal.member_2({"foo": "bar"}, __local0__) - }`), - ruleset: "p", - input: `{"items": [{"foo": "buz"}]}`, // not a match! - expectedRS: []string{ - `p if { __local0__ = input.items; internal.member_2({"foo": "bar"}, __local0__) }`}, - }, - { - note: "wildcard ref equality: input.roles[_] == \"admin\" indexed like \"admin\" in input.roles (match)", - module: module(`package test - p if { - input.roles[_] == "admin" - } - p if { - "admin" in input.roles - }`), - ruleset: "p", - input: `{"roles": ["admin", "user"]}`, - expectedRS: []string{ - `p if { equal(input.roles[_], "admin") }`, - `p if { internal.member_2("admin", input.roles) }`, - }, - }, - { - note: "wildcard ref equality: input.roles[_] == \"admin\" indexed like \"admin\" in input.roles (no match)", - module: module(`package test - p if { - input.roles[_] == "admin" - } - p if { - "admin" in input.roles - }`), - ruleset: "p", - input: `{"roles": ["user", "guest"]}`, - expectedRS: []string{}, - }, - { - note: "wildcard ref equality: reverse order value == ref[_] (match)", - module: module(`package test - p if { - "admin" == input.roles[_] - }`), - ruleset: "p", - input: `{"roles": ["admin", "user"]}`, - expectedRS: []string{ - `p if { equal("admin", input.roles[_]) }`, - }, - }, } for _, tc := range tests { @@ -1270,7 +789,6 @@ func TestBaseDocEqIndexing(t *testing.T) { t.Fatalf("Expected index build to succeed") } - t.Log(index.root.mermaid()) var unknownRefs Set if len(tc.unknowns) > 0 { @@ -1505,39 +1023,6 @@ func TestGetAllRules(t *testing.T) { } } -func TestGetAllRulesInternalMember2(t *testing.T) { - module := module(` - package test - - p if { - x = input.fruit - x in {"apple", "pear"} - } - `) - - index := newBaseDocEqIndex(func(Ref) bool { return false }) - - ok := index.Build(module.Rules) - if !ok { - t.Fatalf("Expected index build to succeed") - } - - result, err := index.AllRules(testResolver{input: MustParseTerm(`{}`)}) - if err != nil { - t.Fatalf("Unexpected error during index lookup: %v", err) - } - - expectedRules := NewRuleSet(module.Rules[0]) - - if !NewRuleSet(result.Rules...).Equal(expectedRules) { - t.Fatalf("Expected rules to be %v but got: %v", expectedRules, result.Rules) - } - - if len(result.Else) > 0 { - t.Fatalf("Expected no else rules but got: %v", result.Else) - } -} - func TestSkipIndexing(t *testing.T) { module := module(`package test