diff --git a/conformance/src/test/java/dev/cel/conformance/BUILD.bazel b/conformance/src/test/java/dev/cel/conformance/BUILD.bazel index c5364b146..4abc705c3 100644 --- a/conformance/src/test/java/dev/cel/conformance/BUILD.bazel +++ b/conformance/src/test/java/dev/cel/conformance/BUILD.bazel @@ -20,10 +20,14 @@ java_library( "//common:compiler_common", "//common:container", "//common:options", + "//common/ast", + "//common/ast:cel_block", + "//common/types", "//common/types:cel_proto_types", "//compiler", "//compiler:compiler_builder", "//extensions", + "//extensions:bindings", "//extensions:optional_library", "//parser:macro", "//parser:parser_builder", @@ -75,6 +79,7 @@ java_library( _ALL_TESTS = [ "@cel_spec//tests/simple:testdata/basic.textproto", "@cel_spec//tests/simple:testdata/bindings_ext.textproto", + "@cel_spec//tests/simple:testdata/block_ext.textproto", "@cel_spec//tests/simple:testdata/comparisons.textproto", "@cel_spec//tests/simple:testdata/conversions.textproto", "@cel_spec//tests/simple:testdata/dynamic.textproto", diff --git a/conformance/src/test/java/dev/cel/conformance/ConformanceTest.java b/conformance/src/test/java/dev/cel/conformance/ConformanceTest.java index db57ccb79..82b4cf812 100644 --- a/conformance/src/test/java/dev/cel/conformance/ConformanceTest.java +++ b/conformance/src/test/java/dev/cel/conformance/ConformanceTest.java @@ -30,16 +30,28 @@ import com.google.protobuf.ExtensionRegistry; import com.google.protobuf.TypeRegistry; import dev.cel.checker.CelChecker; +import dev.cel.checker.CelCheckerBuilder; import dev.cel.common.CelContainer; +import dev.cel.common.CelIssue; import dev.cel.common.CelOptions; import dev.cel.common.CelValidationResult; +import dev.cel.common.CelVarDecl; +import dev.cel.common.ast.CelBlock; +import dev.cel.common.ast.CelConstant; +import dev.cel.common.ast.CelExpr; import dev.cel.common.types.CelProtoTypes; +import dev.cel.common.types.SimpleType; import dev.cel.compiler.CelCompilerFactory; import dev.cel.compiler.CelCompilerLibrary; import dev.cel.expr.conformance.test.SimpleTest; +import dev.cel.extensions.CelBindingsExtensions; import dev.cel.extensions.CelExtensions; import dev.cel.extensions.CelOptionalLibrary; +import dev.cel.parser.CelMacro; +import dev.cel.parser.CelMacroExpander; +import dev.cel.parser.CelMacroExprFactory; import dev.cel.parser.CelParser; +import dev.cel.parser.CelParserBuilder; import dev.cel.parser.CelParserFactory; import dev.cel.parser.CelStandardMacro; import dev.cel.runtime.CelEvaluationException; @@ -50,6 +62,7 @@ import dev.cel.runtime.CelRuntimeImpl; import dev.cel.runtime.CelRuntimeLibrary; import java.util.Map; +import java.util.Optional; import org.junit.runners.model.Statement; // Qualifying proto2/proto3 TestAllTypes makes it less clear. @@ -73,7 +86,8 @@ public final class ConformanceTest extends Statement { CelExtensions.protos(), CelExtensions.sets(OPTIONS), CelExtensions.strings(), - CelOptionalLibrary.INSTANCE); + CelOptionalLibrary.INSTANCE, + new ConformanceBlockLibrary()); private static final ImmutableList CANONICAL_RUNTIME_EXTENSIONS = ImmutableList.of( @@ -206,11 +220,11 @@ public boolean shouldSkip() { @Override public void evaluate() throws Throwable { CelValidationResult response = getParser(test).parse(test.getExpr(), test.getName()); - assertThat(response.hasError()).isFalse(); + assertThat(response.getErrors()).isEmpty(); if (!test.getDisableCheck()) { response = getChecker(test).check(response.getAst()); } - assertThat(response.hasError()).isFalse(); + assertThat(response.getErrors()).isEmpty(); Type resultType = CelProtoTypes.celTypeToType(response.getAst().getResultType()); if (test.getCheckOnly()) { @@ -262,4 +276,103 @@ public void evaluate() throws Throwable { String.format("Unexpected matcher kind: %s", test.getResultMatcherCase())); } } + + /** + * Conformance-only library providing macros for the {@code block_ext} test suite. + * + *

These macros ({@code cel.block}, {@code cel.index}, {@code cel.iterVar}, and {@code + * cel.accuVar}) are strictly used for conformance testing to represent block expressions in text + * form. In production, AST optimization passes (such as common subexpression elimination) + * directly generate the {@code cel.@block} call and {@code @index} / {@code @it} / {@code @ac} + * variable nodes without going through these macros. + */ + private static final class ConformanceBlockLibrary implements CelCompilerLibrary { + private static final int MAX_INDICES = 30; + + @Override + public void setParserOptions(CelParserBuilder parserBuilder) { + parserBuilder.addMacros( + CelMacro.newReceiverMacro("block", 2, ConformanceBlockLibrary::expandBlock), + CelMacro.newReceiverMacro("index", 1, ConformanceBlockLibrary::expandIndex), + CelMacro.newReceiverMacro("iterVar", 2, expandCompreVar("cel.iterVar", "@it")), + CelMacro.newReceiverMacro("accuVar", 2, expandCompreVar("cel.accuVar", "@ac"))); + } + + @Override + public void setCheckerOptions(CelCheckerBuilder checkerBuilder) { + checkerBuilder.addFunctionDeclarations(CelBindingsExtensions.CEL_BLOCK_FUNCTION_DECL); + for (int i = 0; i < MAX_INDICES; i++) { + checkerBuilder.addVarDeclarations( + CelVarDecl.newVarDeclaration(CelBlock.INDEX_PREFIX + i, SimpleType.DYN)); + } + } + + private static Optional expandBlock( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList args) { + if (!isCelNamespace(target)) { + return Optional.empty(); + } + CelExpr bindings = args.get(0); + if (!bindings.exprKind().getKind().equals(CelExpr.ExprKind.Kind.LIST)) { + return Optional.of( + exprFactory.reportError( + CelIssue.formatError( + exprFactory.getSourceLocation(bindings), + "cel.block requires the first arg to be a list literal"))); + } + return Optional.of(exprFactory.newGlobalCall(CelBlock.FUNCTION_NAME, args)); + } + + private static Optional expandIndex( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList args) { + if (!isCelNamespace(target)) { + return Optional.empty(); + } + CelExpr index = args.get(0); + if (!isNonNegativeInt(index)) { + return Optional.of( + exprFactory.reportError( + CelIssue.formatError( + exprFactory.getSourceLocation(index), + "cel.index requires a single non-negative int constant arg"))); + } + return Optional.of( + exprFactory.newIdentifier(CelBlock.INDEX_PREFIX + index.constant().int64Value())); + } + + private static CelMacroExpander expandCompreVar(String macroName, String prefix) { + return (exprFactory, target, args) -> { + if (!isCelNamespace(target)) { + return Optional.empty(); + } + for (CelExpr arg : args) { + if (!isNonNegativeInt(arg)) { + return Optional.of( + exprFactory.reportError( + CelIssue.formatError( + exprFactory.getSourceLocation(arg), + macroName + " requires two non-negative int constant args"))); + } + } + return Optional.of( + exprFactory.newIdentifier( + String.format( + "%s:%d:%d", + prefix, + args.get(0).constant().int64Value(), + args.get(1).constant().int64Value()))); + }; + } + + private static boolean isNonNegativeInt(CelExpr expr) { + return expr.exprKind().getKind().equals(CelExpr.ExprKind.Kind.CONSTANT) + && expr.constant().getKind().equals(CelConstant.Kind.INT64_VALUE) + && expr.constant().int64Value() >= 0; + } + + private static boolean isCelNamespace(CelExpr target) { + return target.exprKind().getKind().equals(CelExpr.ExprKind.Kind.IDENT) + && target.ident().name().equals("cel"); + } + } } diff --git a/extensions/BUILD.bazel b/extensions/BUILD.bazel index dea4cd760..f9c2aee45 100644 --- a/extensions/BUILD.bazel +++ b/extensions/BUILD.bazel @@ -61,3 +61,9 @@ java_library( name = "native", exports = ["//extensions/src/main/java/dev/cel/extensions:native"], ) + +java_library( + name = "bindings", + visibility = ["//:internal"], + exports = ["//extensions/src/main/java/dev/cel/extensions:bindings"], +) diff --git a/extensions/src/main/java/dev/cel/extensions/BUILD.bazel b/extensions/src/main/java/dev/cel/extensions/BUILD.bazel index ba57a07c3..8b7991cc0 100644 --- a/extensions/src/main/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/main/java/dev/cel/extensions/BUILD.bazel @@ -145,6 +145,7 @@ java_library( deps = [ "//common:compiler_common", "//common/ast", + "//common/ast:cel_block", "//common/types", "//compiler:compiler_builder", "//extensions:extension_library", diff --git a/extensions/src/main/java/dev/cel/extensions/CelBindingsExtensions.java b/extensions/src/main/java/dev/cel/extensions/CelBindingsExtensions.java index 0e6537334..9fea7f481 100644 --- a/extensions/src/main/java/dev/cel/extensions/CelBindingsExtensions.java +++ b/extensions/src/main/java/dev/cel/extensions/CelBindingsExtensions.java @@ -23,6 +23,7 @@ import dev.cel.common.CelFunctionDecl; import dev.cel.common.CelIssue; import dev.cel.common.CelOverloadDecl; +import dev.cel.common.ast.CelBlock; import dev.cel.common.ast.CelExpr; import dev.cel.common.types.ListType; import dev.cel.common.types.SimpleType; @@ -59,6 +60,15 @@ static CelExtensionLibrary library() { return LIBRARY; } + public static final CelFunctionDecl CEL_BLOCK_FUNCTION_DECL = + CelFunctionDecl.newFunctionDeclaration( + CelBlock.FUNCTION_NAME, + CelOverloadDecl.newGlobalOverload( + "cel_block_list", + TypeParamType.create("T"), + ListType.create(SimpleType.DYN), + TypeParamType.create("T"))); + @Override public int version() { return 0; @@ -67,14 +77,7 @@ public int version() { @Override public ImmutableSet functions() { // TODO: Add bindings for block once decorator support is available. - return ImmutableSet.of( - CelFunctionDecl.newFunctionDeclaration( - "cel.@block", - CelOverloadDecl.newGlobalOverload( - "cel_block_list", - TypeParamType.create("T"), - ListType.create(SimpleType.DYN), - TypeParamType.create("T")))); + return ImmutableSet.of(CEL_BLOCK_FUNCTION_DECL); } @Override diff --git a/optimizer/src/main/java/dev/cel/optimizer/optimizers/BUILD.bazel b/optimizer/src/main/java/dev/cel/optimizer/optimizers/BUILD.bazel index 1012b19c2..0e6509c44 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/optimizers/BUILD.bazel +++ b/optimizer/src/main/java/dev/cel/optimizer/optimizers/BUILD.bazel @@ -73,6 +73,7 @@ java_library( "//common/navigation:mutable_navigation", "//common/types", "//common/types:type_providers", + "//extensions:bindings", "//optimizer:ast_optimizer", "//optimizer:mutable_ast", "@maven//:com_google_errorprone_error_prone_annotations", diff --git a/optimizer/src/main/java/dev/cel/optimizer/optimizers/SubexpressionOptimizer.java b/optimizer/src/main/java/dev/cel/optimizer/optimizers/SubexpressionOptimizer.java index 6a9860750..6d671b162 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/optimizers/SubexpressionOptimizer.java +++ b/optimizer/src/main/java/dev/cel/optimizer/optimizers/SubexpressionOptimizer.java @@ -34,7 +34,6 @@ import dev.cel.common.CelFunctionDecl; import dev.cel.common.CelMutableAst; import dev.cel.common.CelMutableSource; -import dev.cel.common.CelOverloadDecl; import dev.cel.common.CelSource; import dev.cel.common.CelSource.Extension; import dev.cel.common.CelSource.Extension.Component; @@ -55,8 +54,8 @@ import dev.cel.common.navigation.CelNavigableMutableExpr; import dev.cel.common.navigation.TraversalOrder; import dev.cel.common.types.CelType; -import dev.cel.common.types.ListType; import dev.cel.common.types.SimpleType; +import dev.cel.extensions.CelBindingsExtensions; import dev.cel.optimizer.AstMutator; import dev.cel.optimizer.AstMutator.MangledComprehensionAst; import dev.cel.optimizer.CelAstOptimizer; @@ -98,8 +97,6 @@ public final class SubexpressionOptimizer implements CelAstOptimizer { private static final SubexpressionOptimizer INSTANCE = new SubexpressionOptimizer(SubexpressionOptimizerOptions.newBuilder().build()); private static final String BIND_IDENTIFIER_PREFIX = "@r"; - private static final String CEL_BLOCK_FUNCTION = "cel.@block"; - private static final String BLOCK_INDEX_PREFIX = "@index"; private static final Extension CEL_BLOCK_AST_EXTENSION_TAG = Extension.create("cel_block", Version.of(1L, 1L), Component.COMPONENT_RUNTIME); @@ -165,7 +162,7 @@ private OptimizationResult optimizeUsingCelBlock(CelAbstractSyntaxTree ast, Cel CelMutableExpr targetCseShape = normalizeForEquality(cseCandidates.get(0)); subexpressions.add(cseCandidates.get(0)); - String blockIdentifier = BLOCK_INDEX_PREFIX + blockIdentifierIndex++; + String blockIdentifier = CelBlock.INDEX_PREFIX + blockIdentifierIndex++; // Replace all CSE candidates with new block index identifier astToModify = @@ -217,7 +214,7 @@ private OptimizationResult optimizeUsingCelBlock(CelAbstractSyntaxTree ast, Cel // Wrap the optimized expression in cel.block astToModify = - astMutator.wrapAstWithNewCelBlock(CEL_BLOCK_FUNCTION, astToModify, subexpressions); + astMutator.wrapAstWithNewCelBlock(CelBlock.FUNCTION_NAME, astToModify, subexpressions); astToModify = astMutator.renumberIdsConsecutively(astToModify); // Tag the AST with cel.block designated as an extension @@ -226,7 +223,7 @@ private OptimizationResult optimizeUsingCelBlock(CelAbstractSyntaxTree ast, Cel return OptimizationResult.create( optimizedAst, newVarDecls.build(), - ImmutableList.of(newCelBlockFunctionDecl(ast.getResultType()))); + ImmutableList.of(CelBindingsExtensions.CEL_BLOCK_FUNCTION_DECL)); } /** @@ -595,11 +592,8 @@ private CelMutableExpr normalizeForEquality(CelMutableExpr mutableExpr) { } @VisibleForTesting - static CelFunctionDecl newCelBlockFunctionDecl(CelType resultType) { - return CelFunctionDecl.newFunctionDeclaration( - CEL_BLOCK_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "cel_block_list", resultType, ListType.create(SimpleType.DYN), resultType)); + static CelFunctionDecl newCelBlockFunctionDecl(CelType unusedResultType) { + return CelBindingsExtensions.CEL_BLOCK_FUNCTION_DECL; } /** Options to configure how Common Subexpression Elimination behave. */