Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 5 additions & 51 deletions v1/ast/compare.go
Original file line number Diff line number Diff line change
Expand Up @@ -107,10 +107,10 @@ func Compare(a, b any) int {
case Var:
return VarCompare(a, b.(Var))
case Ref:
return termSliceCompare(a, b.(Ref))
return slices.CompareFunc(a, b.(Ref), TermValueCompare)
case *Array:
b := b.(*Array)
return termSliceCompare(a.elems, b.elems)
return slices.CompareFunc(a.elems, b.elems, TermValueCompare)
case *lazyObj:
return Compare(a.force(), b)
case *object:
Expand All @@ -130,7 +130,7 @@ func Compare(a, b any) int {
b := b.(*SetComprehension)
return a.Compare(b)
case Call:
return termSliceCompare(a, b.(Call))
return slices.CompareFunc(a, b.(Call), TermValueCompare)
case *Expr:
return a.Compare(b.(*Expr))
case *SomeDecl:
Expand All @@ -150,7 +150,7 @@ func Compare(a, b any) int {
case *Rule:
return a.Compare(b.(*Rule))
case Args:
return termSliceCompare(a, b.(Args))
return slices.CompareFunc(a, b.(Args), TermValueCompare)
case *Import:
return a.Compare(b.(*Import))
case *Package:
Expand Down Expand Up @@ -247,52 +247,6 @@ func sortOrder(x any) int {
panic(fmt.Sprintf("illegal value: %T", x))
}

func rulesCompare(a, b []*Rule) int {
minLen := min(len(b), len(a))
for i := range minLen {
if cmp := a[i].Compare(b[i]); cmp != 0 {
return cmp
}
}
if len(a) < len(b) {
return -1
}
if len(b) < len(a) {
return 1
}
return 0
}

func termSliceCompare(a, b []*Term) int {
minLen := min(len(b), len(a))
for i := range minLen {
if cmp := a[i].Value.Compare(b[i].Value); cmp != 0 {
return cmp
}
}
if len(a) < len(b) {
return -1
} else if len(b) < len(a) {
return 1
}
return 0
}

func withSliceCompare(a, b []*With) int {
minLen := min(len(b), len(a))
for i := range minLen {
if cmp := a[i].Compare(b[i]); cmp != 0 {
return cmp
}
}
if len(a) < len(b) {
return -1
} else if len(b) < len(a) {
return 1
}
return 0
}

func VarCompare(a, b Var) int {
if a == b {
return 0
Expand Down Expand Up @@ -329,7 +283,7 @@ func ValueEqual(a, b Value) bool {
}

func RefCompare(a, b Ref) int {
return termSliceCompare(a, b)
return slices.CompareFunc(a, b, TermValueCompare)
}

func RefEqual(a, b Ref) bool {
Expand Down
99 changes: 33 additions & 66 deletions v1/ast/compile.go
Original file line number Diff line number Diff line change
Expand Up @@ -2118,18 +2118,6 @@ func (c *Compiler) getExports() *util.HasherMap[Ref, []Ref] {
return rules
}

func refSliceEqual(a, b []Ref) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if !a[i].Equal(b[i]) {
return false
}
}
return true
}

func hashMapAdd(rules *util.HasherMap[Ref, []Ref], pkg, rule Ref) {
prev, ok := rules.Get(pkg)
if !ok {
Expand Down Expand Up @@ -2640,15 +2628,12 @@ func rewriteTemplateString(tsr *templateStringRewriter, safe VarSet, loc *Locati
// Note: we don't care about not exprs here
vis = ClearOrNewVarVisitor(vis).WithParams(SafetyCheckVisitorParams)
vis.Walk(t)
vars := vis.Vars()
if vars.DiffCount(safe) > 0 {
unsafe := vars.Diff(safe)
for _, v := range unsafe.Sorted() {
if w, ok := tsr.rewritten[v]; ok {
v = w
}
errs = append(errs, NewError(CompileErr, t.Loc(), "var %v is undeclared", v))
unsafe := vis.Vars().DeleteFunc(safe.Contains)
for _, v := range unsafe.Sorted() {
if w, ok := tsr.rewritten[v]; ok {
v = w
}
errs = append(errs, NewError(CompileErr, t.Loc(), "var %v is undeclared", v))
}

loc := t.Loc()
Expand Down Expand Up @@ -2829,15 +2814,12 @@ func rewritePrintCalls(gen *localVarGenerator, getArity func(Ref) int, globals V
// Note: we don't care about not exprs here
vis = vis.Clear().WithParams(SafetyCheckVisitorParams)
vis.Walk(args[j])
vars := vis.Vars()
if vars.DiffCount(safe) > 0 {
unsafe := vars.Diff(safe)
for _, v := range unsafe.Sorted() {
if w, ok := rewritten[v]; ok {
v = w
}
errs = append(errs, NewError(CompileErr, args[j].Loc(), "var %v is undeclared", v))
unsafe := vis.Vars().DeleteFunc(safe.Contains)
for _, v := range unsafe.Sorted() {
if w, ok := rewritten[v]; ok {
v = w
}
errs = append(errs, NewError(CompileErr, args[j].Loc(), "var %v is undeclared", v))
}
}

Expand Down Expand Up @@ -3356,11 +3338,7 @@ func (c *Compiler) rewriteLocalVarsInRule(rule *Rule, unusedArgs VarSet, argsSta
if rule.Head.Value != nil && !IsScalar(rule.Head.Value.Value) {
valueVars := rule.Head.Value.Vars()
vis.vars.Update(valueVars)
for arg := range unusedArgs {
if valueVars.Contains(arg) {
delete(unusedArgs, arg)
}
}
unusedArgs.DeleteFunc(valueVars.Contains)
}

used = vis.Vars()
Expand Down Expand Up @@ -4063,7 +4041,9 @@ func getComprehensionIndex(dbg debug.Debug, arity func(Ref) int, candidates VarS
}

outputs := outputVarsForBody(body, arity, ReservedVars, nil)
unsafe := body.Vars(SafetyCheckVisitorParams).Diff(outputs).Diff(ReservedVars)
unsafe := body.Vars(SafetyCheckVisitorParams).
DeleteFunc(outputs.Contains).
DeleteFunc(ReservedVars.Contains)

if len(unsafe) > 0 {
dbg.Printf("%s: comprehension index: unsafe vars: %v", expr.Location, unsafe)
Expand Down Expand Up @@ -4100,12 +4080,7 @@ func getComprehensionIndex(dbg debug.Debug, arity func(Ref) int, candidates VarS
return nil
}

result := make([]*Term, 0, len(indexVars))
for v := range indexVars {
result = append(result, NewTerm(v))
}
slices.SortFunc(result, TermValueCompare)

result := util.SortedFunc(util.MapKeys(indexVars, ToTerm), TermValueCompare)
debugRes := make([]*Term, len(result))
for i, r := range result {
if o, ok := rwVars[r.Value.(Var)]; ok {
Expand Down Expand Up @@ -4796,10 +4771,7 @@ func (g *Graph) Sort() (sorted []util.T, ok bool) {
temp: map[util.T]struct{}{},
}

nodesList := make([]util.T, 0, len(g.nodes))
for node := range g.nodes {
nodesList = append(nodesList, node)
}
nodesList := util.Keys(g.nodes)
sortGraphNodes(nodesList)
for _, node := range nodesList {
if !sorter.Visit(node) {
Expand Down Expand Up @@ -5010,10 +4982,8 @@ func reorderBodyForSafety(builtins map[string]*Builtin, arity func(Ref) int, glo
for _, e := range body {
vis = vis.Clear().WithParams(SafetyCheckVisitorParamsWithArity(arity))
vis.Walk(e)
for v := range vis.Vars() {
if _, ok := safe[v]; !ok {
unsafe.Add(e, v)
}
for v := range vis.Vars().DeleteFunc(safe.Contains) {
unsafe.Add(e, v)
}
}

Expand All @@ -5034,13 +5004,13 @@ func reorderBodyForSafety(builtins map[string]*Builtin, arity func(Ref) int, glo
// check closures: is this expression closing over variables that
// haven't been made safe by what's already included in `reordered`?
unsafeVarsInClosures(e, unsVis)
cv := unsVis.Vars().Intersect(bodyVars).Diff(globals)
cv := unsVis.Vars().Intersect(bodyVars).DeleteFunc(globals.Contains)
unsVis.Clear()

ob := outputVarsForBody(reordered, arity, safe, vis)

if cv.DiffCount(ob) > 0 {
uv := cv.Diff(ob)
uv := cv.DeleteFunc(ob.Contains)
if uv.Equal(ovs) { // special case "closure-self"
continue
}
Expand Down Expand Up @@ -5300,8 +5270,7 @@ func (xform *bodySafetyTransformer) reorderComprehensionSafety(tv VarSet, body B
bv.Update(xform.globals)

if tv.DiffCount(bv) > 0 {
uv := tv.Diff(bv)
for v := range uv {
for v := range tv.Diff(bv) {
xform.unsafe.Add(xform.current, v)
}
}
Expand Down Expand Up @@ -5361,11 +5330,10 @@ func outputVarsForBody(body Body, arity func(Ref) int, safe VarSet, vis *VarVisi
output := VarSet{}

vis = ClearOrNewVarVisitor(vis)

for _, e := range body {
o.Update(outputVarsForExpr(e, arity, o, output, vis))
}
return o.Diff(safe)
return o.DeleteFunc(safe.Contains)
}

// OutputVarsFromExpr returns all variables which are the "output" for
Expand Down Expand Up @@ -5473,7 +5441,7 @@ func outputVarsForExprCall(expr *Expr, arity int, safe VarSet, terms []*Term, vi
vis = ClearOrNewVarVisitor(vis).WithParams(params)
vis.WalkArgs(Args(terms[:numInputTerms]))

unsafe := vis.Vars().Diff(output).DiffCount(safe)
unsafe := vis.Vars().DeleteFunc(output.Contains).DiffCount(safe)
if unsafe > 0 {
return VarSet{}
}
Expand Down Expand Up @@ -5574,9 +5542,9 @@ func newLocalVarGeneratorForModuleSet(sorted []string, modules map[string]*Modul
return &localVarGenerator{exclude: vis.vars, suffix: LocalVarPrefix}
}

func newLocalVarGenerator(suffix string, node any) *localVarGenerator {
func newLocalVarGenerator(suffix string, body Body) *localVarGenerator {
vis := NewVarVisitor()
vis.Walk(node)
vis.WalkBody(body)
return &localVarGenerator{exclude: vis.vars, suffix: LocalVarPrefix + suffix}
}

Expand All @@ -5599,7 +5567,7 @@ func getGlobals(pkg *Package, rules []Ref, imports []*Import) map[Var]*usedRef {

for _, ref := range rules {
v := ref[0].Value.(Var)
globals[v] = &usedRef{ref: pkg.Path.Append(StringTerm(string(v)))}
globals[v] = &usedRef{ref: pkg.Path.Append(InternedTerm(string(v)))}
}

for _, imp := range imports {
Expand Down Expand Up @@ -6810,13 +6778,12 @@ func checkUnusedDeclaredVars(body Body, stack *localDeclaredVars, used VarSet, c

for v := range used {
if gv, ok := stack.Declared(v); ok {
bodyvars.Add(gv)
} else {
bodyvars.Add(v)
v = gv
}
bodyvars.Add(v)
}

dbv := declared.Diff(bodyvars)
dbv := declared.DeleteFunc(bodyvars.Contains)
if dbv.DiffCount(used) == 0 {
return errs
}
Expand All @@ -6826,7 +6793,7 @@ func checkUnusedDeclaredVars(body Body, stack *localDeclaredVars, used VarSet, c
reversed[v] = k
}

for _, gv := range dbv.Diff(used).Sorted() {
for _, gv := range dbv.DeleteFunc(used.Contains).Sorted() {
rv := reversed[gv]
if !rv.IsGenerated() {
// Scan through body exprs, looking for a match between the
Expand Down Expand Up @@ -7301,7 +7268,7 @@ func validateWith(c *Compiler, unsafeBuiltinsMap map[string]struct{}, expr *Expr
// Ensure that values that are built-ins are rewritten to Ref (not Var)
if v, ok := value.Value.(Var); ok {
if _, ok := c.builtins[v.String()]; ok {
value.Value = Ref([]*Term{NewTerm(v)})
value.Value = Ref([]*Term{NewTerm(value.Value)})
}
}
isBuiltinRefOrVar, err := isBuiltinRefOrVar(c.builtins, unsafeBuiltinsMap, target)
Expand Down Expand Up @@ -7362,8 +7329,8 @@ func validateWith(c *Compiler, unsafeBuiltinsMap map[string]struct{}, expr *Expr
case isBuiltinRefOrVar:
// NOTE(sr): first we ensure that parsed Var builtins (`count`, `concat`, etc)
// are rewritten to their proper Ref convention
if v, ok := target.Value.(Var); ok {
target.Value = Ref([]*Term{NewTerm(v)})
if _, ok := target.Value.(Var); ok {
target.Value = Ref([]*Term{NewTerm(target.Value)})
}

targetRef := target.Value.(Ref)
Expand Down
7 changes: 2 additions & 5 deletions v1/ast/compile_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -444,10 +444,7 @@ func refMapEqual(a, b *util.HasherMap[Ref, []Ref]) bool {
if !ok {
return true
}
if !refSliceEqual(v, v2) {
return true
}
return false
return !slices.EqualFunc(v, v2, RefEqual)
})
}

Expand Down Expand Up @@ -2982,7 +2979,7 @@ func TestCompilerExprExpansion(t *testing.T) {

for _, tc := range tests {
t.Run(tc.note, func(t *testing.T) {
gen := newLocalVarGenerator("", NullTerm())
gen := newLocalVarGenerator("", Body{})
expr := MustParseExpr(tc.input)
result := expandExpr(gen, expr.Copy())
if len(result) != len(tc.expected) {
Expand Down
Loading
Loading