Revert "ast: make rule index track var assignments and x in {...} (#8341)" (#8410)

This reverts commit 32b97ae08c.

Signed-off-by: Sebastian Spaink <sebastianspaink@gmail.com>
This commit is contained in:
Sebastian Spaink
2026-03-06 14:08:34 -06:00
committed by Stephan Renatus
parent 3d57b2ec57
commit c416fd964e
3 changed files with 103 additions and 1043 deletions
+98 -304
View File
@@ -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__ = <something>
//
// 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
// <something with 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) {
-219
View File
@@ -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, "<br/>")
}
func mermaidEscape(s string) string {
s = strings.ReplaceAll(s, `"`, `&quot;`)
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)
}
}
+5 -520
View File
@@ -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