diff --git a/v1/ast/index.go b/v1/ast/index.go index 663072ed83..acddd0776d 100644 --- a/v1/ast/index.go +++ b/v1/ast/index.go @@ -5,7 +5,6 @@ package ast import ( - "fmt" "slices" "sort" "strings" @@ -68,6 +67,7 @@ var ( globMatchRef = GlobMatch.Ref() internalPrintRef = InternalPrint.Ref() internalTestCaseRef = InternalTestCase.Ref() + internalMemberRef = Member.Ref() skipIndexing = NewSet(NewTerm(internalPrintRef), NewTerm(internalTestCaseRef)) ) @@ -87,6 +87,7 @@ 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 { @@ -106,8 +107,9 @@ func (i *baseDocEqIndex) Build(rules []*Rule) bool { } } if !skip { + clear(values) for i := range rule.Body { - indices.Update(rule, rule.Body[i]) + indices.Update(rule, rule.Body[i], values) } } return false @@ -124,7 +126,46 @@ func (i *baseDocEqIndex) Build(rules []*Rule) bool { node := i.root if indices.Indexed(rule) { for _, ref := range indices.Sorted() { - node = node.Insert(ref, indices.Value(rule, ref), indices.Mapper(rule, ref)) + 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 + } + } } } // Insert rule into trie with (insertion order, priority order) @@ -292,7 +333,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) { +func (i *refindices) Update(rule *Rule, expr *Expr, values map[Var]Value) { if len(expr.With) > 0 { // NOTE(tsandall): In the future, we may need to consider expressions @@ -314,30 +355,44 @@ func (i *refindices) Update(rule *Rule, expr *Expr) { // 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) + i.updateEq(rule, ref, anyValue, nil) } } } - 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: + equalish := op.Equal(equalityRef) || // unification, no 3-operands version exists // 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.) - i.updateEq(rule, a.Value, b.Value) + (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) + } 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. @@ -384,16 +439,46 @@ func (i *refindices) Mapper(rule *Rule, ref Ref) *valueMapper { return nil } -func (i *refindices) updateEq(rule *Rule, a, b Value) { +func (i *refindices) updateEq(rule *Rule, a, b Value, constants map[Var]Value) { args := rule.Head.Args - if idx, ok := eqOperandsToRefAndValue(i.isVirtual, args, a, b); ok { - i.insert(rule, idx) - return + if !i.eqOperandsToRefAndValue(rule, args, a, b, constants) { + i.eqOperandsToRefAndValue(rule, args, b, a, constants) } - if idx, ok := eqOperandsToRefAndValue(i.isVirtual, args, b, a); 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 } + + 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) { @@ -409,21 +494,8 @@ 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 _, 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 { + if v, ok := match.Value.(Var); ok { + if ref := resolveVarToRef(i.rules[rule], args, v); ref != nil { i.insert(rule, &refindex{ Ref: ref, Value: arr.Value, @@ -442,14 +514,137 @@ 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) { - i.rules[rule][pos] = index - return + + if ValueEqual(other.Value, index.Value) { + return + } + _, otherValueIsVar := other.Value.(Var) + if !indexValueIsVar && otherValueIsVar { + i.rules[rule][pos] = index + return + } } } @@ -495,7 +690,12 @@ func (tr *trieTraversalResult) Add(t *trieNode) { if !ok { tr.ordering = append(tr.ordering, root) } - tr.unordered[root] = append(nodes, node) + // 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) + } } if t.multiple { tr.multiple = true @@ -523,45 +723,6 @@ 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}) @@ -741,8 +902,19 @@ 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: - return node.array.traverseArray(resolver, tr, value) + 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 Null, Boolean, Number, String: child, ok := node.scalars.Get(value) @@ -755,6 +927,29 @@ 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 @@ -817,31 +1012,42 @@ 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 eqOperandsToRefAndValue(isVirtual func(Ref) bool, args []*Term, a, b Value) (*refindex, bool) { +func (i *refindices) eqOperandsToRefAndValue(rule *Rule, args []*Term, a, b Value, constants map[Var]Value) bool { switch v := a.(type) { case Var: - 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 - } - } + // 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 !RootDocumentNames.Contains(v[0]) { - return nil, false + if !i.isValidIndexRef(v) { + return 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 + + if bvar, ok := b.(Var); ok { // cheaper lookup first: constants + if resolved, ok := constants[bvar]; ok { + b = resolved + } + } else if bval, ok := indexValue(b); ok { + b = bval + } else { + return false } + + i.insert(rule, &refindex{Ref: v, Value: b}) + return true } - return nil, false + return false } func indexValue(b Value) (Value, bool) { diff --git a/v1/ast/index_debug.go b/v1/ast/index_debug.go new file mode 100644 index 0000000000..88d451b175 --- /dev/null +++ b/v1/ast/index_debug.go @@ -0,0 +1,219 @@ +// 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 7f2eca8744..4aadf3b3fc 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 + input.x[_] = 1 # can be indexed since 1.14.0 } { input.x[input.y] = 1 } { @@ -431,10 +431,13 @@ 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))), + 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: "unknown: all", @@ -739,6 +742,484 @@ 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 { @@ -789,6 +1270,7 @@ 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 { @@ -1023,6 +1505,39 @@ 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