From 30f8e6db9acdeaf260cfc34eea297612b38caa26 Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Wed, 12 Aug 2026 13:28:25 -0700 Subject: [PATCH] Add a validation pass for ID uniqueness in optimizers PiperOrigin-RevId: 963628515 --- .../main/java/dev/cel/optimizer/BUILD.bazel | 4 + .../cel/optimizer/CelOptimizerFactory.java | 32 ++- .../dev/cel/optimizer/CelOptimizerImpl.java | 80 +++++++- .../cel/optimizer/CelOptimizerOptions.java | 51 +++++ .../optimizer/CelOptimizerFactoryTest.java | 36 ++++ .../cel/optimizer/CelOptimizerImplTest.java | 188 +++++++++++++++++- 6 files changed, 383 insertions(+), 8 deletions(-) create mode 100644 optimizer/src/main/java/dev/cel/optimizer/CelOptimizerOptions.java diff --git a/optimizer/src/main/java/dev/cel/optimizer/BUILD.bazel b/optimizer/src/main/java/dev/cel/optimizer/BUILD.bazel index e9e8994a2..22dab14f4 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/BUILD.bazel +++ b/optimizer/src/main/java/dev/cel/optimizer/BUILD.bazel @@ -32,12 +32,14 @@ java_library( srcs = [ "CelOptimizer.java", "CelOptimizerBuilder.java", + "CelOptimizerOptions.java", ], tags = [ ], deps = [ ":ast_optimizer", ":optimization_exception", + "//:auto_value", "//common:cel_ast", "@maven//:com_google_errorprone_error_prone_annotations", ], @@ -57,6 +59,8 @@ java_library( "//bundle:cel", "//common:cel_ast", "//common:compiler_common", + "//common/ast", + "//common/navigation", "@maven//:com_google_guava_guava", ], ) diff --git a/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerFactory.java b/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerFactory.java index 1ebfd293e..d82403825 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerFactory.java +++ b/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerFactory.java @@ -25,22 +25,48 @@ /** Factory class for constructing an {@link CelOptimizer} instance. */ public final class CelOptimizerFactory { + private static final CelOptimizerOptions DEFAULT_OPTIMIZER_OPTIONS = + CelOptimizerOptions.newBuilder().build(); + /** Create a new builder for constructing a {@link CelOptimizer} instance. */ public static CelOptimizerBuilder standardCelOptimizerBuilder(Cel cel) { - return CelOptimizerImpl.newBuilder(cel); + return standardCelOptimizerBuilder(cel, DEFAULT_OPTIMIZER_OPTIONS); + } + + /** Create a new builder for constructing a {@link CelOptimizer} instance with custom options. */ + public static CelOptimizerBuilder standardCelOptimizerBuilder( + Cel cel, CelOptimizerOptions optimizerOptions) { + return CelOptimizerImpl.newBuilder(cel, optimizerOptions); } /** Create a new builder for constructing a {@link CelOptimizer} instance. */ public static CelOptimizerBuilder standardCelOptimizerBuilder( CelCompiler celCompiler, CelRuntime celRuntime) { - return standardCelOptimizerBuilder(CelFactory.combine(celCompiler, celRuntime)); + return standardCelOptimizerBuilder(celCompiler, celRuntime, DEFAULT_OPTIMIZER_OPTIONS); + } + + /** Create a new builder for constructing a {@link CelOptimizer} instance with custom options. */ + public static CelOptimizerBuilder standardCelOptimizerBuilder( + CelCompiler celCompiler, CelRuntime celRuntime, CelOptimizerOptions optimizerOptions) { + return standardCelOptimizerBuilder( + CelFactory.combine(celCompiler, celRuntime), optimizerOptions); } /** Create a new builder for constructing a {@link CelOptimizer} instance. */ public static CelOptimizerBuilder standardCelOptimizerBuilder( CelParser celParser, CelChecker celChecker, CelRuntime celRuntime) { return standardCelOptimizerBuilder( - CelCompilerFactory.combine(celParser, celChecker), celRuntime); + celParser, celChecker, celRuntime, DEFAULT_OPTIMIZER_OPTIONS); + } + + /** Create a new builder for constructing a {@link CelOptimizer} instance with custom options. */ + public static CelOptimizerBuilder standardCelOptimizerBuilder( + CelParser celParser, + CelChecker celChecker, + CelRuntime celRuntime, + CelOptimizerOptions optimizerOptions) { + return standardCelOptimizerBuilder( + CelCompilerFactory.combine(celParser, celChecker), celRuntime, optimizerOptions); } private CelOptimizerFactory() {} diff --git a/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerImpl.java b/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerImpl.java index 4ac8764f1..2911d3d4a 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerImpl.java +++ b/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerImpl.java @@ -20,16 +20,25 @@ import dev.cel.bundle.Cel; import dev.cel.common.CelAbstractSyntaxTree; import dev.cel.common.CelValidationException; +import dev.cel.common.ast.CelExpr; +import dev.cel.common.ast.CelExpr.ExprKind.Kind; +import dev.cel.common.navigation.CelNavigableAst; +import dev.cel.common.navigation.CelNavigableExpr; import dev.cel.optimizer.CelAstOptimizer.OptimizationResult; import java.util.Arrays; +import java.util.HashMap; +import java.util.Map; final class CelOptimizerImpl implements CelOptimizer { private final Cel cel; private final ImmutableSet astOptimizers; + private final CelOptimizerOptions optimizerOptions; - CelOptimizerImpl(Cel cel, ImmutableSet astOptimizers) { + CelOptimizerImpl( + Cel cel, ImmutableSet astOptimizers, CelOptimizerOptions optimizerOptions) { this.cel = cel; this.astOptimizers = astOptimizers; + this.optimizerOptions = optimizerOptions; } @Override @@ -52,6 +61,9 @@ public CelAbstractSyntaxTree optimize(CelAbstractSyntaxTree ast) throws CelOptim .build(); } optimizedAst = celOptimizerEnv.check(result.optimizedAst()).getAst(); + if (optimizerOptions.enableAstValidation()) { + assertAstIdCorrectness(optimizedAst); + } } } catch (CelValidationException e) { throw new CelOptimizationException( @@ -63,18 +75,78 @@ public CelAbstractSyntaxTree optimize(CelAbstractSyntaxTree ast) throws CelOptim return optimizedAst; } + private static void assertAstIdCorrectness(CelAbstractSyntaxTree ast) { + Map allExprs = new HashMap<>(); + CelNavigableAst.fromAst(ast) + .getRoot() + .allNodes() + .forEach( + navExpr -> { + CelExpr expr = navExpr.expr(); + CelExpr existing = allExprs.put(expr.id(), expr); + if (existing != null) { + throw new IllegalStateException( + String.format("Duplicate expr ID %d detected in the AST.", expr.id())); + } + }); + + for (CelExpr macroCall : ast.getSource().getMacroCalls().values()) { + if (macroCall.id() != 0) { + throw new IllegalStateException( + String.format("Expected macro call root ID to be 0, but was %d.", macroCall.id())); + } + CelNavigableExpr.fromExpr(macroCall) + .descendants() + .forEach( + navExpr -> { + CelExpr macroExpr = navExpr.expr(); + CelExpr astExpr = allExprs.get(macroExpr.id()); + // A node may not exist in the AST if it is a synthetic macro node or was eliminated + // during optimization passes. + if (astExpr == null) { + return; + } + + if (astExpr.exprKind().getKind().equals(Kind.COMPREHENSION)) { + if (!macroExpr.exprKind().getKind().equals(Kind.NOT_SET)) { + throw new IllegalStateException( + String.format( + "Expected macro call node %d to be NOT_SET for comprehension, but" + + " was %s.", + macroExpr.id(), macroExpr.exprKind().getKind())); + } + } else if (!macroExpr.exprKind().getKind().equals(astExpr.exprKind().getKind())) { + throw new IllegalStateException( + String.format( + "Macro call node %d kind mismatch: expected %s (from AST), but was %s" + + " (in macro call).", + macroExpr.id(), + astExpr.exprKind().getKind(), + macroExpr.exprKind().getKind())); + } + }); + } + } + /** Create a new builder for constructing a {@link CelOptimizer} instance. */ static CelOptimizerImpl.Builder newBuilder(Cel cel) { - return new CelOptimizerImpl.Builder(cel); + return newBuilder(cel, CelOptimizerOptions.newBuilder().build()); + } + + /** Create a new builder for constructing a {@link CelOptimizer} instance with custom options. */ + static CelOptimizerImpl.Builder newBuilder(Cel cel, CelOptimizerOptions optimizerOptions) { + return new CelOptimizerImpl.Builder(cel, optimizerOptions); } /** Builder class for {@link CelOptimizerImpl}. */ static final class Builder implements CelOptimizerBuilder { private final Cel cel; + private final CelOptimizerOptions optimizerOptions; private final ImmutableSet.Builder astOptimizers; - private Builder(Cel cel) { + private Builder(Cel cel, CelOptimizerOptions optimizerOptions) { this.cel = cel; + this.optimizerOptions = checkNotNull(optimizerOptions); this.astOptimizers = ImmutableSet.builder(); } @@ -93,7 +165,7 @@ public CelOptimizerBuilder addAstOptimizers(Iterable astOptimiz @Override public CelOptimizer build() { - return new CelOptimizerImpl(cel, astOptimizers.build()); + return new CelOptimizerImpl(cel, astOptimizers.build(), optimizerOptions); } } } diff --git a/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerOptions.java b/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerOptions.java new file mode 100644 index 000000000..888683298 --- /dev/null +++ b/optimizer/src/main/java/dev/cel/optimizer/CelOptimizerOptions.java @@ -0,0 +1,51 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package dev.cel.optimizer; + +import com.google.auto.value.AutoValue; + +/** Options to configure how {@link CelOptimizer} behaves. */ +@AutoValue +public abstract class CelOptimizerOptions { + + /** + * Returns true if AST validation is enabled. When enabled, each optimizer pass verifies AST + * invariants (such as expression ID uniqueness and macro source consistency) after type-checking. + */ + public abstract boolean enableAstValidation(); + + /** Builder for configuring the {@link CelOptimizerOptions}. */ + @AutoValue.Builder + public abstract static class Builder { + + /** + * Enables or disables post-pass AST validation. When enabled, each optimizer pass verifies that + * expression IDs are unique and macro calls in the AST source are consistent with the + * expression nodes. + */ + public abstract Builder enableAstValidation(boolean value); + + public abstract CelOptimizerOptions build(); + + Builder() {} + } + + /** Returns a new options builder with recommended defaults pre-configured. */ + public static Builder newBuilder() { + return new AutoValue_CelOptimizerOptions.Builder().enableAstValidation(false); + } + + CelOptimizerOptions() {} +} diff --git a/optimizer/src/test/java/dev/cel/optimizer/CelOptimizerFactoryTest.java b/optimizer/src/test/java/dev/cel/optimizer/CelOptimizerFactoryTest.java index 41c7ecd74..146102995 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/CelOptimizerFactoryTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/CelOptimizerFactoryTest.java @@ -39,6 +39,19 @@ public void standardCelOptimizerBuilder_withParserCheckerAndRuntime() { assertThat(builder.build()).isNotNull(); } + @Test + public void standardCelOptimizerBuilder_withParserCheckerRuntimeAndOptions() { + CelOptimizerBuilder builder = + CelOptimizerFactory.standardCelOptimizerBuilder( + CelParserFactory.standardCelParserBuilder().build(), + CelCompilerFactory.standardCelCheckerBuilder().build(), + CelRuntimeFactory.standardCelRuntimeBuilder().build(), + CelOptimizerOptions.newBuilder().enableAstValidation(true).build()); + + assertThat(builder).isNotNull(); + assertThat(builder.build()).isNotNull(); + } + @Test public void standardCelOptimizerBuilder_withCompilerAndRuntime() { CelOptimizerBuilder builder = @@ -50,6 +63,18 @@ public void standardCelOptimizerBuilder_withCompilerAndRuntime() { assertThat(builder.build()).isNotNull(); } + @Test + public void standardCelOptimizerBuilder_withCompilerRuntimeAndOptions() { + CelOptimizerBuilder builder = + CelOptimizerFactory.standardCelOptimizerBuilder( + CelCompilerFactory.standardCelCompilerBuilder().build(), + CelRuntimeFactory.standardCelRuntimeBuilder().build(), + CelOptimizerOptions.newBuilder().enableAstValidation(true).build()); + + assertThat(builder).isNotNull(); + assertThat(builder.build()).isNotNull(); + } + @Test public void standardCelOptimizerBuilder_withCel() { CelOptimizerBuilder builder = @@ -58,4 +83,15 @@ public void standardCelOptimizerBuilder_withCel() { assertThat(builder).isNotNull(); assertThat(builder.build()).isNotNull(); } + + @Test + public void standardCelOptimizerBuilder_withCelAndOptions() { + CelOptimizerBuilder builder = + CelOptimizerFactory.standardCelOptimizerBuilder( + CelFactory.standardCelBuilder().build(), + CelOptimizerOptions.newBuilder().enableAstValidation(true).build()); + + assertThat(builder).isNotNull(); + assertThat(builder.build()).isNotNull(); + } } diff --git a/optimizer/src/test/java/dev/cel/optimizer/CelOptimizerImplTest.java b/optimizer/src/test/java/dev/cel/optimizer/CelOptimizerImplTest.java index cb0bff6c6..4373e7fe4 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/CelOptimizerImplTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/CelOptimizerImplTest.java @@ -17,15 +17,20 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.Assert.assertThrows; +import com.google.common.collect.ImmutableList; import dev.cel.bundle.Cel; import dev.cel.bundle.CelFactory; import dev.cel.common.CelAbstractSyntaxTree; +import dev.cel.common.CelOptions; import dev.cel.common.CelSource; import dev.cel.common.CelValidationException; +import dev.cel.common.ast.CelConstant; import dev.cel.common.ast.CelExpr; import dev.cel.optimizer.CelAstOptimizer.OptimizationResult; +import dev.cel.parser.CelStandardMacro; import java.util.ArrayList; import java.util.List; +import java.util.Optional; import org.junit.Test; import org.junit.runner.RunWith; import org.junit.runners.JUnit4; @@ -33,7 +38,11 @@ @RunWith(JUnit4.class) public class CelOptimizerImplTest { - private static final Cel CEL = CelFactory.standardCelBuilder().build(); + private static final Cel CEL = + CelFactory.standardCelBuilder() + .setOptions(CelOptions.current().populateMacroCalls(true).build()) + .setStandardMacros(CelStandardMacro.STANDARD_MACROS) + .build(); @Test public void constructCelOptimizer_success() { @@ -131,4 +140,181 @@ public void optimizedAst_failsToTypeCheck_throwsException() { + " 'undeclared_ident' (in container '')"); assertThat(e).hasCauseThat().isInstanceOf(CelValidationException.class); } + + @Test + public void optimize_duplicateExprId_throwsException() { + CelOptimizer celOptimizer = + CelOptimizerImpl.newBuilder( + CEL, CelOptimizerOptions.newBuilder().enableAstValidation(true).build()) + .addAstOptimizers( + (navigableAst, cel) -> + OptimizationResult.create( + CelAbstractSyntaxTree.newParsedAst( + CelExpr.ofCall( + 1, + Optional.empty(), + "_+_", + ImmutableList.of( + CelExpr.ofConstant(1, CelConstant.ofValue(1L)), + CelExpr.ofConstant(2, CelConstant.ofValue(2L)))), + CelSource.newBuilder().build()))) + .build(); + + CelOptimizationException e = + assertThrows( + CelOptimizationException.class, + () -> celOptimizer.optimize(CEL.compile("1 + 2").getAst())); + + assertThat(e) + .hasMessageThat() + .isEqualTo("Optimization failure: Duplicate expr ID 1 detected in the AST."); + assertThat(e).hasCauseThat().isInstanceOf(IllegalStateException.class); + } + + @Test + public void optimize_macroCallRootIdNonZero_throwsException() { + CelOptimizer celOptimizer = + CelOptimizerImpl.newBuilder( + CEL, CelOptimizerOptions.newBuilder().enableAstValidation(true).build()) + .addAstOptimizers( + (navigableAst, cel) -> + OptimizationResult.create( + CelAbstractSyntaxTree.newParsedAst( + CelExpr.ofConstant(1, CelConstant.ofValue(1L)), + CelSource.newBuilder() + .addMacroCalls( + 1L, + CelExpr.ofCall( + 10L, + Optional.empty(), + "has", + ImmutableList.of( + CelExpr.ofConstant(1L, CelConstant.ofValue(1L))))) + .build()))) + .build(); + + CelOptimizationException e = + assertThrows( + CelOptimizationException.class, () -> celOptimizer.optimize(CEL.compile("1").getAst())); + + assertThat(e) + .hasMessageThat() + .isEqualTo("Optimization failure: Expected macro call root ID to be 0, but was 10."); + assertThat(e).hasCauseThat().isInstanceOf(IllegalStateException.class); + } + + @Test + public void optimize_macroCallKindMismatch_throwsException() { + CelOptimizer celOptimizer = + CelOptimizerImpl.newBuilder( + CEL, CelOptimizerOptions.newBuilder().enableAstValidation(true).build()) + .addAstOptimizers( + (navigableAst, cel) -> + OptimizationResult.create( + CelAbstractSyntaxTree.newParsedAst( + CelExpr.ofConstant(1, CelConstant.ofValue(1L)), + CelSource.newBuilder() + .addMacroCalls( + 1L, + CelExpr.ofCall( + 0L, + Optional.empty(), + "has", + ImmutableList.of(CelExpr.ofIdent(1L, "x")))) + .build()))) + .build(); + + CelOptimizationException e = + assertThrows( + CelOptimizationException.class, () -> celOptimizer.optimize(CEL.compile("1").getAst())); + + assertThat(e) + .hasMessageThat() + .isEqualTo( + "Optimization failure: Macro call node 1 kind mismatch: expected CONSTANT (from AST)," + + " but was IDENT (in macro call)."); + assertThat(e).hasCauseThat().isInstanceOf(IllegalStateException.class); + } + + @Test + public void optimize_macroCallComprehensionKindNotSetMismatch_throwsException() throws Exception { + CelAbstractSyntaxTree astWithComprehension = CEL.compile("[1].all(x, x > 0)").getAst(); + long compId = astWithComprehension.getExpr().id(); + + CelOptimizer celOptimizer = + CelOptimizerImpl.newBuilder( + CEL, CelOptimizerOptions.newBuilder().enableAstValidation(true).build()) + .addAstOptimizers( + (navigableAst, cel) -> + OptimizationResult.create( + CelAbstractSyntaxTree.newParsedAst( + astWithComprehension.getExpr(), + CelSource.newBuilder() + .addMacroCalls( + compId, + CelExpr.ofCall( + 0L, + Optional.empty(), + "all", + ImmutableList.of( + CelExpr.ofIdent(compId, "not_set_expected")))) + .build()))) + .build(); + + CelOptimizationException e = + assertThrows( + CelOptimizationException.class, () -> celOptimizer.optimize(astWithComprehension)); + + assertThat(e) + .hasMessageThat() + .isEqualTo( + String.format( + "Optimization failure: Expected macro call node %d to be NOT_SET for comprehension," + + " but was IDENT.", + compId)); + assertThat(e).hasCauseThat().isInstanceOf(IllegalStateException.class); + } + + @Test + public void optimize_macroCallComprehensionKindNotSet_success() throws Exception { + CelAbstractSyntaxTree astWithComprehension = CEL.compile("[1].all(x, x > 0)").getAst(); + long compId = astWithComprehension.getExpr().id(); + + CelOptimizer celOptimizer = + CelOptimizerImpl.newBuilder(CEL) + .addAstOptimizers( + (navigableAst, cel) -> + OptimizationResult.create( + CelAbstractSyntaxTree.newParsedAst( + astWithComprehension.getExpr(), + CelSource.newBuilder() + .addMacroCalls( + compId, + CelExpr.ofCall( + 0L, + Optional.empty(), + "all", + ImmutableList.of(CelExpr.ofNotSet(compId)))) + .build()))) + .build(); + + CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(astWithComprehension); + + assertThat(optimizedAst).isNotNull(); + } + + @Test + public void optimize_validMacroCalls_success() throws Exception { + CelAbstractSyntaxTree ast = CEL.compile("[1, 2, 3].all(x, x > 0)").getAst(); + + CelOptimizer celOptimizer = + CelOptimizerImpl.newBuilder(CEL) + .addAstOptimizers((navigableAst, cel) -> OptimizationResult.create(navigableAst)) + .build(); + + CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast); + + assertThat(optimizedAst).isNotNull(); + assertThat(optimizedAst.getSource().getMacroCalls()).hasSize(1); + } }