From 4ba19e392152956558c96f2d67be2ea140eb8949 Mon Sep 17 00:00:00 2001 From: Dmitri Plotnikov Date: Fri, 2 Oct 2026 15:41:39 -0700 Subject: [PATCH] Split CelComprehensionsExtensions into CelComprehensionsCompilerLibrary and CelComprehensionsRuntimeLibrary to allow its usage in Lite runtime. PiperOrigin-RevId: 992578130 --- extensions/BUILD.bazel | 17 + .../main/java/dev/cel/extensions/BUILD.bazel | 66 ++- .../CelComprehensionsCompilerLibrary.java | 505 ++++++++++++++++++ .../CelComprehensionsExtensions.java | 489 ++--------------- .../CelComprehensionsRuntimeLibrary.java | 175 ++++++ .../dev/cel/extensions/CelExtensions.java | 36 ++ .../test/java/dev/cel/extensions/BUILD.bazel | 5 + .../CelComprehensionsExtensionsTest.java | 188 +++++++ .../src/main/java/dev/cel/runtime/BUILD.bazel | 21 +- .../CelInternalLiteRuntimeLibrary.java | 35 ++ .../cel/runtime/CelLiteRuntimeFactory.java | 2 +- ...ntimeImpl.java => CelLiteRuntimeImpl.java} | 19 +- .../src/test/java/dev/cel/runtime/BUILD.bazel | 1 + .../runtime/CelLiteRuntimeAndroidTest.java | 18 +- .../java/dev/cel/testing/compiled/BUILD.bazel | 7 + .../resources/environment/all_extensions.yaml | 1 + 16 files changed, 1128 insertions(+), 457 deletions(-) create mode 100644 extensions/src/main/java/dev/cel/extensions/CelComprehensionsCompilerLibrary.java create mode 100644 extensions/src/main/java/dev/cel/extensions/CelComprehensionsRuntimeLibrary.java create mode 100644 runtime/src/main/java/dev/cel/runtime/CelInternalLiteRuntimeLibrary.java rename runtime/src/main/java/dev/cel/runtime/{LiteRuntimeImpl.java => CelLiteRuntimeImpl.java} (96%) diff --git a/extensions/BUILD.bazel b/extensions/BUILD.bazel index cfe6e8fca..a494fe82e 100644 --- a/extensions/BUILD.bazel +++ b/extensions/BUILD.bazel @@ -92,6 +92,23 @@ java_library( exports = ["//extensions/src/main/java/dev/cel/extensions:comprehensions"], ) +java_library( + name = "comprehensions_compiler_library", + visibility = ["//:internal"], + exports = ["//extensions/src/main/java/dev/cel/extensions:comprehensions_compiler_library"], +) + +java_library( + name = "comprehensions_runtime_library", + visibility = ["//:internal"], + exports = ["//extensions/src/main/java/dev/cel/extensions:comprehensions_runtime_library"], +) + +cel_android_library( + name = "comprehensions_runtime_library_android", + exports = ["//extensions/src/main/java/dev/cel/extensions:comprehensions_runtime_library_android"], +) + java_library( name = "native", exports = ["//extensions/src/main/java/dev/cel/extensions:native"], diff --git a/extensions/src/main/java/dev/cel/extensions/BUILD.bazel b/extensions/src/main/java/dev/cel/extensions/BUILD.bazel index d664397c0..d9e5c556c 100644 --- a/extensions/src/main/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/main/java/dev/cel/extensions/BUILD.bazel @@ -546,25 +546,81 @@ cel_android_library( java_library( name = "comprehensions", srcs = ["CelComprehensionsExtensions.java"], + tags = [ + ], deps = [ + ":comprehensions_compiler_library", + ":comprehensions_runtime_library", + ":extension_library", "//checker:checker_builder", - "//common:compiler_common", - "//common:operator", + "//common:cel_function_decl", "//common:options", + "//compiler:compiler_builder", + "//parser:macro", + "//parser:parser_builder", + "//runtime", + "//runtime:runtime_equality", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven//:com_google_guava_guava", + ], +) + +java_library( + name = "comprehensions_compiler_library", + srcs = ["CelComprehensionsCompilerLibrary.java"], + tags = [ + "alt_dep=//extensions:comprehensions_compiler_library", + ], + deps = [ + ":extension_library", + "//checker:checker_builder", + "//common:cel_function_decl", + "//common:cel_issue", + "//common:cel_overload_decl", + "//common:operator", "//common/ast", "//common/types", - "//common/values:mutable_map_value", "//compiler:compiler_builder", - "//extensions:extension_library", "//parser:macro", "//parser:parser_builder", - "//runtime", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven//:com_google_guava_guava", + ], +) + +java_library( + name = "comprehensions_runtime_library", + srcs = ["CelComprehensionsRuntimeLibrary.java"], + tags = [ + "alt_dep=//extensions:comprehensions_runtime_library", + ], + deps = [ + "//common:options", + "//common/values:mutable_map_value", "//runtime:function_binding", + "//runtime:lite_runtime", "//runtime:runtime_equality", + "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", ], ) +cel_android_library( + name = "comprehensions_runtime_library_android", + srcs = ["CelComprehensionsRuntimeLibrary.java"], + tags = [ + ], + deps = [ + "//common:options", + "//common/values:mutable_map_value_android", + "//runtime:function_binding_android", + "//runtime:lite_runtime_android", + "//runtime:runtime_equality_android", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven_android//:com_google_guava_guava", + ], +) + java_library( name = "native", srcs = ["CelNativeTypesExtensions.java"], diff --git a/extensions/src/main/java/dev/cel/extensions/CelComprehensionsCompilerLibrary.java b/extensions/src/main/java/dev/cel/extensions/CelComprehensionsCompilerLibrary.java new file mode 100644 index 000000000..2f33f6d7f --- /dev/null +++ b/extensions/src/main/java/dev/cel/extensions/CelComprehensionsCompilerLibrary.java @@ -0,0 +1,505 @@ +// Copyright 2025 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.extensions; + +import static com.google.common.base.Preconditions.checkArgument; +import static com.google.common.base.Preconditions.checkNotNull; + +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableSet; +import com.google.errorprone.annotations.Immutable; +import dev.cel.checker.CelCheckerBuilder; +import dev.cel.common.CelFunctionDecl; +import dev.cel.common.CelIssue; +import dev.cel.common.CelOverloadDecl; +import dev.cel.common.Operator; +import dev.cel.common.ast.CelExpr; +import dev.cel.common.types.MapType; +import dev.cel.common.types.TypeParamType; +import dev.cel.compiler.CelCompilerLibrary; +import dev.cel.parser.CelMacro; +import dev.cel.parser.CelMacroExprFactory; +import dev.cel.parser.CelParserBuilder; +import java.util.Arrays; +import java.util.Collection; +import java.util.Optional; + +/** Internal implementation of CEL two variable comprehensions compile-time extensions. */ +@Immutable +public final class CelComprehensionsCompilerLibrary + implements CelCompilerLibrary, CelExtensionLibrary.FeatureSet { + + private static final String MAP_INSERT_FUNCTION = "cel.@mapInsert"; + private static final String MAP_INSERT_OVERLOAD_MAP_MAP = "cel_@mapInsert_map_map"; + private static final String MAP_INSERT_OVERLOAD_KEY_VALUE = "cel_@mapInsert_map_key_value"; + + private static final class Types { + private static final TypeParamType TYPE_PARAM_K = TypeParamType.create("K"); + private static final TypeParamType TYPE_PARAM_V = TypeParamType.create("V"); + private static final MapType MAP_KV_TYPE = MapType.create(TYPE_PARAM_K, TYPE_PARAM_V); + } + + /** Enumeration of functions for Comprehensions compiler extension. */ + public enum Function { + MAP_INSERT( + CelFunctionDecl.newFunctionDeclaration( + MAP_INSERT_FUNCTION, + CelOverloadDecl.newGlobalOverload( + MAP_INSERT_OVERLOAD_MAP_MAP, + "Returns a map that's the result of merging given two maps.", + Types.MAP_KV_TYPE, + Types.MAP_KV_TYPE, + Types.MAP_KV_TYPE), + CelOverloadDecl.newGlobalOverload( + MAP_INSERT_OVERLOAD_KEY_VALUE, + "Adds the given key-value pair to the map.", + Types.MAP_KV_TYPE, + Types.MAP_KV_TYPE, + Types.TYPE_PARAM_K, + Types.TYPE_PARAM_V))); + + private final CelFunctionDecl functionDecl; + + public CelFunctionDecl functionDecl() { + return functionDecl; + } + + public CelFunctionDecl getFunctionDecl() { + return functionDecl; + } + + public String getFunction() { + return functionDecl.name(); + } + + Function(CelFunctionDecl functionDecl) { + this.functionDecl = functionDecl; + } + } + + private static ImmutableSet getFunctionsForVersion(int version) { + switch (version) { + case 0: + case Integer.MAX_VALUE: + return ImmutableSet.of(Function.MAP_INSERT); + default: + throw new IllegalArgumentException( + "Unsupported 'comprehensions' extension version " + version); + } + } + + /** Returns the latest version of the 'comprehensions' compile-time extensions. */ + public static CelComprehensionsCompilerLibrary comprehensions() { + return comprehensions(Integer.MAX_VALUE); + } + + /** Returns the specified version of the 'comprehensions' compile-time extensions. */ + public static CelComprehensionsCompilerLibrary comprehensions(int version) { + return new CelComprehensionsCompilerLibrary(version); + } + + /** Returns the 'comprehensions' compile-time extensions with only the specified functions. */ + public static CelComprehensionsCompilerLibrary comprehensions(Function... functions) { + return comprehensions(Arrays.asList(functions)); + } + + /** Returns the 'comprehensions' compile-time extensions with only the specified functions. */ + public static CelComprehensionsCompilerLibrary comprehensions(Collection functions) { + return new CelComprehensionsCompilerLibrary(functions); + } + + private static final class Library + implements CelExtensionLibrary { + private final CelComprehensionsCompilerLibrary version0; + + Library() { + version0 = new CelComprehensionsCompilerLibrary(0); + } + + @Override + public String name() { + return "comprehensions"; + } + + @Override + public ImmutableSet versions() { + return ImmutableSet.of(version0); + } + } + + private static final Library LIBRARY = new Library(); + + public static CelExtensionLibrary library() { + return LIBRARY; + } + + private final int version; + private final ImmutableSet functions; + + CelComprehensionsCompilerLibrary(int version) { + this(version, getFunctionsForVersion(version)); + } + + CelComprehensionsCompilerLibrary(Collection functions) { + this(-1, functions); + } + + private CelComprehensionsCompilerLibrary(int version, Collection functions) { + this.version = version; + this.functions = ImmutableSet.copyOf(functions); + } + + @Override + public void setCheckerOptions(CelCheckerBuilder checkerBuilder) { + functions.forEach(function -> checkerBuilder.addFunctionDeclarations(function.functionDecl())); + } + + @Override + public int version() { + return version; + } + + @Override + public ImmutableSet macros() { + return ImmutableSet.of( + CelMacro.newReceiverMacro( + Operator.ALL.getFunction(), 3, CelComprehensionsCompilerLibrary::expandAllMacro), + CelMacro.newReceiverMacro( + Operator.EXISTS.getFunction(), 3, CelComprehensionsCompilerLibrary::expandExistsMacro), + CelMacro.newReceiverMacro( + Operator.EXISTS_ONE.getFunction(), + 3, + CelComprehensionsCompilerLibrary::expandExistsOneMacro), + CelMacro.newReceiverMacro( + Operator.EXISTS_ONE_NEW.getFunction(), + 3, + CelComprehensionsCompilerLibrary::expandExistsOneMacro), + CelMacro.newReceiverMacro( + "transformList", 3, CelComprehensionsCompilerLibrary::transformListMacro), + CelMacro.newReceiverMacro( + "transformList", 4, CelComprehensionsCompilerLibrary::transformListMacro), + CelMacro.newReceiverMacro( + "transformMap", 3, CelComprehensionsCompilerLibrary::transformMapMacro), + CelMacro.newReceiverMacro( + "transformMap", 4, CelComprehensionsCompilerLibrary::transformMapMacro), + CelMacro.newReceiverMacro( + "transformMapEntry", 3, CelComprehensionsCompilerLibrary::transformMapEntryMacro), + CelMacro.newReceiverMacro( + "transformMapEntry", 4, CelComprehensionsCompilerLibrary::transformMapEntryMacro)); + } + + @Override + public void setParserOptions(CelParserBuilder parserBuilder) { + parserBuilder.addMacros(macros()); + } + + private static Optional expandAllMacro( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { + checkNotNull(exprFactory); + checkNotNull(target); + checkArgument(arguments.size() == 3); + CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); + if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg0); + } + CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); + if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg1); + } + CelExpr arg2 = checkNotNull(arguments.get(2)); + CelExpr accuInit = exprFactory.newBoolLiteral(true); + CelExpr condition = + exprFactory.newGlobalCall( + Operator.NOT_STRICTLY_FALSE.getFunction(), + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); + CelExpr step = + exprFactory.newGlobalCall( + Operator.LOGICAL_AND.getFunction(), + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), + arg2); + CelExpr result = exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()); + return Optional.of( + exprFactory.fold( + arg0.ident().name(), + arg1.ident().name(), + target, + exprFactory.getAccumulatorVarName(), + accuInit, + condition, + step, + result)); + } + + private static Optional expandExistsMacro( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { + checkNotNull(exprFactory); + checkNotNull(target); + checkArgument(arguments.size() == 3); + CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); + if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg0); + } + CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); + if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg1); + } + CelExpr arg2 = checkNotNull(arguments.get(2)); + CelExpr accuInit = exprFactory.newBoolLiteral(false); + CelExpr condition = + exprFactory.newGlobalCall( + Operator.NOT_STRICTLY_FALSE.getFunction(), + exprFactory.newGlobalCall( + Operator.LOGICAL_NOT.getFunction(), + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()))); + CelExpr step = + exprFactory.newGlobalCall( + Operator.LOGICAL_OR.getFunction(), + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), + arg2); + CelExpr result = exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()); + return Optional.of( + exprFactory.fold( + arg0.ident().name(), + arg1.ident().name(), + target, + exprFactory.getAccumulatorVarName(), + accuInit, + condition, + step, + result)); + } + + private static Optional expandExistsOneMacro( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { + checkNotNull(exprFactory); + checkNotNull(target); + checkArgument(arguments.size() == 3); + CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); + if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg0); + } + CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); + if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg1); + } + CelExpr arg2 = checkNotNull(arguments.get(2)); + CelExpr accuInit = exprFactory.newIntLiteral(0); + CelExpr condition = exprFactory.newBoolLiteral(true); + CelExpr step = + exprFactory.newGlobalCall( + Operator.CONDITIONAL.getFunction(), + arg2, + exprFactory.newGlobalCall( + Operator.ADD.getFunction(), + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), + exprFactory.newIntLiteral(1)), + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); + CelExpr result = + exprFactory.newGlobalCall( + Operator.EQUALS.getFunction(), + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), + exprFactory.newIntLiteral(1)); + return Optional.of( + exprFactory.fold( + arg0.ident().name(), + arg1.ident().name(), + target, + exprFactory.getAccumulatorVarName(), + accuInit, + condition, + step, + result)); + } + + private static Optional transformListMacro( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { + checkNotNull(exprFactory); + checkNotNull(target); + checkArgument(arguments.size() == 3 || arguments.size() == 4); + CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); + if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg0); + } + CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); + if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg1); + } + CelExpr transform; + CelExpr filter = null; + if (arguments.size() == 4) { + filter = checkNotNull(arguments.get(2)); + transform = checkNotNull(arguments.get(3)); + } else { + transform = checkNotNull(arguments.get(2)); + } + CelExpr accuInit = exprFactory.newList(); + CelExpr condition = exprFactory.newBoolLiteral(true); + CelExpr step = + exprFactory.newGlobalCall( + Operator.ADD.getFunction(), + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), + exprFactory.newList(transform)); + if (filter != null) { + step = + exprFactory.newGlobalCall( + Operator.CONDITIONAL.getFunction(), + filter, + step, + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); + } + return Optional.of( + exprFactory.fold( + arg0.ident().name(), + arg1.ident().name(), + target, + exprFactory.getAccumulatorVarName(), + accuInit, + condition, + step, + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()))); + } + + private static Optional transformMapMacro( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { + checkNotNull(exprFactory); + checkNotNull(target); + checkArgument(arguments.size() == 3 || arguments.size() == 4); + CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); + if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg0); + } + CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); + if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg1); + } + CelExpr transform; + CelExpr filter = null; + if (arguments.size() == 4) { + filter = checkNotNull(arguments.get(2)); + transform = checkNotNull(arguments.get(3)); + } else { + transform = checkNotNull(arguments.get(2)); + } + CelExpr accuInit = exprFactory.newMap(); + CelExpr condition = exprFactory.newBoolLiteral(true); + CelExpr step = + exprFactory.newGlobalCall( + MAP_INSERT_FUNCTION, + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), + arg0, + transform); + if (filter != null) { + step = + exprFactory.newGlobalCall( + Operator.CONDITIONAL.getFunction(), + filter, + step, + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); + } + return Optional.of( + exprFactory.fold( + arg0.ident().name(), + arg1.ident().name(), + target, + exprFactory.getAccumulatorVarName(), + accuInit, + condition, + step, + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()))); + } + + private static Optional transformMapEntryMacro( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { + checkNotNull(exprFactory); + checkNotNull(target); + checkArgument(arguments.size() == 3 || arguments.size() == 4); + CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); + if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg0); + } + CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); + if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of(arg1); + } + CelExpr transform; + CelExpr filter = null; + if (arguments.size() == 4) { + filter = checkNotNull(arguments.get(2)); + transform = checkNotNull(arguments.get(3)); + } else { + transform = checkNotNull(arguments.get(2)); + } + CelExpr accuInit = exprFactory.newMap(); + CelExpr condition = exprFactory.newBoolLiteral(true); + CelExpr step = + exprFactory.newGlobalCall( + MAP_INSERT_FUNCTION, + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), + transform); + if (filter != null) { + step = + exprFactory.newGlobalCall( + Operator.CONDITIONAL.getFunction(), + filter, + step, + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); + } + return Optional.of( + exprFactory.fold( + arg0.ident().name(), + arg1.ident().name(), + target, + exprFactory.getAccumulatorVarName(), + accuInit, + condition, + step, + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()))); + } + + private static CelExpr validatedIterationVariable( + CelMacroExprFactory exprFactory, CelExpr argument) { + CelExpr arg = checkNotNull(argument); + if (!isSimpleIdentifier(arg)) { + return reportArgumentError(exprFactory, arg); + } else if (arg.exprKind().ident().name().equals(exprFactory.getAccumulatorVarName()) + || arg.exprKind().ident().name().equals("__result__")) { + return reportAccumulatorOverwriteError(exprFactory, arg); + } else { + return arg; + } + } + + private static boolean isSimpleIdentifier(CelExpr expr) { + return expr.getKind() == CelExpr.ExprKind.Kind.IDENT + && !expr.ident().name().isEmpty() + && !expr.ident().name().startsWith("."); + } + + private static CelExpr reportArgumentError(CelMacroExprFactory exprFactory, CelExpr argument) { + return exprFactory.reportError( + CelIssue.formatError( + exprFactory.getSourceLocation(argument), "The argument must be a simple name")); + } + + private static CelExpr reportAccumulatorOverwriteError( + CelMacroExprFactory exprFactory, CelExpr argument) { + return exprFactory.reportError( + CelIssue.formatError( + exprFactory.getSourceLocation(argument), + String.format( + "The iteration variable %s overwrites accumulator variable", + argument.ident().name()))); + } +} diff --git a/extensions/src/main/java/dev/cel/extensions/CelComprehensionsExtensions.java b/extensions/src/main/java/dev/cel/extensions/CelComprehensionsExtensions.java index 8a415608d..e94f037b0 100644 --- a/extensions/src/main/java/dev/cel/extensions/CelComprehensionsExtensions.java +++ b/extensions/src/main/java/dev/cel/extensions/CelComprehensionsExtensions.java @@ -14,96 +14,73 @@ package dev.cel.extensions; -import static com.google.common.base.Preconditions.checkArgument; -import static com.google.common.base.Preconditions.checkNotNull; +import static com.google.common.collect.ImmutableSet.toImmutableSet; -import com.google.common.collect.ImmutableList; -import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; +import com.google.errorprone.annotations.Immutable; import dev.cel.checker.CelCheckerBuilder; import dev.cel.common.CelFunctionDecl; -import dev.cel.common.CelIssue; import dev.cel.common.CelOptions; -import dev.cel.common.CelOverloadDecl; -import dev.cel.common.Operator; -import dev.cel.common.ast.CelExpr; -import dev.cel.common.types.MapType; -import dev.cel.common.types.TypeParamType; -import dev.cel.common.values.MutableMapValue; import dev.cel.compiler.CelCompilerLibrary; import dev.cel.parser.CelMacro; -import dev.cel.parser.CelMacroExprFactory; import dev.cel.parser.CelParserBuilder; -import dev.cel.runtime.CelFunctionBinding; import dev.cel.runtime.CelInternalRuntimeLibrary; import dev.cel.runtime.CelRuntimeBuilder; import dev.cel.runtime.RuntimeEquality; -import java.util.Map; -import java.util.Optional; +import java.util.Set; /** Internal implementation of CEL two variable comprehensions extensions. */ +@Immutable public final class CelComprehensionsExtensions implements CelCompilerLibrary, CelInternalRuntimeLibrary, CelExtensionLibrary.FeatureSet { - private static final String MAP_INSERT_FUNCTION = "cel.@mapInsert"; - private static final String MAP_INSERT_OVERLOAD_MAP_MAP = "cel_@mapInsert_map_map"; - private static final String MAP_INSERT_OVERLOAD_KEY_VALUE = "cel_@mapInsert_map_key_value"; - - private static final class Types { - private static final TypeParamType TYPE_PARAM_K = TypeParamType.create("K"); - private static final TypeParamType TYPE_PARAM_V = TypeParamType.create("V"); - private static final MapType MAP_KV_TYPE = MapType.create(TYPE_PARAM_K, TYPE_PARAM_V); - } - /** Enumeration of functions for Comprehensions extension. */ public enum Function { MAP_INSERT( - CelFunctionDecl.newFunctionDeclaration( - MAP_INSERT_FUNCTION, - CelOverloadDecl.newGlobalOverload( - MAP_INSERT_OVERLOAD_MAP_MAP, - "Returns a map that's the result of merging given two maps.", - Types.MAP_KV_TYPE, - Types.MAP_KV_TYPE, - Types.MAP_KV_TYPE), - CelOverloadDecl.newGlobalOverload( - MAP_INSERT_OVERLOAD_KEY_VALUE, - "Adds the given key-value pair to the map.", - Types.MAP_KV_TYPE, - Types.MAP_KV_TYPE, - Types.TYPE_PARAM_K, - Types.TYPE_PARAM_V))); + CelComprehensionsCompilerLibrary.Function.MAP_INSERT, + CelComprehensionsRuntimeLibrary.Function.MAP_INSERT); - private final CelFunctionDecl functionDecl; + private final CelComprehensionsCompilerLibrary.Function compilerFunction; + private final CelComprehensionsRuntimeLibrary.Function runtimeFunction; public CelFunctionDecl functionDecl() { - return functionDecl; + return compilerFunction.functionDecl(); + } + + public CelFunctionDecl getFunctionDecl() { + return compilerFunction.getFunctionDecl(); } - String getFunction() { - return functionDecl.name(); + public String getFunction() { + return compilerFunction.getFunction(); } - Function(CelFunctionDecl functionDecl) { - this.functionDecl = functionDecl; + Function( + CelComprehensionsCompilerLibrary.Function compilerFunction, + CelComprehensionsRuntimeLibrary.Function runtimeFunction) { + this.compilerFunction = compilerFunction; + this.runtimeFunction = runtimeFunction; } } private static final class Library implements CelExtensionLibrary { - private final CelComprehensionsExtensions version0; + private final ImmutableSet versions; Library() { - version0 = new CelComprehensionsExtensions(); + versions = + CelComprehensionsCompilerLibrary.library().versions().stream() + .map(CelComprehensionsExtensions::new) + .collect(toImmutableSet()); } @Override public String name() { - return "comprehensions"; + return CelComprehensionsCompilerLibrary.library().name(); } @Override public ImmutableSet versions() { - return ImmutableSet.of(version0); + return versions; } } @@ -113,415 +90,55 @@ static CelExtensionLibrary library() { return LIBRARY; } - private final ImmutableSet functions; + private final CelComprehensionsCompilerLibrary compilerLibrary; + private final CelComprehensionsRuntimeLibrary runtimeLibrary; CelComprehensionsExtensions() { - this.functions = ImmutableSet.of(Function.MAP_INSERT); + this(CelComprehensionsCompilerLibrary.comprehensions()); } - @Override - public void setCheckerOptions(CelCheckerBuilder checkerBuilder) { - functions.forEach(function -> checkerBuilder.addFunctionDeclarations(function.functionDecl())); + CelComprehensionsExtensions(Set functions) { + this.compilerLibrary = + new CelComprehensionsCompilerLibrary( + functions.stream().map(f -> f.compilerFunction).collect(toImmutableSet())); + this.runtimeLibrary = + new CelComprehensionsRuntimeLibrary( + functions.stream().map(f -> f.runtimeFunction).collect(toImmutableSet())); } - @Override - public void setRuntimeOptions(CelRuntimeBuilder runtimeBuilder) { - throw new UnsupportedOperationException("Unsupported"); - } - - @Override - public void setRuntimeOptions( - CelRuntimeBuilder runtimeBuilder, RuntimeEquality runtimeEquality, CelOptions celOptions) { - runtimeBuilder.addFunctionBindings( - CelFunctionBinding.fromOverloads( - MAP_INSERT_FUNCTION, - CelFunctionBinding.from( - MAP_INSERT_OVERLOAD_MAP_MAP, - Map.class, - Map.class, - (map1, map2) -> mapInsertMap(map1, map2, runtimeEquality)), - CelFunctionBinding.from( - MAP_INSERT_OVERLOAD_KEY_VALUE, - ImmutableList.of(Map.class, Object.class, Object.class), - args -> mapInsertKeyValue(args, runtimeEquality)))); + private CelComprehensionsExtensions(CelComprehensionsCompilerLibrary compilerLibrary) { + this.compilerLibrary = compilerLibrary; + this.runtimeLibrary = CelComprehensionsRuntimeLibrary.comprehensions(compilerLibrary.version()); } @Override public int version() { - return 0; + return compilerLibrary.version(); } @Override public ImmutableSet macros() { - return ImmutableSet.of( - CelMacro.newReceiverMacro( - Operator.ALL.getFunction(), 3, CelComprehensionsExtensions::expandAllMacro), - CelMacro.newReceiverMacro( - Operator.EXISTS.getFunction(), 3, CelComprehensionsExtensions::expandExistsMacro), - CelMacro.newReceiverMacro( - Operator.EXISTS_ONE.getFunction(), - 3, - CelComprehensionsExtensions::expandExistsOneMacro), - CelMacro.newReceiverMacro( - Operator.EXISTS_ONE_NEW.getFunction(), - 3, - CelComprehensionsExtensions::expandExistsOneMacro), - CelMacro.newReceiverMacro( - "transformList", 3, CelComprehensionsExtensions::transformListMacro), - CelMacro.newReceiverMacro( - "transformList", 4, CelComprehensionsExtensions::transformListMacro), - CelMacro.newReceiverMacro( - "transformMap", 3, CelComprehensionsExtensions::transformMapMacro), - CelMacro.newReceiverMacro( - "transformMap", 4, CelComprehensionsExtensions::transformMapMacro), - CelMacro.newReceiverMacro( - "transformMapEntry", 3, CelComprehensionsExtensions::transformMapEntryMacro), - CelMacro.newReceiverMacro( - "transformMapEntry", 4, CelComprehensionsExtensions::transformMapEntryMacro)); + return compilerLibrary.macros(); } @Override public void setParserOptions(CelParserBuilder parserBuilder) { - parserBuilder.addMacros(macros()); - } - - private static Map mapInsertMap( - Map targetMap, Map mapToMerge, RuntimeEquality equality) { - for (Object key : mapToMerge.keySet()) { - checkArgument( - !equality.findInMap(targetMap, key).isPresent(), - "insert failed: key '%s' already exists", - key); - } - - if (targetMap instanceof MutableMapValue) { - MutableMapValue wrapper = (MutableMapValue) targetMap; - wrapper.putAll(mapToMerge); - return wrapper; - } - - return ImmutableMap.builderWithExpectedSize(targetMap.size() + mapToMerge.size()) - .putAll(targetMap) - .putAll(mapToMerge) - .buildOrThrow(); - } - - private static Map mapInsertKeyValue(Object[] args, RuntimeEquality equality) { - Map mapArg = (Map) args[0]; - Object key = args[1]; - Object value = args[2]; - - checkArgument( - !equality.findInMap(mapArg, key).isPresent(), - "insert failed: key '%s' already exists", - key); - - if (mapArg instanceof MutableMapValue) { - MutableMapValue mutableMap = (MutableMapValue) mapArg; - mutableMap.put(key, value); - return mutableMap; - } - - ImmutableMap.Builder builder = - ImmutableMap.builderWithExpectedSize(mapArg.size() + 1); - return builder.put(key, value).putAll(mapArg).buildOrThrow(); - } - - private static Optional expandAllMacro( - CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { - checkNotNull(exprFactory); - checkNotNull(target); - checkArgument(arguments.size() == 3); - CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); - if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg0); - } - CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); - if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg1); - } - CelExpr arg2 = checkNotNull(arguments.get(2)); - CelExpr accuInit = exprFactory.newBoolLiteral(true); - CelExpr condition = - exprFactory.newGlobalCall( - Operator.NOT_STRICTLY_FALSE.getFunction(), - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); - CelExpr step = - exprFactory.newGlobalCall( - Operator.LOGICAL_AND.getFunction(), - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), - arg2); - CelExpr result = exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()); - return Optional.of( - exprFactory.fold( - arg0.ident().name(), - arg1.ident().name(), - target, - exprFactory.getAccumulatorVarName(), - accuInit, - condition, - step, - result)); - } - - private static Optional expandExistsMacro( - CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { - checkNotNull(exprFactory); - checkNotNull(target); - checkArgument(arguments.size() == 3); - CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); - if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg0); - } - CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); - if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg1); - } - CelExpr arg2 = checkNotNull(arguments.get(2)); - CelExpr accuInit = exprFactory.newBoolLiteral(false); - CelExpr condition = - exprFactory.newGlobalCall( - Operator.NOT_STRICTLY_FALSE.getFunction(), - exprFactory.newGlobalCall( - Operator.LOGICAL_NOT.getFunction(), - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()))); - CelExpr step = - exprFactory.newGlobalCall( - Operator.LOGICAL_OR.getFunction(), - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), - arg2); - CelExpr result = exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()); - return Optional.of( - exprFactory.fold( - arg0.ident().name(), - arg1.ident().name(), - target, - exprFactory.getAccumulatorVarName(), - accuInit, - condition, - step, - result)); - } - - private static Optional expandExistsOneMacro( - CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { - checkNotNull(exprFactory); - checkNotNull(target); - checkArgument(arguments.size() == 3); - CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); - if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg0); - } - CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); - if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg1); - } - CelExpr arg2 = checkNotNull(arguments.get(2)); - CelExpr accuInit = exprFactory.newIntLiteral(0); - CelExpr condition = exprFactory.newBoolLiteral(true); - CelExpr step = - exprFactory.newGlobalCall( - Operator.CONDITIONAL.getFunction(), - arg2, - exprFactory.newGlobalCall( - Operator.ADD.getFunction(), - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), - exprFactory.newIntLiteral(1)), - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); - CelExpr result = - exprFactory.newGlobalCall( - Operator.EQUALS.getFunction(), - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), - exprFactory.newIntLiteral(1)); - return Optional.of( - exprFactory.fold( - arg0.ident().name(), - arg1.ident().name(), - target, - exprFactory.getAccumulatorVarName(), - accuInit, - condition, - step, - result)); - } - - private static Optional transformListMacro( - CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { - checkNotNull(exprFactory); - checkNotNull(target); - checkArgument(arguments.size() == 3 || arguments.size() == 4); - CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); - if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg0); - } - CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); - if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg1); - } - CelExpr transform; - CelExpr filter = null; - if (arguments.size() == 4) { - filter = checkNotNull(arguments.get(2)); - transform = checkNotNull(arguments.get(3)); - } else { - transform = checkNotNull(arguments.get(2)); - } - CelExpr accuInit = exprFactory.newList(); - CelExpr condition = exprFactory.newBoolLiteral(true); - CelExpr step = - exprFactory.newGlobalCall( - Operator.ADD.getFunction(), - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), - exprFactory.newList(transform)); - if (filter != null) { - step = - exprFactory.newGlobalCall( - Operator.CONDITIONAL.getFunction(), - filter, - step, - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); - } - return Optional.of( - exprFactory.fold( - arg0.ident().name(), - arg1.ident().name(), - target, - exprFactory.getAccumulatorVarName(), - accuInit, - condition, - step, - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()))); - } - - private static Optional transformMapMacro( - CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { - checkNotNull(exprFactory); - checkNotNull(target); - checkArgument(arguments.size() == 3 || arguments.size() == 4); - CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); - if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg0); - } - CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); - if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg1); - } - CelExpr transform; - CelExpr filter = null; - if (arguments.size() == 4) { - filter = checkNotNull(arguments.get(2)); - transform = checkNotNull(arguments.get(3)); - } else { - transform = checkNotNull(arguments.get(2)); - } - CelExpr accuInit = exprFactory.newMap(); - CelExpr condition = exprFactory.newBoolLiteral(true); - CelExpr step = - exprFactory.newGlobalCall( - MAP_INSERT_FUNCTION, - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), - arg0, - transform); - if (filter != null) { - step = - exprFactory.newGlobalCall( - Operator.CONDITIONAL.getFunction(), - filter, - step, - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); - } - return Optional.of( - exprFactory.fold( - arg0.ident().name(), - arg1.ident().name(), - target, - exprFactory.getAccumulatorVarName(), - accuInit, - condition, - step, - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()))); + compilerLibrary.setParserOptions(parserBuilder); } - private static Optional transformMapEntryMacro( - CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { - checkNotNull(exprFactory); - checkNotNull(target); - checkArgument(arguments.size() == 3 || arguments.size() == 4); - CelExpr arg0 = validatedIterationVariable(exprFactory, arguments.get(0)); - if (arg0.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg0); - } - CelExpr arg1 = validatedIterationVariable(exprFactory, arguments.get(1)); - if (arg1.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of(arg1); - } - CelExpr transform; - CelExpr filter = null; - if (arguments.size() == 4) { - filter = checkNotNull(arguments.get(2)); - transform = checkNotNull(arguments.get(3)); - } else { - transform = checkNotNull(arguments.get(2)); - } - CelExpr accuInit = exprFactory.newMap(); - CelExpr condition = exprFactory.newBoolLiteral(true); - CelExpr step = - exprFactory.newGlobalCall( - MAP_INSERT_FUNCTION, - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), - transform); - if (filter != null) { - step = - exprFactory.newGlobalCall( - Operator.CONDITIONAL.getFunction(), - filter, - step, - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); - } - return Optional.of( - exprFactory.fold( - arg0.ident().name(), - arg1.ident().name(), - target, - exprFactory.getAccumulatorVarName(), - accuInit, - condition, - step, - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()))); - } - - private static CelExpr validatedIterationVariable( - CelMacroExprFactory exprFactory, CelExpr argument) { - CelExpr arg = checkNotNull(argument); - if (!isSimpleIdentifier(arg)) { - return reportArgumentError(exprFactory, arg); - } else if (arg.exprKind().ident().name().equals(exprFactory.getAccumulatorVarName()) - || arg.exprKind().ident().name().equals("__result__")) { - return reportAccumulatorOverwriteError(exprFactory, arg); - } else { - return arg; - } - } - - private static boolean isSimpleIdentifier(CelExpr expr) { - return expr.getKind() == CelExpr.ExprKind.Kind.IDENT - && !expr.ident().name().isEmpty() - && !expr.ident().name().startsWith("."); + @Override + public void setCheckerOptions(CelCheckerBuilder checkerBuilder) { + compilerLibrary.setCheckerOptions(checkerBuilder); } - private static CelExpr reportArgumentError(CelMacroExprFactory exprFactory, CelExpr argument) { - return exprFactory.reportError( - CelIssue.formatError( - exprFactory.getSourceLocation(argument), "The argument must be a simple name")); + @Override + public void setRuntimeOptions(CelRuntimeBuilder runtimeBuilder) { + throw new UnsupportedOperationException("Unsupported"); } - private static CelExpr reportAccumulatorOverwriteError( - CelMacroExprFactory exprFactory, CelExpr argument) { - return exprFactory.reportError( - CelIssue.formatError( - exprFactory.getSourceLocation(argument), - String.format( - "The iteration variable %s overwrites accumulator variable", - argument.ident().name()))); + @Override + public void setRuntimeOptions( + CelRuntimeBuilder runtimeBuilder, RuntimeEquality runtimeEquality, CelOptions celOptions) { + runtimeBuilder.addFunctionBindings(runtimeLibrary.newFunctionBindings(runtimeEquality)); } } diff --git a/extensions/src/main/java/dev/cel/extensions/CelComprehensionsRuntimeLibrary.java b/extensions/src/main/java/dev/cel/extensions/CelComprehensionsRuntimeLibrary.java new file mode 100644 index 000000000..2ec6ca988 --- /dev/null +++ b/extensions/src/main/java/dev/cel/extensions/CelComprehensionsRuntimeLibrary.java @@ -0,0 +1,175 @@ +// Copyright 2024 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.extensions; + +import static com.google.common.base.Preconditions.checkArgument; + +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; +import com.google.common.collect.ImmutableSet; +import com.google.errorprone.annotations.Immutable; +import dev.cel.common.CelOptions; +import dev.cel.common.values.MutableMapValue; +import dev.cel.runtime.CelFunctionBinding; +import dev.cel.runtime.CelInternalLiteRuntimeLibrary; +import dev.cel.runtime.CelLiteRuntimeBuilder; +import dev.cel.runtime.RuntimeEquality; +import java.util.Map; +import java.util.Set; + +/** Runtime implementation of CEL two variable comprehensions extensions. */ +@Immutable +public final class CelComprehensionsRuntimeLibrary implements CelInternalLiteRuntimeLibrary { + + private static final String MAP_INSERT_FUNCTION = "cel.@mapInsert"; + private static final String MAP_INSERT_OVERLOAD_MAP_MAP = "cel_@mapInsert_map_map"; + private static final String MAP_INSERT_OVERLOAD_KEY_VALUE = "cel_@mapInsert_map_key_value"; + + /** Enumeration of functions for Comprehensions runtime extension. */ + public enum Function { + MAP_INSERT("cel.@mapInsert"); + + private final String functionName; + + public String getFunction() { + return functionName; + } + + Function(String functionName) { + this.functionName = functionName; + } + } + + private static final CelComprehensionsRuntimeLibrary VERSION_0 = + new CelComprehensionsRuntimeLibrary(ImmutableSet.of(Function.MAP_INSERT)); + + private static ImmutableSet getFunctionsForVersion(int version) { + switch (version) { + case 0: + case Integer.MAX_VALUE: + return VERSION_0.functions; + default: + throw new IllegalArgumentException( + "Unsupported 'comprehensions' extension version " + version); + } + } + + /** Returns the latest version of the 'comprehensions' runtime functions. */ + public static CelComprehensionsRuntimeLibrary comprehensions() { + return VERSION_0; + } + + /** Returns the specified version of the 'comprehensions' runtime functions. */ + public static CelComprehensionsRuntimeLibrary comprehensions(int version) { + if (version == 0 || version == Integer.MAX_VALUE) { + return VERSION_0; + } + return new CelComprehensionsRuntimeLibrary(getFunctionsForVersion(version)); + } + + /** Returns the 'comprehensions' runtime functions with only the specified functions. */ + public static CelComprehensionsRuntimeLibrary comprehensions(Function... functions) { + return comprehensions(ImmutableSet.copyOf(functions)); + } + + /** Returns the 'comprehensions' runtime functions with only the specified functions. */ + public static CelComprehensionsRuntimeLibrary comprehensions(Set functions) { + return new CelComprehensionsRuntimeLibrary(functions); + } + + private final ImmutableSet functions; + + CelComprehensionsRuntimeLibrary(Set functions) { + this.functions = ImmutableSet.copyOf(functions); + } + + ImmutableSet functions() { + return functions; + } + + @Override + public void setRuntimeOptions(CelLiteRuntimeBuilder runtimeBuilder) { + throw new UnsupportedOperationException("Unsupported"); + } + + @Override + public void setRuntimeOptions( + CelLiteRuntimeBuilder runtimeBuilder, + RuntimeEquality runtimeEquality, + CelOptions celOptions) { + runtimeBuilder.addFunctionBindings(newFunctionBindings(runtimeEquality)); + } + + public ImmutableList newFunctionBindings(RuntimeEquality runtimeEquality) { + ImmutableList.Builder bindings = ImmutableList.builder(); + if (functions.contains(Function.MAP_INSERT)) { + bindings.addAll( + CelFunctionBinding.fromOverloads( + MAP_INSERT_FUNCTION, + CelFunctionBinding.from( + MAP_INSERT_OVERLOAD_MAP_MAP, + Map.class, + Map.class, + (map1, map2) -> mapInsertMap(map1, map2, runtimeEquality)), + CelFunctionBinding.from( + MAP_INSERT_OVERLOAD_KEY_VALUE, + ImmutableList.of(Map.class, Object.class, Object.class), + args -> mapInsertKeyValue(args, runtimeEquality)))); + } + return bindings.build(); + } + + private static Map mapInsertMap( + Map targetMap, Map mapToMerge, RuntimeEquality equality) { + for (Object key : mapToMerge.keySet()) { + checkArgument( + !equality.findInMap(targetMap, key).isPresent(), + "insert failed: key '%s' already exists", + key); + } + + if (targetMap instanceof MutableMapValue) { + MutableMapValue wrapper = (MutableMapValue) targetMap; + wrapper.putAll(mapToMerge); + return wrapper; + } + + return ImmutableMap.builderWithExpectedSize(targetMap.size() + mapToMerge.size()) + .putAll(targetMap) + .putAll(mapToMerge) + .buildOrThrow(); + } + + private static Map mapInsertKeyValue(Object[] args, RuntimeEquality equality) { + Map mapArg = (Map) args[0]; + Object key = args[1]; + Object value = args[2]; + + checkArgument( + !equality.findInMap(mapArg, key).isPresent(), + "insert failed: key '%s' already exists", + key); + + if (mapArg instanceof MutableMapValue) { + MutableMapValue mutableMap = (MutableMapValue) mapArg; + mutableMap.put(key, value); + return mutableMap; + } + + ImmutableMap.Builder builder = + ImmutableMap.builderWithExpectedSize(mapArg.size() + 1); + return builder.put(key, value).putAll(mapArg).buildOrThrow(); + } +} diff --git a/extensions/src/main/java/dev/cel/extensions/CelExtensions.java b/extensions/src/main/java/dev/cel/extensions/CelExtensions.java index fc9a5874c..95b843913 100644 --- a/extensions/src/main/java/dev/cel/extensions/CelExtensions.java +++ b/extensions/src/main/java/dev/cel/extensions/CelExtensions.java @@ -386,6 +386,41 @@ public static CelComprehensionsExtensions comprehensions() { return COMPREHENSIONS_EXTENSIONS; } + /** + * Extended functions for Two Variable Comprehensions Expressions. + * + *

Refer to README.md for functions available in each version. + */ + public static CelComprehensionsExtensions comprehensions(int version) { + return CelComprehensionsExtensions.library().version(version); + } + + /** + * Extended functions for Two Variable Comprehensions Expressions. + * + *

Refer to README.md for available functions. + * + *

This will include only the specific functions denoted by {@link + * CelComprehensionsExtensions.Function}. + */ + public static CelComprehensionsExtensions comprehensions( + CelComprehensionsExtensions.Function... functions) { + return comprehensions(ImmutableSet.copyOf(functions)); + } + + /** + * Extended functions for Two Variable Comprehensions Expressions. + * + *

Refer to README.md for available functions. + * + *

This will include only the specific functions denoted by {@link + * CelComprehensionsExtensions.Function}. + */ + public static CelComprehensionsExtensions comprehensions( + Set functions) { + return new CelComprehensionsExtensions(functions); + } + /** * Extensions for supporting native Java types (POJOs) in CEL. * @@ -452,6 +487,7 @@ public static CelExtensionLibrary getE return CelSetsExtensions.library(options); case "strings": return CelStringExtensions.library(); + case "two-var-comprehensions": case "comprehensions": return CelComprehensionsExtensions.library(); // TODO: add support for remaining standard extensions diff --git a/extensions/src/test/java/dev/cel/extensions/BUILD.bazel b/extensions/src/test/java/dev/cel/extensions/BUILD.bazel index 0e3bca0fd..4168d0dfb 100644 --- a/extensions/src/test/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/test/java/dev/cel/extensions/BUILD.bazel @@ -34,6 +34,9 @@ java_library( "//compiler", "//compiler:compiler_builder", "//extensions", + "//extensions:comprehensions", + "//extensions:comprehensions_compiler_library", + "//extensions:comprehensions_runtime_library", "//extensions:encoders_compiler_library", "//extensions:encoders_runtime_library", "//extensions:extension_library", @@ -61,6 +64,8 @@ java_library( "//runtime:lite_runtime", "//runtime:lite_runtime_factory", "//runtime:partial_vars", + "//runtime:runtime_equality", + "//runtime:runtime_helpers", "//runtime:unknown_attributes", "//testing:cel_runtime_flavor", "//validator", diff --git a/extensions/src/test/java/dev/cel/extensions/CelComprehensionsExtensionsTest.java b/extensions/src/test/java/dev/cel/extensions/CelComprehensionsExtensionsTest.java index bb1784755..1a31cb3dc 100644 --- a/extensions/src/test/java/dev/cel/extensions/CelComprehensionsExtensionsTest.java +++ b/extensions/src/test/java/dev/cel/extensions/CelComprehensionsExtensionsTest.java @@ -19,10 +19,13 @@ import static org.junit.Assert.assertThrows; import com.google.common.base.Throwables; +import com.google.common.collect.ImmutableMap; +import com.google.common.collect.ImmutableSet; import com.google.testing.junit.testparameterinjector.TestParameter; import com.google.testing.junit.testparameterinjector.TestParameterInjector; import com.google.testing.junit.testparameterinjector.TestParameters; import dev.cel.bundle.Cel; +import dev.cel.bundle.CelFactory; import dev.cel.common.CelAbstractSyntaxTree; import dev.cel.common.CelFunctionDecl; import dev.cel.common.CelOptions; @@ -33,11 +36,21 @@ import dev.cel.common.exceptions.CelIndexOutOfBoundsException; import dev.cel.common.types.SimpleType; import dev.cel.common.types.TypeParamType; +import dev.cel.compiler.CelCompiler; +import dev.cel.compiler.CelCompilerFactory; import dev.cel.parser.CelMacro; import dev.cel.parser.CelStandardMacro; import dev.cel.parser.CelUnparser; import dev.cel.parser.CelUnparserFactory; import dev.cel.runtime.CelEvaluationException; +import dev.cel.runtime.CelLiteRuntime; +import dev.cel.runtime.CelLiteRuntimeBuilder; +import dev.cel.runtime.CelLiteRuntimeFactory; +import dev.cel.runtime.CelRuntime; +import dev.cel.runtime.CelRuntimeBuilder; +import dev.cel.runtime.CelRuntimeFactory; +import dev.cel.runtime.RuntimeEquality; +import dev.cel.runtime.RuntimeHelpers; import dev.cel.testing.CelRuntimeFlavor; import org.junit.Assume; import org.junit.Test; @@ -363,5 +376,180 @@ public void mutableMapValue_select_missingKeyException() throws Exception { assertThat(e).hasCauseThat().hasMessageThat().contains("key 'b' is not present in map."); } + @Test + public void separateLibraryAndRuntime_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelComprehensionsCompilerLibrary.comprehensions()) + .build(); + CelLiteRuntime celLiteRuntime = + CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .addLibraries(CelComprehensionsRuntimeLibrary.comprehensions()) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("[1, 2, 3].transformMap(i, v, v)").getAst(); + Object result = celLiteRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableMap.of(0L, 1L, 1L, 2L, 2L, 3L)); + } + + @Test + public void separateLibraryAndRuntime_withFunctionBindings_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelComprehensionsCompilerLibrary.comprehensions()) + .build(); + RuntimeEquality runtimeEquality = + RuntimeEquality.create(RuntimeHelpers.create(), CelOptions.DEFAULT); + CelRuntime celRuntime = + CelRuntimeFactory.standardCelRuntimeBuilder() + .addFunctionBindings( + CelComprehensionsRuntimeLibrary.comprehensions() + .newFunctionBindings(runtimeEquality)) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("[1, 2, 3].transformMap(i, v, v + 1)").getAst(); + Object result = celRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableMap.of(0L, 2L, 1L, 3L, 2L, 4L)); + } + + @Test + public void separateLibraryAndRuntime_unsupportedVersion_throws() { + assertThrows( + IllegalArgumentException.class, () -> CelComprehensionsCompilerLibrary.comprehensions(99)); + assertThrows( + IllegalArgumentException.class, () -> CelComprehensionsRuntimeLibrary.comprehensions(99)); + } + + @Test + public void comprehensions_subsetOfFunctions_success() throws Exception { + Cel cel = + CelFactory.standardCelBuilder() + .addCompilerLibraries( + CelExtensions.comprehensions(CelComprehensionsExtensions.Function.MAP_INSERT)) + .addRuntimeLibraries( + CelExtensions.comprehensions(CelComprehensionsExtensions.Function.MAP_INSERT)) + .build(); + + CelAbstractSyntaxTree ast = cel.compile("[1, 2, 3].transformMap(i, v, v + 1)").getAst(); + Object result = cel.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableMap.of(0L, 2L, 1L, 3L, 2L, 4L)); + assertThat( + CelComprehensionsCompilerLibrary.comprehensions( + CelComprehensionsCompilerLibrary.Function.MAP_INSERT) + .version()) + .isEqualTo(-1); + assertThat( + CelExtensions.comprehensions(CelComprehensionsExtensions.Function.MAP_INSERT).version()) + .isEqualTo(-1); + } + + @Test + public void comprehensions_setOfFunctions_success() throws Exception { + Cel cel = + CelFactory.standardCelBuilder() + .addCompilerLibraries( + CelExtensions.comprehensions( + ImmutableSet.of(CelComprehensionsExtensions.Function.MAP_INSERT))) + .addRuntimeLibraries( + CelExtensions.comprehensions( + ImmutableSet.of(CelComprehensionsExtensions.Function.MAP_INSERT))) + .build(); + + CelAbstractSyntaxTree ast = cel.compile("[1, 2, 3].transformMap(i, v, v + 1)").getAst(); + Object result = cel.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableMap.of(0L, 2L, 1L, 3L, 2L, 4L)); + assertThat( + CelComprehensionsCompilerLibrary.comprehensions( + ImmutableSet.of(CelComprehensionsCompilerLibrary.Function.MAP_INSERT)) + .version()) + .isEqualTo(-1); + assertThat( + new CelComprehensionsCompilerLibrary( + ImmutableSet.of(CelComprehensionsCompilerLibrary.Function.MAP_INSERT)) + .version()) + .isEqualTo(-1); + assertThat( + new CelComprehensionsExtensions( + ImmutableSet.of(CelComprehensionsExtensions.Function.MAP_INSERT)) + .version()) + .isEqualTo(-1); + assertThat( + CelExtensions.comprehensions( + ImmutableSet.of(CelComprehensionsExtensions.Function.MAP_INSERT)) + .version()) + .isEqualTo(-1); + } + + @Test + public void comprehensions_versioned_success() throws Exception { + Cel cel = + CelFactory.standardCelBuilder() + .addCompilerLibraries(CelExtensions.comprehensions(0)) + .addRuntimeLibraries(CelExtensions.comprehensions(0)) + .build(); + + CelAbstractSyntaxTree ast = cel.compile("[1, 2, 3].transformMap(i, v, v + 1)").getAst(); + Object result = cel.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableMap.of(0L, 2L, 1L, 3L, 2L, 4L)); + assertThat(CelComprehensionsCompilerLibrary.comprehensions(0).version()).isEqualTo(0); + assertThat(CelExtensions.comprehensions(0).version()).isEqualTo(0); + } + + @Test + public void comprehensions_noArgConstructor_success() { + CelComprehensionsExtensions extensions = new CelComprehensionsExtensions(); + assertThat(extensions.version()) + .isEqualTo(CelComprehensionsCompilerLibrary.comprehensions().version()); + assertThat(extensions.macros()).isNotEmpty(); + } + + @Test + public void runtimeLibrary_constructorsAndFactories() { + CelComprehensionsRuntimeLibrary lib1 = + new CelComprehensionsRuntimeLibrary( + ImmutableSet.of(CelComprehensionsRuntimeLibrary.Function.MAP_INSERT)); + assertThat(lib1.functions()) + .containsExactly(CelComprehensionsRuntimeLibrary.Function.MAP_INSERT); + + CelComprehensionsRuntimeLibrary lib2 = CelComprehensionsRuntimeLibrary.comprehensions(); + assertThat(lib2.functions()) + .containsExactly(CelComprehensionsRuntimeLibrary.Function.MAP_INSERT); + + CelComprehensionsRuntimeLibrary lib3 = CelComprehensionsRuntimeLibrary.comprehensions(0); + assertThat(lib3.functions()) + .containsExactly(CelComprehensionsRuntimeLibrary.Function.MAP_INSERT); + + CelComprehensionsRuntimeLibrary lib4 = + CelComprehensionsRuntimeLibrary.comprehensions( + CelComprehensionsRuntimeLibrary.Function.MAP_INSERT); + assertThat(lib4.functions()) + .containsExactly(CelComprehensionsRuntimeLibrary.Function.MAP_INSERT); + + CelComprehensionsRuntimeLibrary lib5 = + CelComprehensionsRuntimeLibrary.comprehensions( + ImmutableSet.of(CelComprehensionsRuntimeLibrary.Function.MAP_INSERT)); + assertThat(lib5.functions()) + .containsExactly(CelComprehensionsRuntimeLibrary.Function.MAP_INSERT); + } + + @Test + public void runtimeOptions_unsupported_throws() { + CelComprehensionsExtensions extensions = new CelComprehensionsExtensions(); + CelRuntimeBuilder runtimeBuilder = CelRuntimeFactory.standardCelRuntimeBuilder(); + assertThrows( + UnsupportedOperationException.class, () -> extensions.setRuntimeOptions(runtimeBuilder)); + + CelComprehensionsRuntimeLibrary runtimeLibrary = + CelComprehensionsRuntimeLibrary.comprehensions(); + CelLiteRuntimeBuilder liteRuntimeBuilder = CelLiteRuntimeFactory.newLiteRuntimeBuilder(); + assertThrows( + UnsupportedOperationException.class, + () -> runtimeLibrary.setRuntimeOptions(liteRuntimeBuilder)); + } } diff --git a/runtime/src/main/java/dev/cel/runtime/BUILD.bazel b/runtime/src/main/java/dev/cel/runtime/BUILD.bazel index 158c932c7..156323849 100644 --- a/runtime/src/main/java/dev/cel/runtime/BUILD.bazel +++ b/runtime/src/main/java/dev/cel/runtime/BUILD.bazel @@ -34,6 +34,7 @@ DESCRIPTOR_MESSAGE_PROVIDER_SOURCES = [ # keep sorted LITE_RUNTIME_SOURCES = [ + "CelInternalLiteRuntimeLibrary.java", "CelLiteRuntime.java", "CelLiteRuntimeBuilder.java", "CelLiteRuntimeLibrary.java", @@ -41,7 +42,7 @@ LITE_RUNTIME_SOURCES = [ # keep sorted LITE_RUNTIME_IMPL_SOURCES = [ - "LiteRuntimeImpl.java", + "CelLiteRuntimeImpl.java", ] # keep sorted @@ -107,6 +108,7 @@ java_library( "@maven//:com_google_guava_guava", "@maven//:com_google_protobuf_protobuf_java", "@maven//:org_jspecify_jspecify", + "@maven_android//:com_google_protobuf_protobuf_javalite", ], ) @@ -162,8 +164,8 @@ java_library( ":runtime_helpers", "//common/annotations", "@maven//:com_google_guava_guava", - "@maven//:com_google_protobuf_protobuf_java", "@maven//:org_jspecify_jspecify", + "@maven_android//:com_google_protobuf_protobuf_javalite", ], ) @@ -212,6 +214,7 @@ java_library( "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", "@maven//:com_google_protobuf_protobuf_java", + "@maven_android//:com_google_protobuf_protobuf_javalite", ], ) @@ -318,6 +321,7 @@ java_library( "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", "@maven//:com_google_protobuf_protobuf_java", + "@maven_android//:com_google_protobuf_protobuf_javalite", ], ) @@ -375,7 +379,7 @@ java_library( "//common/internal:comparison_functions", "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", - "@maven//:com_google_protobuf_protobuf_java", + "@maven_android//:com_google_protobuf_protobuf_javalite", ], ) @@ -458,6 +462,7 @@ java_library( "@maven//:com_google_protobuf_protobuf_java", "@maven//:com_google_re2j_re2j", "@maven//:org_threeten_threeten_extra", + "@maven_android//:com_google_protobuf_protobuf_javalite", ], ) @@ -474,6 +479,7 @@ java_library( "//common/annotations", "//common/internal:dynamic_proto", "@maven//:com_google_protobuf_protobuf_java", + "@maven_android//:com_google_protobuf_protobuf_javalite", ], ) @@ -853,6 +859,7 @@ java_library( "//common/values:proto_message_value_provider", "//runtime:activation", "//runtime:interpretable", + "//runtime:program", "//runtime:proto_message_activation_factory", "//runtime/planner:planned_program", "//runtime/planner:program_planner", @@ -951,7 +958,6 @@ java_library( "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", "@maven//:com_google_protobuf_protobuf_java", - "@maven//:org_jspecify_jspecify", ], ) @@ -964,6 +970,7 @@ java_library( ":evaluation_exception", ":function_binding", ":program", + ":runtime_equality", "//common:cel_ast", "//common:container", "//common:options", @@ -1084,7 +1091,8 @@ java_library( deps = [ ":unknown_attributes", "//common:cel_ast", - "//common:compiler_common", + "//common:cel_validation_exception", + "//common:cel_validation_result", "//common:operator", "//common/ast", "//parser:parser_builder", @@ -1180,6 +1188,7 @@ cel_android_library( ":evaluation_exception", ":function_binding_android", ":program_android", + ":runtime_equality_android", "//common:cel_ast_android", "//common:container_android", "//common:options", @@ -1378,7 +1387,6 @@ java_library( "//:auto_value", "@maven//:com_google_code_findbugs_annotations", "@maven//:com_google_errorprone_error_prone_annotations", - "@maven//:org_jspecify_jspecify", ], ) @@ -1393,7 +1401,6 @@ cel_android_library( "//:auto_value", "@maven//:com_google_code_findbugs_annotations", "@maven//:com_google_errorprone_error_prone_annotations", - "@maven//:org_jspecify_jspecify", ], ) diff --git a/runtime/src/main/java/dev/cel/runtime/CelInternalLiteRuntimeLibrary.java b/runtime/src/main/java/dev/cel/runtime/CelInternalLiteRuntimeLibrary.java new file mode 100644 index 000000000..1bda3c8b6 --- /dev/null +++ b/runtime/src/main/java/dev/cel/runtime/CelInternalLiteRuntimeLibrary.java @@ -0,0 +1,35 @@ +// 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.runtime; + +import dev.cel.common.CelOptions; +import dev.cel.common.annotations.Internal; + +/** + * CelInternalLiteRuntimeLibrary defines the interface to extend functionalities beyond the CEL + * standard functions for {@link CelLiteRuntime}, with access to runtime internals. This is not + * intended for general use. + * + *

CEL Library Internals. Do Not Use. + */ +@Internal +public interface CelInternalLiteRuntimeLibrary extends CelLiteRuntimeLibrary { + + /** + * Configures the runtime to support the library implementation, such as adding function bindings. + */ + void setRuntimeOptions( + CelLiteRuntimeBuilder runtimeBuilder, RuntimeEquality runtimeEquality, CelOptions celOptions); +} diff --git a/runtime/src/main/java/dev/cel/runtime/CelLiteRuntimeFactory.java b/runtime/src/main/java/dev/cel/runtime/CelLiteRuntimeFactory.java index 260a62980..4d24a680a 100644 --- a/runtime/src/main/java/dev/cel/runtime/CelLiteRuntimeFactory.java +++ b/runtime/src/main/java/dev/cel/runtime/CelLiteRuntimeFactory.java @@ -22,7 +22,7 @@ public final class CelLiteRuntimeFactory { /** Create a new builder for constructing a {@code CelLiteRuntime} instance. */ public static CelLiteRuntimeBuilder newLiteRuntimeBuilder() { - return LiteRuntimeImpl.newBuilder(); + return CelLiteRuntimeImpl.newBuilder(); } private CelLiteRuntimeFactory() {} diff --git a/runtime/src/main/java/dev/cel/runtime/LiteRuntimeImpl.java b/runtime/src/main/java/dev/cel/runtime/CelLiteRuntimeImpl.java similarity index 96% rename from runtime/src/main/java/dev/cel/runtime/LiteRuntimeImpl.java rename to runtime/src/main/java/dev/cel/runtime/CelLiteRuntimeImpl.java index 0c13f8a32..46f25ce02 100644 --- a/runtime/src/main/java/dev/cel/runtime/LiteRuntimeImpl.java +++ b/runtime/src/main/java/dev/cel/runtime/CelLiteRuntimeImpl.java @@ -37,7 +37,7 @@ import java.util.Optional; @ThreadSafe -final class LiteRuntimeImpl implements CelLiteRuntime { +final class CelLiteRuntimeImpl implements CelLiteRuntime { private final ProgramPlanner planner; private final CelOptions celOptions; private final ImmutableList customFunctionBindings; @@ -177,14 +177,21 @@ private static void assertAllowedCelOptions(CelOptions celOptions) { @Override public CelLiteRuntime build() { assertAllowedCelOptions(celOptions); + RuntimeHelpers runtimeHelpers = RuntimeHelpers.create(); + RuntimeEquality runtimeEquality = RuntimeEquality.create(runtimeHelpers, celOptions); ImmutableSet runtimeLibs = runtimeLibrariesBuilder.build(); - runtimeLibs.forEach(lib -> lib.setRuntimeOptions(this)); + for (CelLiteRuntimeLibrary lib : runtimeLibs) { + if (lib instanceof CelInternalLiteRuntimeLibrary) { + ((CelInternalLiteRuntimeLibrary) lib) + .setRuntimeOptions(this, runtimeEquality, celOptions); + } else { + lib.setRuntimeOptions(this); + } + } ImmutableMap.Builder functionBindingsBuilder = ImmutableMap.builder(); - RuntimeHelpers runtimeHelpers = RuntimeHelpers.create(); - RuntimeEquality runtimeEquality = RuntimeEquality.create(runtimeHelpers, celOptions); ImmutableSet standardFunctions = standardFunctionBuilder.build(); if (!standardFunctions.isEmpty()) { for (CelStandardFunction standardFunction : standardFunctions) { @@ -235,7 +242,7 @@ public CelLiteRuntime build() { CelAsyncEvaluationOptions.defaultOptions(), /* asyncExecutor= */ null); - return new LiteRuntimeImpl( + return new CelLiteRuntimeImpl( planner, celOptions, customFunctionBindings.values(), @@ -272,7 +279,7 @@ static CelLiteRuntimeBuilder newBuilder() { return new Builder(); } - private LiteRuntimeImpl( + private CelLiteRuntimeImpl( ProgramPlanner planner, CelOptions celOptions, Iterable customFunctionBindings, diff --git a/runtime/src/test/java/dev/cel/runtime/BUILD.bazel b/runtime/src/test/java/dev/cel/runtime/BUILD.bazel index 54f936cb6..0371947e5 100644 --- a/runtime/src/test/java/dev/cel/runtime/BUILD.bazel +++ b/runtime/src/test/java/dev/cel/runtime/BUILD.bazel @@ -194,6 +194,7 @@ cel_android_local_test( "//common/values:cel_byte_string", "//common/values:cel_value_provider_android", "//common/values:proto_message_lite_value_provider_android", + "//extensions:comprehensions_runtime_library_android", "//extensions:encoders_runtime_library_android", "//extensions:lists_runtime_library_android", "//extensions:math_runtime_library_android", diff --git a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java index 27ae853ae..7b7119147 100644 --- a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java +++ b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java @@ -45,6 +45,7 @@ import dev.cel.expr.conformance.proto3.NestedTestAllTypesCelLiteDescriptor; import dev.cel.expr.conformance.proto3.TestAllTypes; import dev.cel.expr.conformance.proto3.TestAllTypesCelLiteDescriptor; +import dev.cel.extensions.CelComprehensionsRuntimeLibrary; import dev.cel.extensions.CelEncoderRuntimeLibrary; import dev.cel.extensions.CelListsRuntimeLibrary; import dev.cel.extensions.CelMathRuntimeLibrary; @@ -144,8 +145,8 @@ public void toRuntimeBuilder_propertiesCopied() { .addLibraries(runtimeExtension); CelLiteRuntime runtime = runtimeBuilder.build(); - LiteRuntimeImpl.Builder newRuntimeBuilder = - (LiteRuntimeImpl.Builder) runtime.toRuntimeBuilder(); + CelLiteRuntimeImpl.Builder newRuntimeBuilder = + (CelLiteRuntimeImpl.Builder) runtime.toRuntimeBuilder(); assertThat(newRuntimeBuilder.celOptions).isEqualTo(celOptions); assertThat(newRuntimeBuilder.celValueProvider).isSameInstanceAs(celValueProvider); @@ -792,6 +793,19 @@ public void eval_listsExtension() throws Exception { assertThat(runtime.createProgram(ast).eval()).isEqualTo(ImmutableList.of(1L, 2L)); } + @Test + public void eval_comprehensionsExtension() throws Exception { + CelLiteRuntime runtime = + CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .addLibraries(CelComprehensionsRuntimeLibrary.comprehensions()) + .build(); + // Expr: [1, 2, 3].transformMap(i, v, v) + CelAbstractSyntaxTree ast = readCheckedExpr("compiled_comprehensions_transform_map"); + + assertThat(runtime.createProgram(ast).eval()) + .isEqualTo(ImmutableMap.of(0L, 1L, 1L, 2L, 2L, 3L)); + } + @Test public void eval_regexExtension() throws Exception { CelLiteRuntime runtime = diff --git a/testing/src/main/java/dev/cel/testing/compiled/BUILD.bazel b/testing/src/main/java/dev/cel/testing/compiled/BUILD.bazel index 92786afc9..104633e5d 100644 --- a/testing/src/main/java/dev/cel/testing/compiled/BUILD.bazel +++ b/testing/src/main/java/dev/cel/testing/compiled/BUILD.bazel @@ -44,6 +44,7 @@ java_library( resources = [ ":compiled_comprehension", ":compiled_comprehension_exists", + ":compiled_comprehensions_transform_map", ":compiled_custom_functions", ":compiled_encoders_encode", ":compiled_extended_env", @@ -110,6 +111,12 @@ compile_cel( expression = "cel.bind(x, 10, math.greatest([1,x])) < int(' 11 '.trim()) && optional.none().orValue(true) && [].flatten() == []", ) +compile_cel( + name = "compiled_comprehensions_transform_map", + environment = "//testing/environment:all_extensions", + expression = "[1, 2, 3].transformMap(i, v, v)", +) + compile_cel( name = "compiled_encoders_encode", environment = "//testing/environment:all_extensions", diff --git a/testing/src/test/resources/environment/all_extensions.yaml b/testing/src/test/resources/environment/all_extensions.yaml index 034a051b1..80622bacc 100644 --- a/testing/src/test/resources/environment/all_extensions.yaml +++ b/testing/src/test/resources/environment/all_extensions.yaml @@ -15,6 +15,7 @@ name: "all-extensions" extensions: - name: "bindings" +- name: "two-var-comprehensions" - name: "encoders" - name: "lists" - name: "math"