eval: catch booleans trying to escape to the heap (#8968)

We do this in the hottest spots already, but might as well do it
wherever we can. Had to deploy a few tricks in some places, but I think
the code is better off too.

Signed-off-by: Anders Eknert <anders.eknert@apple.com>
This commit is contained in:
Anders Eknert
2026-08-10 17:57:17 +02:00
committed by GitHub
parent 8db1554985
commit e5b3e1cfc1
+55 -75
View File
@@ -192,6 +192,7 @@ func (e *eval) closure(query ast.Body, cpy *eval) {
cpy.queryID = cpy.queryIDFact.Next()
cpy.parent = e
cpy.findOne = false
cpy.defined = false
}
// childWithBindingSizeHint creates a child evaluator with bindings pre-sized for the expected number of variables.
@@ -203,6 +204,7 @@ func (e *eval) childWithBindingSizeHint(query ast.Body, cpy *eval, sizeHint int)
cpy.bindings = newBindingsWithSize(cpy.queryID, e.instr, sizeHint)
cpy.parent = e
cpy.findOne = false
cpy.defined = false
}
func (e *eval) next(iter evalIterator) error {
@@ -418,15 +420,12 @@ func (e *eval) evalExpr(iter evalIterator) error {
})
}
func (e *eval) evalStep(iter evalIterator) error {
func (e *eval) evalStep(iter evalIterator) (err error) {
expr := e.query[e.index]
if expr.Negated {
return e.evalNot(iter)
}
var err error
// NOTE(æ): the reason why there's one branch for the tracing case and one almost
// identical branch below for when tracing is disabled is that the tracing case
// allocates wildly. These allocations are cause by the "defined" boolean variable
@@ -629,7 +628,6 @@ func (e *eval) fmtVar() string {
func (e *eval) evalNot(iter evalIterator) error {
expr := e.query[e.index]
if e.unknown(expr, e.bindings) {
return e.setupAndEvalNotPartial(iter)
}
@@ -639,7 +637,6 @@ func (e *eval) evalNot(iter evalIterator) error {
defer evalPool.Put(child)
e.closure(negation, child)
if e.traceEnabled {
child.traceEnter(negation)
}
@@ -660,8 +657,6 @@ func (e *eval) evalNot(iter evalIterator) error {
return iter(e)
}
child.defined = false
e.traceFail(expr)
return nil
}
@@ -3003,7 +2998,6 @@ func (h *evalVirtualPartialCacheHint) keyWithoutScope() ast.Ref {
}
func (e evalVirtualPartial) eval(iter unifyIterator) error {
unknown := e.e.unknown(e.ref[:e.pos+1], e.bindings)
if len(e.ref) == e.pos+1 {
@@ -3038,7 +3032,6 @@ func maxRefLength(rules []*ast.Rule, ceil int) int {
}
func (e evalVirtualPartial) evalEachRule(iter unifyIterator, unknown bool) error {
if e.ir.Empty() {
return nil
}
@@ -3095,7 +3088,6 @@ func (e evalVirtualPartial) evalEachRule(iter unifyIterator, unknown bool) error
}
func (e evalVirtualPartial) evalAllRules(iter unifyIterator, rules []*ast.Rule) error {
cacheKey := e.plugged[:e.pos+1]
result, _ := e.e.virtualCache.Get(cacheKey)
if result != nil {
@@ -3128,17 +3120,16 @@ func (e evalVirtualPartial) evalAllRulesNoCache(rules []*ast.Rule) (*ast.Term, e
for _, rule := range rules {
e.e.childWithBindingSizeHint(rule.Body, child, ast.EstimateBodyBindingCount(rule.Body))
child.traceEnter(rule)
err := child.eval(func(*eval) error {
err := child.eval(func(*eval) (err error) {
child.traceExit(rule)
e.e.evaluated.Record(rule)
var err error
result, _, err = e.reduce(rule, child.bindings, result, &visitedRefs)
if err != nil {
return err
if err == nil && child.traceEnabled {
child.traceRedo(rule)
}
child.traceRedo(rule)
return nil
return err
})
if err != nil {
@@ -3159,14 +3150,20 @@ func wrapInObjects(leaf *ast.Term, ref ast.Ref) *ast.Term {
return ast.ObjectTerm(ast.Item(key, val))
}
func (e evalVirtualPartial) evalOneRulePreUnify(iter unifyIterator, rule *ast.Rule, result *ast.Term, unknown bool, visitedRefs *[]ast.Ref) (*ast.Term, error) {
func (e evalVirtualPartial) evalOneRulePreUnify(
iter unifyIterator,
rule *ast.Rule,
result *ast.Term,
unknown bool,
visitedRefs *[]ast.Ref,
) (*ast.Term, error) {
child := evalPool.Get()
defer evalPool.Put(child)
e.e.childWithBindingSizeHint(rule.Body, child, ast.EstimateBodyBindingCount(rule.Body))
child.traceEnter(rule)
var defined bool
if child.traceEnabled {
child.traceEnter(rule)
}
headKey := rule.Head.Key
if headKey == nil {
@@ -3174,11 +3171,17 @@ func (e evalVirtualPartial) evalOneRulePreUnify(iter unifyIterator, rule *ast.Ru
}
// Walk the dynamic portion of rule ref and key to unify vars
err := child.biunifyRuleHead(e.pos+1, e.ref, rule, e.bindings, child.bindings, func(_ int) error {
defined = true
return child.eval(func(child *eval) error {
child.traceExit(rule)
err := child.biunifyRuleHead(e.pos+1, e.ref, rule, e.bindings, child.bindings, func(int) error {
child.defined = true
return child.eval(func(child *eval) (err error) {
if child.traceEnabled {
child.traceExit(rule)
defer func() {
if err == nil {
child.traceRedo(rule)
}
}()
}
term := rule.Head.Value
if term == nil {
@@ -3187,7 +3190,6 @@ func (e evalVirtualPartial) evalOneRulePreUnify(iter unifyIterator, rule *ast.Ru
if unknown {
term, termbindings := child.bindings.apply(term)
if rule.Head.RuleKind() == ast.MultiValue {
term = ast.SetTerm(term)
}
@@ -3195,37 +3197,24 @@ func (e evalVirtualPartial) evalOneRulePreUnify(iter unifyIterator, rule *ast.Ru
objRef := rule.Ref()[e.pos+1:]
term = wrapInObjects(term, objRef)
err := e.evalTerm(iter, e.pos+1, term, termbindings)
if err != nil {
return err
}
err = e.evalTerm(iter, e.pos+1, term, termbindings)
} else {
var dup bool
var err error
result, dup, err = e.reduce(rule, child.bindings, result, visitedRefs)
if err != nil {
return err
} else if !unknown && dup {
if err == nil && !unknown && dup && child.traceEnabled {
child.traceDuplicate(rule)
return nil
}
}
child.traceRedo(rule)
return nil
return err
})
})
if err != nil {
return nil, err
}
if !defined {
if err == nil && child.traceEnabled && !child.defined {
child.traceFail(rule)
}
return result, nil
return result, err
}
func (e *eval) biunifyRuleHead(pos int, ref ast.Ref, rule *ast.Rule, refBindings, ruleBindings *bindings, iter unifyRefIterator) error {
@@ -3256,34 +3245,30 @@ func (e *eval) biunifyDynamicRef(pos int, a, b ast.Ref, b1, b2 *bindings, iter u
func (e evalVirtualPartial) evalOneRulePostUnify(iter unifyIterator, rule *ast.Rule) error {
child := evalPool.Get()
defer evalPool.Put(child)
defer func() {
if child.traceEnabled && !child.defined {
child.traceFail(rule)
}
evalPool.Put(child)
}()
e.e.childWithBindingSizeHint(rule.Body, child, ast.EstimateBodyBindingCount(rule.Body))
if e.e.traceEnabled {
child.traceEnter(rule)
}
child.traceEnter(rule)
var defined bool
err := child.eval(func(child *eval) error {
defined = true
return e.e.biunifyRuleHead(e.pos+1, e.ref, rule, e.bindings, child.bindings, func(_ int) error {
return e.evalOneRuleContinue(iter, rule, child)
return child.eval(func(next *eval) error {
child.defined = true
return e.e.biunifyRuleHead(e.pos+1, e.ref, rule, e.bindings, next.bindings, func(int) error {
return e.evalOneRuleContinue(iter, rule, next)
})
})
if err != nil {
return err
}
if !defined {
child.traceFail(rule)
}
return nil
}
func (e evalVirtualPartial) evalOneRuleContinue(iter unifyIterator, rule *ast.Rule, child *eval) error {
child.traceExit(rule)
if child.traceEnabled {
child.traceExit(rule)
}
term := rule.Head.Value
if term == nil {
@@ -3291,7 +3276,6 @@ func (e evalVirtualPartial) evalOneRuleContinue(iter unifyIterator, rule *ast.Ru
}
term, termbindings := child.bindings.apply(term)
if rule.Head.RuleKind() == ast.MultiValue {
term = ast.SetTerm(term)
}
@@ -3300,16 +3284,14 @@ func (e evalVirtualPartial) evalOneRuleContinue(iter unifyIterator, rule *ast.Ru
term = wrapInObjects(term, objRef)
err := e.evalTerm(iter, e.pos+1, term, termbindings)
if err != nil {
return err
if child.traceEnabled && err == nil {
child.traceRedo(rule)
}
child.traceRedo(rule)
return nil
return err
}
func (e evalVirtualPartial) partialEvalSupport(iter unifyIterator) error {
path := e.e.namespaceRef(e.plugged[:e.pos+1])
term := ast.NewTerm(e.e.namespaceRef(e.ref))
@@ -4327,7 +4309,6 @@ func (e evalNot) eval(iter evalIterator) error {
defer evalPool.Put(child)
e.e.closure(e.not.Body, child)
if e.e.traceEnabled {
child.traceEnter(e.not.Body)
}
@@ -4545,13 +4526,12 @@ func evalLogicalOperand(parent *eval, body ast.Body) (bool, error) {
child.traceEnter(body)
}
defined := false
err := child.eval(func(*eval) error {
if parent.traceEnabled {
child.traceExit(body)
child.traceRedo(body)
}
defined = true
child.defined = true
return nil
})
@@ -4561,7 +4541,7 @@ func evalLogicalOperand(parent *eval, body ast.Body) (bool, error) {
return false, err
}
return defined, nil
return child.defined, nil
}
func plugBody(e *eval, body ast.Body) error {