diff --git a/bundle/src/main/java/dev/cel/bundle/CelBuilder.java b/bundle/src/main/java/dev/cel/bundle/CelBuilder.java index f603b479f..a45f846e4 100644 --- a/bundle/src/main/java/dev/cel/bundle/CelBuilder.java +++ b/bundle/src/main/java/dev/cel/bundle/CelBuilder.java @@ -211,6 +211,9 @@ public interface CelBuilder { @CanIgnoreReturnValue CelBuilder setValueProvider(CelValueProvider celValueProvider); + /** Returns the configured {@link CelValueProvider}, or null if not set. */ + CelValueProvider valueProvider(); + /** * Set the {@code typeProvider} for use with type-checking expressions. * diff --git a/bundle/src/main/java/dev/cel/bundle/CelImpl.java b/bundle/src/main/java/dev/cel/bundle/CelImpl.java index f0db128c1..999f1573a 100644 --- a/bundle/src/main/java/dev/cel/bundle/CelImpl.java +++ b/bundle/src/main/java/dev/cel/bundle/CelImpl.java @@ -317,6 +317,11 @@ public CelBuilder setValueProvider(CelValueProvider celValueProvider) { return this; } + @Override + public CelValueProvider valueProvider() { + return runtimeBuilder.valueProvider(); + } + @Override @Deprecated public Builder setTypeProvider(TypeProvider typeProvider) { 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 35476a792..da722d521 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/optimizers/BUILD.bazel +++ b/optimizer/src/main/java/dev/cel/optimizer/optimizers/BUILD.bazel @@ -31,6 +31,10 @@ java_library( "//common/navigation:common", "//common/navigation:mutable_navigation", "//common/types", + "//common/types:type_providers", + "//common/values", + "//common/values:cel_value", + "//common/values:cel_value_provider", "//extensions:optional_library", "//optimizer:ast_optimizer", "//optimizer:mutable_ast", @@ -40,6 +44,7 @@ java_library( "//runtime:unknown_attributes", "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", + "@maven//:org_jspecify_jspecify", ], ) diff --git a/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java b/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java index 35d181905..0fcbb497c 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java +++ b/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java @@ -24,6 +24,7 @@ import com.google.common.collect.ImmutableSet; import com.google.errorprone.annotations.CanIgnoreReturnValue; import dev.cel.bundle.Cel; +import dev.cel.bundle.CelBuilder; import dev.cel.common.CelAbstractSyntaxTree; import dev.cel.common.CelMutableAst; import dev.cel.common.CelSource; @@ -42,7 +43,13 @@ import dev.cel.common.navigation.CelNavigableMutableAst; import dev.cel.common.navigation.CelNavigableMutableExpr; import dev.cel.common.navigation.TraversalOrder; +import dev.cel.common.types.CelType; +import dev.cel.common.types.CelTypeProvider; import dev.cel.common.types.SimpleType; +import dev.cel.common.types.StructType; +import dev.cel.common.values.CelValue; +import dev.cel.common.values.CelValueProvider; +import dev.cel.common.values.StructValue; import dev.cel.extensions.CelOptionalLibrary.Function; import dev.cel.optimizer.AstMutator; import dev.cel.optimizer.CelAstOptimizer; @@ -59,6 +66,7 @@ import java.util.List; import java.util.Map; import java.util.Optional; +import org.jspecify.annotations.Nullable; /** * Performs optimization for inlining constant scalar and aggregate literal values within function @@ -95,8 +103,16 @@ private static CelMutableExpr newOptionalNoneExpr() { @Override public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel) throws CelOptimizationException { + CelBuilder builder = cel.toCelBuilder(); + CelValueProvider valueProvider; + try { + valueProvider = builder.valueProvider(); + } catch (UnsupportedOperationException e) { + // Legacy runtime does not support valueProvider and may throw. + valueProvider = null; + } // Override the environment's expected type to generally allow all subtrees to be folded. - Cel optimizerEnv = cel.toCelBuilder().setResultType(SimpleType.DYN).build(); + Cel optimizerEnv = builder.setResultType(SimpleType.DYN).build(); CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast); int iterCount = 0; @@ -123,7 +139,7 @@ public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel) if (!mutatedResult.isPresent()) { // Evaluate the call then fold try { - mutatedResult = maybeFold(optimizerEnv, mutableAst, foldableExpr); + mutatedResult = maybeFold(optimizerEnv, valueProvider, mutableAst, foldableExpr); } catch (CelEvaluationException e) { throw new CelOptimizationException( "Constant folding failure. Failed to evaluate subtree due to: " + e.getMessage(), @@ -290,7 +306,10 @@ private static boolean isNestedComprehension(CelNavigableMutableExpr expr) { } private Optional maybeFold( - Cel cel, CelMutableAst mutableAst, CelNavigableMutableExpr node) + Cel cel, + CelValueProvider valueProvider, + CelMutableAst mutableAst, + CelNavigableMutableExpr node) throws CelOptimizationException, CelEvaluationException { Object result; try { @@ -305,10 +324,12 @@ private Optional maybeFold( // ex2: optional.ofNonZeroValue(5) -> optional.of(5) if (result instanceof Optional) { Optional optResult = ((Optional) result); - return maybeRewriteOptional(optResult, mutableAst, node.expr()); + return maybeRewriteOptional( + cel.getTypeProvider(), valueProvider, optResult, mutableAst, node.expr()); } - CelMutableExpr adaptedResult = maybeAdaptEvaluatedResult(result).orElse(null); + CelMutableExpr adaptedResult = + maybeAdaptEvaluatedResult(cel.getTypeProvider(), valueProvider, result).orElse(null); if (adaptedResult == null) { return Optional.empty(); } @@ -316,14 +337,20 @@ private Optional maybeFold( return Optional.of(astMutator.replaceSubtree(mutableAst, adaptedResult, node.id())); } - private Optional maybeAdaptEvaluatedResult(Object result) { + private Optional maybeAdaptEvaluatedResult( + CelTypeProvider typeProvider, @Nullable CelValueProvider valueProvider, Object result) { + if (valueProvider != null && !(result instanceof CelValue)) { + result = valueProvider.celValueConverter().toRuntimeValue(result); + } + if (CelConstant.isConstantValue(result)) { return Optional.of(CelMutableExpr.ofConstant(CelConstant.ofObjectValue(result))); } else if (result instanceof Collection) { Collection collection = (Collection) result; List listElements = new ArrayList<>(); for (Object evaluatedElement : collection) { - CelMutableExpr adaptedExpr = maybeAdaptEvaluatedResult(evaluatedElement).orElse(null); + CelMutableExpr adaptedExpr = + maybeAdaptEvaluatedResult(typeProvider, valueProvider, evaluatedElement).orElse(null); if (adaptedExpr == null) { return Optional.empty(); } @@ -335,11 +362,13 @@ private Optional maybeAdaptEvaluatedResult(Object result) { Map map = (Map) result; List mapEntries = new ArrayList<>(); for (Map.Entry entry : map.entrySet()) { - CelMutableExpr adaptedKey = maybeAdaptEvaluatedResult(entry.getKey()).orElse(null); + CelMutableExpr adaptedKey = + maybeAdaptEvaluatedResult(typeProvider, valueProvider, entry.getKey()).orElse(null); if (adaptedKey == null) { return Optional.empty(); } - CelMutableExpr adaptedValue = maybeAdaptEvaluatedResult(entry.getValue()).orElse(null); + CelMutableExpr adaptedValue = + maybeAdaptEvaluatedResult(typeProvider, valueProvider, entry.getValue()).orElse(null); if (adaptedValue == null) { return Optional.empty(); } @@ -364,6 +393,31 @@ private Optional maybeAdaptEvaluatedResult(Object result) { CelMutableExpr.ofConstant(CelConstant.ofValue(timestampStrArg))); return Optional.of(CelMutableExpr.ofCall(timestampCall)); + } else if (result instanceof StructValue) { + @SuppressWarnings("unchecked") // Unchecked: StructValue only supports String keys. + StructValue structValue = (StructValue) result; + List structEntries = new ArrayList<>(); + + String typeName = structValue.celType().name(); + CelType optType = typeProvider.findType(typeName).orElse(null); + if (!(optType instanceof StructType)) { + return Optional.empty(); + } + StructType structType = (StructType) optType; + for (String fieldName : structType.fieldNames()) { + Optional fieldOpt = structValue.find(fieldName); + if (!fieldOpt.isPresent()) { + continue; + } + CelMutableExpr adaptedFieldExpr = + maybeAdaptEvaluatedResult(typeProvider, valueProvider, fieldOpt.get()).orElse(null); + if (adaptedFieldExpr == null) { + return Optional.empty(); + } + structEntries.add(CelMutableStruct.Entry.create(0, fieldName, adaptedFieldExpr)); + } + return Optional.of( + CelMutableExpr.ofStruct(CelMutableStruct.create(structType.name(), structEntries))); } // Evaluated result cannot be folded (e.g: unknowns) @@ -371,7 +425,11 @@ private Optional maybeAdaptEvaluatedResult(Object result) { } private Optional maybeRewriteOptional( - Optional optResult, CelMutableAst mutableAst, CelMutableExpr expr) { + CelTypeProvider typeProvider, + CelValueProvider valueProvider, + Optional optResult, + CelMutableAst mutableAst, + CelMutableExpr expr) { Object unwrappedResult = optResult.orElse(null); if (unwrappedResult == null) { if (isCallToFunction(expr, Function.OPTIONAL_NONE.getFunction())) { @@ -387,7 +445,8 @@ private Optional maybeRewriteOptional( return Optional.empty(); } - CelMutableExpr adaptedResult = maybeAdaptEvaluatedResult(unwrappedResult).orElse(null); + CelMutableExpr adaptedResult = + maybeAdaptEvaluatedResult(typeProvider, valueProvider, unwrappedResult).orElse(null); if (adaptedResult == null) { // Evaluated result is not an adaptable constant. Leave the optional as is. return Optional.empty(); diff --git a/optimizer/src/test/java/dev/cel/optimizer/optimizers/BUILD.bazel b/optimizer/src/test/java/dev/cel/optimizer/optimizers/BUILD.bazel index 53d72de67..c912d9570 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/optimizers/BUILD.bazel +++ b/optimizer/src/test/java/dev/cel/optimizer/optimizers/BUILD.bazel @@ -42,8 +42,10 @@ java_library( "@maven//:junit_junit", "@maven//:com_google_testparameterinjector_test_parameter_injector", "//:java_truth", + "@cel_spec//proto/cel/expr/conformance/proto2:test_all_types_java_proto", "@cel_spec//proto/cel/expr/conformance/proto3:test_all_types_java_proto", "@maven//:com_google_guava_guava", + "@maven//:com_google_protobuf_protobuf_java", ], ) diff --git a/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java b/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java index ec4ffd6bc..367ba1724 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java @@ -18,6 +18,8 @@ import static org.junit.Assert.assertThrows; import com.google.common.collect.ImmutableList; +import com.google.protobuf.Duration; +import com.google.protobuf.Timestamp; import com.google.testing.junit.testparameterinjector.TestParameter; import com.google.testing.junit.testparameterinjector.TestParameterInjector; import com.google.testing.junit.testparameterinjector.TestParameters; @@ -31,6 +33,9 @@ import dev.cel.common.types.ListType; import dev.cel.common.types.MapType; import dev.cel.common.types.SimpleType; +import dev.cel.common.types.StructTypeReference; +import dev.cel.expr.conformance.proto2.TestAllTypes.NestedMessage; +import dev.cel.expr.conformance.proto2.TestAllTypesExtensions; import dev.cel.expr.conformance.proto3.TestAllTypes; import dev.cel.extensions.CelExtensions; import dev.cel.extensions.CelOptionalLibrary; @@ -91,6 +96,7 @@ private static Cel setupEnv(CelBuilder celBuilder) { .addFunctionBindings( CelFunctionBinding.from("get_true_overload", ImmutableList.of(), unused -> true)) .addMessageTypes(TestAllTypes.getDescriptor()) + .addFileTypes(TestAllTypesExtensions.getDescriptor()) .setContainer(CelContainer.ofName("cel.expr.conformance.proto3")) .setOptions(CEL_OPTIONS) .addCompilerLibraries( @@ -196,6 +202,8 @@ private static Cel setupEnv(CelBuilder celBuilder) { @TestParameters( "{source: 'TestAllTypes{single_nested_message: TestAllTypes.NestedMessage{bb:" + " 42}}.single_nested_message.bb', expected: '42'}") + @TestParameters("{source: 'TestAllTypes{single_int64: 1 + 2 + 3}.single_int64', expected: '6'}") + @TestParameters("{source: 'TestAllTypes{single_int64: 3}.single_int64', expected: '3'}") @TestParameters("{source: '{\"a\": 1}[\"a\"]', expected: '1'}") @TestParameters("{source: '{\"a\": {\"b\": 2}}[\"a\"][\"b\"]', expected: '2'}") @TestParameters("{source: '{\"hello\": \"world\"}.hello == x', expected: '\"world\" == x'}") @@ -301,6 +309,52 @@ public void constantFold_success(String source, String expected) throws Exceptio assertThat(CEL_UNPARSER.unparse(optimizedAst)).isEqualTo(expected); } + @Test + @TestParameters( + "{source: 'TestAllTypes{single_int32: 3}', " + + " expected: 'cel.expr.conformance.proto3.TestAllTypes{single_int32: 3}'}") + @TestParameters( + "{source: 'TestAllTypes{single_float: 1.5}', " + + " expected: 'cel.expr.conformance.proto3.TestAllTypes{single_float: 1.5}'}") + @TestParameters( + "{source: 'TestAllTypes{single_nested_message: TestAllTypes.NestedMessage{bb: 42}}', " + + " expected: 'cel.expr.conformance.proto3.TestAllTypes{single_nested_message:" + + " cel.expr.conformance.proto3.TestAllTypes.NestedMessage{bb: 42}}'}") + @TestParameters( + "{source: 'TestAllTypes{repeated_int32: [1, 2, 3]}', " + + " expected: 'cel.expr.conformance.proto3.TestAllTypes{repeated_int32: [1, 2, 3]}'}") + @TestParameters( + "{source: 'TestAllTypes{map_int32_int64: {1: 2}}', " + + " expected: 'cel.expr.conformance.proto3.TestAllTypes{map_int32_int64: {1: 2}}'}") + @TestParameters( + "{source: 'TestAllTypes{single_any: google.protobuf.Any{type_url:" + + " \"type.googleapis.com/google.protobuf.Int32Value\", value: b\"\\010\\001\"}}', " + + " expected: 'cel.expr.conformance.proto3.TestAllTypes{single_any:" + + " google.protobuf.Any{type_url:" + + " \"type.googleapis.com/google.protobuf.Int32Value\", value: b\"\\010\\001\"}}'}") + @TestParameters( + "{source: '[TestAllTypes{single_int32: 42}][0]', " + + " expected: 'cel.expr.conformance.proto3.TestAllTypes{single_int32: 42}'}") + @TestParameters( + "{source: 'TestAllTypes{single_bytes: b\"\\010\\001\"}', " + + " expected: 'cel.expr.conformance.proto3.TestAllTypes{single_bytes: b\"\\010\\001\"}'}") + @TestParameters( + "{source: 'TestAllTypes{standalone_enum: 1}', " + + " expected: 'cel.expr.conformance.proto3.TestAllTypes{standalone_enum: 1}'}") + public void constantFold_protoMessageLiteral_success(String source, String expected) + throws Exception { + // Legacy runtime does not support adapting protobuf messages into CelValue (via + // CelValueProvider). + if (runtimeFlavor.equals(CelRuntimeFlavor.LEGACY)) { + return; + } + CelAbstractSyntaxTree ast = cel.compile(source).getAst(); + + CelAbstractSyntaxTree optimizedAst = celOptimizer.optimize(ast); + + assertThat(CEL_UNPARSER.unparse(optimizedAst)).isEqualTo(expected); + } + @Test @TestParameters("{source: '[1 + 1, 1 + 2].exists(i, i < 10)', expected: 'true'}") @TestParameters("{source: '[1, 1 + 1, 1 + 2, 2 + 3].exists(i, i < 10)', expected: 'true'}") @@ -467,6 +521,204 @@ public void constantFold_addFoldableFunction_success() throws Exception { assertThat(CEL_UNPARSER.unparse(optimizedAst)).isEqualTo("true"); } + @Test + public void constantFold_protoMessage_success() throws Exception { + // Legacy runtime does not support adapting protobuf messages into CelValue (via + // CelValueProvider). + if (runtimeFlavor.equals(CelRuntimeFlavor.LEGACY)) { + return; + } + Cel customCel = + cel.toCelBuilder() + .addFunctionDeclarations( + CelFunctionDecl.newFunctionDeclaration( + "get_test_all_types", + CelOverloadDecl.newGlobalOverload( + "get_test_all_types_overload", + StructTypeReference.create("cel.expr.conformance.proto3.TestAllTypes")))) + .addFunctionBindings( + CelFunctionBinding.from( + "get_test_all_types_overload", + ImmutableList.of(), + unused -> + TestAllTypes.newBuilder() + .setSingleInt32(1) + .setSingleString("hello") + .build())) + .build(); + CelAbstractSyntaxTree ast = customCel.compile("get_test_all_types()").getAst(); + ConstantFoldingOptions options = + ConstantFoldingOptions.newBuilder().addFoldableFunctions("get_test_all_types").build(); + CelOptimizer optimizer = + CelOptimizerFactory.standardCelOptimizerBuilder(customCel) + .addAstOptimizers(ConstantFoldingOptimizer.newInstance(options)) + .build(); + + CelAbstractSyntaxTree optimizedAst = optimizer.optimize(ast); + + assertThat(CEL_UNPARSER.unparse(optimizedAst)) + .isEqualTo( + "cel.expr.conformance.proto3.TestAllTypes{single_int32: 1, single_string: \"hello\"}"); + } + + @Test + public void constantFold_protoMessage_complexFields_success() throws Exception { + // Legacy runtime does not support adapting protobuf messages into CelValue (via + // CelValueProvider). + if (runtimeFlavor.equals(CelRuntimeFlavor.LEGACY)) { + return; + } + Cel customCel = + cel.toCelBuilder() + .addFunctionDeclarations( + CelFunctionDecl.newFunctionDeclaration( + "get_test_all_types_complex", + CelOverloadDecl.newGlobalOverload( + "get_test_all_types_complex_overload", + StructTypeReference.create("cel.expr.conformance.proto3.TestAllTypes")))) + .addFunctionBindings( + CelFunctionBinding.from( + "get_test_all_types_complex_overload", + ImmutableList.of(), + unused -> + TestAllTypes.newBuilder() + .setSingleUint32(123) + .setSingleUint64(456L) + .setSingleDuration(Duration.newBuilder().setSeconds(10).build()) + .setSingleTimestamp(Timestamp.newBuilder().setSeconds(10).build()) + .addRepeatedNestedMessage( + TestAllTypes.NestedMessage.newBuilder().setBb(99).build()) + .build())) + .build(); + CelAbstractSyntaxTree ast = customCel.compile("get_test_all_types_complex()").getAst(); + ConstantFoldingOptions options = + ConstantFoldingOptions.newBuilder() + .addFoldableFunctions("get_test_all_types_complex") + .build(); + CelOptimizer optimizer = + CelOptimizerFactory.standardCelOptimizerBuilder(customCel) + .addAstOptimizers(ConstantFoldingOptimizer.newInstance(options)) + .build(); + + CelAbstractSyntaxTree optimizedAst = optimizer.optimize(ast); + + assertThat(CEL_UNPARSER.unparse(optimizedAst)).contains("123u"); + } + + @Test + public void constantFold_proto2Message_success() throws Exception { + // Legacy runtime does not support adapting protobuf messages into CelValue (via + // CelValueProvider). + if (runtimeFlavor.equals(CelRuntimeFlavor.LEGACY)) { + return; + } + Cel customCel = + cel.toCelBuilder() + .addFunctionDeclarations( + CelFunctionDecl.newFunctionDeclaration( + "get_test_all_types_proto2", + CelOverloadDecl.newGlobalOverload( + "get_test_all_types_proto2_overload", + StructTypeReference.create("cel.expr.conformance.proto2.TestAllTypes")))) + .addFunctionBindings( + CelFunctionBinding.from( + "get_test_all_types_proto2_overload", + ImmutableList.of(), + unused -> + dev.cel.expr.conformance.proto2.TestAllTypes.newBuilder() + .setSingleInt32(2) + .setExtension(TestAllTypesExtensions.int32Ext, 3) + .build())) + .build(); + CelAbstractSyntaxTree ast = customCel.compile("get_test_all_types_proto2()").getAst(); + ConstantFoldingOptions options = + ConstantFoldingOptions.newBuilder() + .addFoldableFunctions("get_test_all_types_proto2") + .build(); + CelOptimizer optimizer = + CelOptimizerFactory.standardCelOptimizerBuilder(customCel) + .addAstOptimizers(ConstantFoldingOptimizer.newInstance(options)) + .build(); + + CelAbstractSyntaxTree optimizedAst = optimizer.optimize(ast); + + assertThat(CEL_UNPARSER.unparse(optimizedAst)) + .isEqualTo("cel.expr.conformance.proto2.TestAllTypes{single_int32: 2}"); + } + + @Test + public void constantFold_proto2Message_complexFields_success() throws Exception { + // Legacy runtime does not support adapting protobuf messages into CelValue (via + // CelValueProvider). + if (runtimeFlavor.equals(CelRuntimeFlavor.LEGACY)) { + return; + } + Cel customCel = + cel.toCelBuilder() + .addFunctionDeclarations( + CelFunctionDecl.newFunctionDeclaration( + "get_test_all_types_proto2_complex", + CelOverloadDecl.newGlobalOverload( + "get_test_all_types_proto2_complex_overload", + StructTypeReference.create("cel.expr.conformance.proto2.TestAllTypes")))) + .addFunctionBindings( + CelFunctionBinding.from( + "get_test_all_types_proto2_complex_overload", + ImmutableList.of(), + unused -> + dev.cel.expr.conformance.proto2.TestAllTypes.newBuilder() + .setSingleUint32(123) + .setSingleUint64(456L) + .addRepeatedNestedMessage(NestedMessage.newBuilder().setBb(99).build()) + .build())) + .build(); + CelAbstractSyntaxTree ast = customCel.compile("get_test_all_types_proto2_complex()").getAst(); + ConstantFoldingOptions options = + ConstantFoldingOptions.newBuilder() + .addFoldableFunctions("get_test_all_types_proto2_complex") + .build(); + CelOptimizer optimizer = + CelOptimizerFactory.standardCelOptimizerBuilder(customCel) + .addAstOptimizers(ConstantFoldingOptimizer.newInstance(options)) + .build(); + + CelAbstractSyntaxTree optimizedAst = optimizer.optimize(ast); + + assertThat(CEL_UNPARSER.unparse(optimizedAst)).contains("123u"); + } + + @Test + public void constantFold_functionReturningUnregisteredMessage_doesNotFold() throws Exception { + Cel customCel = + runtimeFlavor + .builder() + .addVar("x", SimpleType.DYN) + .addFunctionDeclarations( + CelFunctionDecl.newFunctionDeclaration( + "get_unregistered_message", + CelOverloadDecl.newGlobalOverload( + "get_unregistered_message_overload", SimpleType.ANY))) + .addFunctionBindings( + CelFunctionBinding.from( + "get_unregistered_message_overload", + ImmutableList.of(), + unused -> TestAllTypes.getDefaultInstance())) + .build(); + ConstantFoldingOptions options = + ConstantFoldingOptions.newBuilder() + .addFoldableFunctions("get_unregistered_message") + .build(); + CelOptimizer optimizer = + CelOptimizerFactory.standardCelOptimizerBuilder(customCel) + .addAstOptimizers(ConstantFoldingOptimizer.newInstance(options)) + .build(); + CelAbstractSyntaxTree ast = customCel.compile("get_unregistered_message()").getAst(); + + CelAbstractSyntaxTree optimizedAst = optimizer.optimize(ast); + + assertThat(CEL_UNPARSER.unparse(optimizedAst)).isEqualTo("get_unregistered_message()"); + } + @Test public void constantFold_withExpectedResultTypeSet_success() throws Exception { Cel cel = runtimeFlavor.builder().setResultType(SimpleType.STRING).build(); diff --git a/runtime/src/main/java/dev/cel/runtime/CelRuntimeBuilder.java b/runtime/src/main/java/dev/cel/runtime/CelRuntimeBuilder.java index 87f11fde2..e284b374c 100644 --- a/runtime/src/main/java/dev/cel/runtime/CelRuntimeBuilder.java +++ b/runtime/src/main/java/dev/cel/runtime/CelRuntimeBuilder.java @@ -167,6 +167,9 @@ public interface CelRuntimeBuilder { @CanIgnoreReturnValue CelRuntimeBuilder setValueProvider(CelValueProvider celValueProvider); + /** Returns the configured {@link CelValueProvider}, or null if not set. */ + CelValueProvider valueProvider(); + /** Enable or disable the standard CEL library functions and variables. */ @CanIgnoreReturnValue CelRuntimeBuilder setStandardEnvironmentEnabled(boolean value); diff --git a/runtime/src/main/java/dev/cel/runtime/CelRuntimeImpl.java b/runtime/src/main/java/dev/cel/runtime/CelRuntimeImpl.java index b02f64b61..f934108e0 100644 --- a/runtime/src/main/java/dev/cel/runtime/CelRuntimeImpl.java +++ b/runtime/src/main/java/dev/cel/runtime/CelRuntimeImpl.java @@ -286,7 +286,8 @@ public abstract static class Builder implements CelRuntimeBuilder { abstract CelTypeProvider typeProvider(); - abstract CelValueProvider valueProvider(); + @Override + public abstract CelValueProvider valueProvider(); abstract CelStandardFunctions standardFunctions(); @@ -503,6 +504,7 @@ public CelRuntime build() { protoMessageValueProvider = CombinedCelValueProvider.combine(protoMessageValueProvider, valueProvider()); } + setValueProvider(protoMessageValueProvider); CelValueConverter celValueConverter = protoMessageValueProvider.celValueConverter(); CelTypeProvider messageTypeProvider = diff --git a/runtime/src/main/java/dev/cel/runtime/CelRuntimeLegacyImpl.java b/runtime/src/main/java/dev/cel/runtime/CelRuntimeLegacyImpl.java index c5e06d013..cad7e74f8 100644 --- a/runtime/src/main/java/dev/cel/runtime/CelRuntimeLegacyImpl.java +++ b/runtime/src/main/java/dev/cel/runtime/CelRuntimeLegacyImpl.java @@ -207,6 +207,11 @@ public CelRuntimeBuilder setValueProvider(CelValueProvider celValueProvider) { "setValueProvider is not supported for legacy runtime"); } + @Override + public CelValueProvider valueProvider() { + throw new UnsupportedOperationException("valueProvider is not supported for legacy runtime"); + } + @Override public CelRuntimeBuilder setTypeFactory(Function typeFactory) { this.customTypeFactory = typeFactory;