Skip to content
Merged
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
47 changes: 45 additions & 2 deletions cel/folding.go
Original file line number Diff line number Diff line change
Expand Up @@ -222,7 +222,7 @@ func maybePruneBranches(ctx *OptimizerContext, expr ast.NavigableExpr) bool {
return true
}
needle := args[0]
if (needle.Kind() == ast.LiteralKind || needle.Kind() == ast.IdentKind) && haystack.Kind() == ast.ListKind {
if (needle.Kind() == ast.LiteralKind || isSelfEqualIdent(needle)) && haystack.Kind() == ast.ListKind {
needleIsLit := needle.Kind() == ast.LiteralKind
needleLitVal := needle.AsLiteral()
needleIdentVal := needle.AsIdent()
Expand Down Expand Up @@ -595,7 +595,7 @@ func constantCallMatcher(e ast.NavigableExpr) bool {
return true
}
needle := children[0]
if (needle.Kind() == ast.LiteralKind || needle.Kind() == ast.IdentKind) && haystack.Kind() == ast.ListKind {
if (needle.Kind() == ast.LiteralKind || isSelfEqualIdent(needle)) && haystack.Kind() == ast.ListKind {
needleIsLit := needle.Kind() == ast.LiteralKind
needleLitVal := needle.AsLiteral()
needleIdentVal := needle.AsIdent()
Expand All @@ -619,6 +619,49 @@ func constantCallMatcher(e ast.NavigableExpr) bool {
return true
}

// isSelfEqualIdent indicates whether the expression is an identifier whose static type
// guarantees that its runtime value is equal to itself.
//
// Matching an identifier against a list element by name only proves list membership when the
// value the name resolves to is self-equal. A double may be NaN, which is not equal to itself,
// and dyn, abstract, and struct types may all hold a NaN at runtime, so the check is limited
// to the scalar types which cannot, and to the aggregate types whose type parameters are
// themselves self-equal.
func isSelfEqualIdent(e ast.Expr) bool {
Comment thread
TristonianJones marked this conversation as resolved.
if e.Kind() != ast.IdentKind {
return false
}
nav, ok := e.(ast.NavigableExpr)
if !ok {
return false
}
return isSelfEqualType(nav.Type())
}

// isSelfEqualType indicates whether all runtime values of the given type are equal to themselves.
func isSelfEqualType(t *types.Type) bool {
if t == nil {
return false
}
switch t.Kind() {
case types.BoolKind, types.BytesKind, types.DurationKind, types.IntKind,
Comment thread
TristonianJones marked this conversation as resolved.
types.NullTypeKind, types.StringKind, types.TimestampKind, types.TypeKind,
types.UintKind:
return true
case types.ListKind, types.MapKind:
// Aggregates compare element-wise, so they are self-equal exactly when their type
// parameters are. A list(dyn) or map(string, double) may still contain a NaN.
for _, p := range t.Parameters() {
if !isSelfEqualType(p) {
return false
}
}
return true
default:
return false
}
}

func isExprConstantOfKind(e ast.Expr, t *types.Type) bool {
return e.Kind() == ast.LiteralKind && e.AsLiteral().Type() == t
}
Expand Down
116 changes: 115 additions & 1 deletion cel/folding_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ package cel

import (
"errors"
"math"
"reflect"
"sort"
"testing"
Expand Down Expand Up @@ -378,7 +379,7 @@ func TestConstantFoldingOptimizer(t *testing.T) {
},
{
expr: `x in [1, 2, x]`,
folded: `true`,
folded: `x in [1, 2, x]`,
Comment thread
TristonianJones marked this conversation as resolved.
},
{
expr: `[1, 2].filter(x, x in [1, 2])`,
Expand Down Expand Up @@ -549,6 +550,119 @@ func TestConstantFoldingOptimizer(t *testing.T) {
}
}

// TestConstantFoldingInListIdent checks which identifier needles may be matched against a list
// element by name.
//
// The rewrite is only sound when the identifier's static type guarantees that its runtime value
// is equal to itself. A NaN double is not, so any type which can carry a NaN must be left alone:
// the vars case of each such entry evaluates the optimized program with a NaN binding and expects
// the same result the unoptimized expression produces.
func TestConstantFoldingInListIdent(t *testing.T) {
nan := math.NaN()
tests := []struct {
expr string
folded string
// vars, when set, is evaluated against the optimized AST and must yield false.
vars map[string]any
}{
// Scalar types which cannot hold a NaN.
{expr: `b in [b]`, folded: `true`},
{expr: `by in [by]`, folded: `true`},
{expr: `du in [du]`, folded: `true`},
{expr: `i in [1, 2, i]`, folded: `true`},
{expr: `n in [n]`, folded: `true`},
{expr: `s in [s]`, folded: `true`},
{expr: `ts in [ts]`, folded: `true`},
{expr: `ty in [ty]`, folded: `true`},
{expr: `u in [u]`, folded: `true`},
// Aggregate types compare element-wise, so they are self-equal exactly when their type
// parameters are.
{expr: `li in [li]`, folded: `true`},
{expr: `lli in [lli]`, folded: `true`},
{expr: `msi in [msi]`, folded: `true`},
{expr: `ld in [ld]`, folded: `ld in [ld]`, vars: map[string]any{"ld": []float64{nan}}},
{expr: `lx in [lx]`, folded: `lx in [lx]`, vars: map[string]any{"lx": []any{nan}}},
{expr: `msd in [msd]`, folded: `msd in [msd]`, vars: map[string]any{"msd": map[string]float64{"a": nan}}},
// Doubles and dyn may be a NaN directly.
{expr: `d in [d]`, folded: `d in [d]`, vars: map[string]any{"d": nan}},
{expr: `x in [1, 2, x]`, folded: `x in [1, 2, x]`, vars: map[string]any{"x": types.Double(nan)}},
// Abstract and struct types are left alone as their contents are not inspected.
{expr: `oi in [oi]`, folded: `oi in [oi]`},
{expr: `o in [o]`, folded: `o in [o]`},
// Literal needles use CEL equality rather than a name match, so they are unaffected.
{expr: `1 in [1, 2]`, folded: `true`},
{expr: `1.0 in [d, 1.0]`, folded: `true`},
}
env, err := NewEnv(
OptionalTypes(),
Types(&proto3pb.TestAllTypes{}),
Variable("b", BoolType),
Variable("by", BytesType),
Variable("du", DurationType),
Variable("i", IntType),
Variable("n", NullType),
Variable("s", StringType),
Variable("ts", TimestampType),
Variable("ty", TypeType),
Variable("u", UintType),
Variable("li", ListType(IntType)),
Variable("lli", ListType(ListType(IntType))),
Variable("msi", MapType(StringType, IntType)),
Variable("ld", ListType(DoubleType)),
Variable("lx", ListType(DynType)),
Variable("msd", MapType(StringType, DoubleType)),
Variable("d", DoubleType),
Variable("x", DynType),
Variable("oi", OptionalType(IntType)),
Variable("o", ObjectType("google.expr.proto3.test.TestAllTypes")),
)
if err != nil {
t.Fatalf("NewEnv() failed: %v", err)
}
for _, tst := range tests {
tc := tst
t.Run(tc.expr, func(t *testing.T) {
checked, iss := env.Compile(tc.expr)
if iss.Err() != nil {
t.Fatalf("Compile() failed: %v", iss.Err())
}
folder, err := NewConstantFoldingOptimizer()
if err != nil {
t.Fatalf("NewConstantFoldingOptimizer() failed: %v", err)
}
opt, err := NewStaticOptimizer(folder)
if err != nil {
t.Fatalf("NewStaticOptimizer() failed: %v", err)
}
optimized, iss := opt.Optimize(env, checked)
if iss.Err() != nil {
t.Fatalf("Optimize() generated an invalid AST: %v", iss.Err())
}
folded, err := AstToString(optimized)
if err != nil {
t.Fatalf("AstToString() failed: %v", err)
}
if folded != tc.folded {
t.Errorf("got %q, wanted %q", folded, tc.folded)
}
if tc.vars == nil {
return
}
prg, err := env.Program(optimized)
if err != nil {
t.Fatalf("Program() failed: %v", err)
}
out, _, err := prg.Eval(tc.vars)
if err != nil {
t.Fatalf("Eval() failed: %v", err)
}
if out != types.False {
t.Errorf("got %v, wanted false since NaN is not equal to itself", out)
}
})
}
}

func TestConstantFoldingCallsWithSideEffects(t *testing.T) {
tests := []struct {
expr string
Expand Down