From 14b79d68c27f381a2d2f05fea964348a4523a50b Mon Sep 17 00:00:00 2001 From: Dmitri Plotnikov Date: Fri, 2 Oct 2026 15:18:40 -0700 Subject: [PATCH] Split CelListsExtensions into CelListsCompilerLibrary and CelListsRuntimeLibrary to allow its usage in Lite runtime. PiperOrigin-RevId: 992566081 --- extensions/BUILD.bazel | 22 + .../main/java/dev/cel/extensions/BUILD.bazel | 67 ++- .../extensions/CelListsCompilerLibrary.java | 305 ++++++++++ .../cel/extensions/CelListsExtensions.java | 547 +++--------------- .../extensions/CelListsRuntimeLibrary.java | 405 +++++++++++++ .../test/java/dev/cel/extensions/BUILD.bazel | 3 + .../extensions/CelListsExtensionsTest.java | 143 +++++ .../src/test/java/dev/cel/runtime/BUILD.bazel | 1 + .../runtime/CelLiteRuntimeAndroidTest.java | 13 + .../java/dev/cel/testing/compiled/BUILD.bazel | 7 + 10 files changed, 1030 insertions(+), 483 deletions(-) create mode 100644 extensions/src/main/java/dev/cel/extensions/CelListsCompilerLibrary.java create mode 100644 extensions/src/main/java/dev/cel/extensions/CelListsRuntimeLibrary.java diff --git a/extensions/BUILD.bazel b/extensions/BUILD.bazel index 4e15feb90..8ef633901 100644 --- a/extensions/BUILD.bazel +++ b/extensions/BUILD.bazel @@ -119,3 +119,25 @@ cel_android_library( name = "encoders_runtime_library_android", exports = ["//extensions/src/main/java/dev/cel/extensions:encoders_runtime_library_android"], ) + +java_library( + name = "lists", + exports = ["//extensions/src/main/java/dev/cel/extensions:lists"], +) + +java_library( + name = "lists_compiler_library", + visibility = ["//:internal"], + exports = ["//extensions/src/main/java/dev/cel/extensions:lists_compiler_library"], +) + +java_library( + name = "lists_runtime_library", + visibility = ["//:internal"], + exports = ["//extensions/src/main/java/dev/cel/extensions:lists_runtime_library"], +) + +cel_android_library( + name = "lists_runtime_library_android", + exports = ["//extensions/src/main/java/dev/cel/extensions:lists_runtime_library_android"], +) diff --git a/extensions/src/main/java/dev/cel/extensions/BUILD.bazel b/extensions/src/main/java/dev/cel/extensions/BUILD.bazel index e5892b686..1f48552a4 100644 --- a/extensions/src/main/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/main/java/dev/cel/extensions/BUILD.bazel @@ -405,26 +405,81 @@ java_library( tags = [ ], deps = [ + ":extension_library", + ":lists_compiler_library", + ":lists_runtime_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 = "lists_compiler_library", + srcs = ["CelListsCompilerLibrary.java"], + tags = [ + ], + deps = [ + ":extension_library", + "//checker:checker_builder", + "//common:cel_function_decl", + "//common:cel_issue", + "//common:cel_overload_decl", + "//common:operator", "//common/ast", - "//common/internal:comparison_functions", "//common/types", "//common/types:type_providers", - "//common/values:cel_byte_string", "//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 = "lists_runtime_library", + srcs = ["CelListsRuntimeLibrary.java"], + tags = [ + ], + deps = [ + "//common:options", + "//common/internal:comparison_functions", + "//common/values:cel_byte_string", "//runtime:function_binding", + "//runtime:lite_runtime", "//runtime:runtime_equality", + "//runtime:runtime_helpers", + "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", ], ) +cel_android_library( + name = "lists_runtime_library_android", + srcs = ["CelListsRuntimeLibrary.java"], + tags = [ + ], + deps = [ + "//common:options", + "//common/internal:comparison_functions_android", + "//common/values:cel_byte_string", + "//runtime:function_binding_android", + "//runtime:lite_runtime_android", + "//runtime:runtime_equality_android", + "//runtime:runtime_helpers_android", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven_android//:com_google_guava_guava", + ], +) + java_library( name = "regex", srcs = ["CelRegexExtensions.java"], diff --git a/extensions/src/main/java/dev/cel/extensions/CelListsCompilerLibrary.java b/extensions/src/main/java/dev/cel/extensions/CelListsCompilerLibrary.java new file mode 100644 index 000000000..03edaddd1 --- /dev/null +++ b/extensions/src/main/java/dev/cel/extensions/CelListsCompilerLibrary.java @@ -0,0 +1,305 @@ +// 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 static com.google.common.base.Preconditions.checkNotNull; +import static com.google.common.collect.ImmutableSet.toImmutableSet; + +import com.google.common.base.Ascii; +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.CelType; +import dev.cel.common.types.ListType; +import dev.cel.common.types.SimpleType; +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.Optional; +import java.util.Set; + +/** Internal implementation of CEL lists compile-time extensions. */ +@Immutable +public final class CelListsCompilerLibrary + implements CelCompilerLibrary, CelExtensionLibrary.FeatureSet { + + private static final String UNUSED_ITER_VAR = "#unused"; + private static final String SORT_BY_INPUT_VAR = "@__sortBy_input__"; + + /** Supported functions for Lists extension library. */ + public enum Function { + SLICE( + CelFunctionDecl.newFunctionDeclaration( + "slice", + CelOverloadDecl.newMemberOverload( + "list_slice", + "Returns a new sub-list using the indices provided", + ListType.create(TypeParamType.create("T")), + ListType.create(TypeParamType.create("T")), + SimpleType.INT, + SimpleType.INT))), + FLATTEN( + CelFunctionDecl.newFunctionDeclaration( + "flatten", + CelOverloadDecl.newMemberOverload( + "list_flatten", + "Flattens a list by a single level", + ListType.create(TypeParamType.create("T")), + ListType.create(ListType.create(TypeParamType.create("T")))), + CelOverloadDecl.newMemberOverload( + "list_flatten_list_int", + "Flattens a list to the specified level. A negative depth value flattens the list" + + " recursively to its deepest level.", + ListType.create(SimpleType.DYN), + ListType.create(SimpleType.DYN), + SimpleType.INT))), + RANGE( + CelFunctionDecl.newFunctionDeclaration( + "lists.range", + CelOverloadDecl.newGlobalOverload( + "lists_range", + "Returns a list of integers from 0 to n-1.", + ListType.create(SimpleType.INT), + SimpleType.INT))), + DISTINCT( + CelFunctionDecl.newFunctionDeclaration( + "distinct", + CelOverloadDecl.newMemberOverload( + "list_distinct", + "Returns the distinct elements of a list", + ListType.create(TypeParamType.create("T")), + ListType.create(TypeParamType.create("T"))))), + REVERSE( + CelFunctionDecl.newFunctionDeclaration( + "reverse", + CelOverloadDecl.newMemberOverload( + "list_reverse", + "Returns the elements of a list in reverse order", + ListType.create(TypeParamType.create("T")), + ListType.create(TypeParamType.create("T"))))), + SORT( + CelFunctionDecl.newFunctionDeclaration( + "sort", + CelOverloadDecl.newMemberOverload( + "list_sort", + "Sorts a list with comparable elements.", + ListType.create(TypeParamType.create("T")), + ListType.create(TypeParamType.create("T"))))), + SORT_BY(createSortByFunctionDecl(comparableSortKeyTypes())); + + private static ImmutableList comparableSortKeyTypes() { + return ImmutableList.of( + SimpleType.INT, + SimpleType.UINT, + SimpleType.DOUBLE, + SimpleType.BOOL, + SimpleType.STRING, + SimpleType.BYTES, + SimpleType.DURATION, + SimpleType.TIMESTAMP); + } + + private static CelFunctionDecl createSortByFunctionDecl(ImmutableList keyTypes) { + ImmutableList.Builder overloads = ImmutableList.builder(); + for (CelType type : keyTypes) { + String typeName = Ascii.toLowerCase(type.kind().name()); + overloads.add( + CelOverloadDecl.newMemberOverload( + String.format("list_%s_sortByAssociatedKeys", typeName), + "Sorts a list by an associated list of keys. Used by the 'sortBy' macro", + ListType.create(TypeParamType.create("T")), + ListType.create(TypeParamType.create("T")), + ListType.create(type))); + } + return CelFunctionDecl.newFunctionDeclaration("@sortByAssociatedKeys", overloads.build()); + } + + private final CelFunctionDecl functionDecl; + + public String getFunction() { + return functionDecl.name(); + } + + public CelFunctionDecl getFunctionDecl() { + return functionDecl; + } + + Function(CelFunctionDecl functionDecl) { + this.functionDecl = functionDecl; + } + } + + private static ImmutableSet getFunctionsForVersion(int version) { + switch (version) { + case 0: + return ImmutableSet.of(Function.SLICE); + case 1: + return ImmutableSet.of(Function.SLICE, Function.FLATTEN); + case 2: + case Integer.MAX_VALUE: + return ImmutableSet.copyOf(Function.values()); + default: + throw new IllegalArgumentException("Unsupported 'lists' extension version " + version); + } + } + + private static final class Library implements CelExtensionLibrary { + private final CelListsCompilerLibrary version0 = + new CelListsCompilerLibrary(0, getFunctionsForVersion(0)); + private final CelListsCompilerLibrary version1 = + new CelListsCompilerLibrary(1, getFunctionsForVersion(1)); + private final CelListsCompilerLibrary version2 = + new CelListsCompilerLibrary(2, getFunctionsForVersion(2)); + + @Override + public String name() { + return "lists"; + } + + @Override + public ImmutableSet versions() { + return ImmutableSet.of(version0, version1, version2); + } + } + + private static final Library LIBRARY = new Library(); + + public static CelExtensionLibrary library() { + return LIBRARY; + } + + /** Returns the latest version of the 'lists' compiler extension. */ + public static CelListsCompilerLibrary lists() { + return library().latest(); + } + + /** Returns the specified version of the 'lists' compiler extension. */ + public static CelListsCompilerLibrary lists(int version) { + return library().version(version); + } + + /** Returns the 'lists' compiler extension with only the specified functions. */ + public static CelListsCompilerLibrary lists(Function... functions) { + return lists(ImmutableSet.copyOf(functions)); + } + + /** Returns the 'lists' compiler extension with only the specified functions. */ + public static CelListsCompilerLibrary lists(Set functions) { + return new CelListsCompilerLibrary(functions); + } + + private final ImmutableSet functions; + private final int version; + + CelListsCompilerLibrary(Set functions) { + this(-1, functions); + } + + private CelListsCompilerLibrary(int version, Set functions) { + this.version = version; + this.functions = ImmutableSet.copyOf(functions); + } + + @Override + public int version() { + return version; + } + + @Override + public ImmutableSet functions() { + return functions.stream().map(Function::getFunctionDecl).collect(toImmutableSet()); + } + + @Override + public ImmutableSet macros() { + if (functions.contains(Function.SORT_BY)) { + return ImmutableSet.of( + CelMacro.newReceiverMacro("sortBy", 2, CelListsCompilerLibrary::sortByMacro)); + } + return ImmutableSet.of(); + } + + @Override + public void setParserOptions(CelParserBuilder parserBuilder) { + parserBuilder.addMacros(macros()); + } + + @Override + public void setCheckerOptions(CelCheckerBuilder checkerBuilder) { + functions.forEach(function -> checkerBuilder.addFunctionDeclarations(function.functionDecl)); + } + + private static Optional sortByMacro( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { + checkNotNull(exprFactory); + checkNotNull(target); + checkArgument(arguments.size() == 2); + CelExpr varIdent = checkNotNull(arguments.get(0)); + if (varIdent.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { + return Optional.of( + exprFactory.reportError( + CelIssue.formatError( + exprFactory.getSourceLocation(varIdent), + "sortBy(var, ...) variable name must be a simple identifier"))); + } + + String varName = varIdent.ident().name(); + CelExpr sortKeyExpr = checkNotNull(arguments.get(1)); + + // Build map comprehension: @__sortBy_input__.map(varName, sortKeyExpr) + CelExpr targetIdent = exprFactory.newIdentifier(SORT_BY_INPUT_VAR); + CelExpr mapStep = + exprFactory.newGlobalCall( + Operator.ADD.getFunction(), + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), + exprFactory.newList(sortKeyExpr)); + CelExpr mapCompr = + exprFactory.fold( + varName, + targetIdent, + exprFactory.getAccumulatorVarName(), + exprFactory.newList(), + exprFactory.newBoolLiteral(true), + mapStep, + exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); + + // Build call: @__sortBy_input__.@sortByAssociatedKeys(mapCompr) + CelExpr callExpr = + exprFactory.newReceiverCall( + Function.SORT_BY.getFunction(), exprFactory.newIdentifier(SORT_BY_INPUT_VAR), mapCompr); + + // Build bind: cel.bind(@__sortBy_input__, target, callExpr) + CelExpr bindExpr = + exprFactory.fold( + UNUSED_ITER_VAR, + exprFactory.newList(), + SORT_BY_INPUT_VAR, + target, + exprFactory.newBoolLiteral(false), + exprFactory.newIdentifier(SORT_BY_INPUT_VAR), + callExpr); + + return Optional.of(bindExpr); + } +} diff --git a/extensions/src/main/java/dev/cel/extensions/CelListsExtensions.java b/extensions/src/main/java/dev/cel/extensions/CelListsExtensions.java index a45407180..9c4aac7f2 100644 --- a/extensions/src/main/java/dev/cel/extensions/CelListsExtensions.java +++ b/extensions/src/main/java/dev/cel/extensions/CelListsExtensions.java @@ -14,269 +14,134 @@ 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.base.Ascii; -import com.google.common.base.Preconditions; -import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableSet; -import com.google.common.collect.Lists; -import com.google.common.collect.Sets; +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.internal.ComparisonFunctions; -import dev.cel.common.types.CelType; -import dev.cel.common.types.ListType; -import dev.cel.common.types.SimpleType; -import dev.cel.common.types.TypeParamType; -import dev.cel.common.values.CelByteString; 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.Arrays; -import java.util.Collection; -import java.util.Comparator; -import java.util.Iterator; -import java.util.List; -import java.util.Optional; import java.util.Set; /** Internal implementation of CEL lists extensions. */ +@Immutable public final class CelListsExtensions implements CelCompilerLibrary, CelInternalRuntimeLibrary, CelExtensionLibrary.FeatureSet { - private static final CelObjectComparator OBJECT_COMPARATOR = new CelObjectComparator(); - private static final String UNUSED_ITER_VAR = "#unused"; - private static final String SORT_BY_INPUT_VAR = "@__sortBy_input__"; - /** Supported functions for Lists extension library. */ - @SuppressWarnings({"unchecked"}) // Unchecked: Type-checker guarantees casting safety. public enum Function { - // Note! Creating dependencies on the outer class may cause circular initialization issues. - SLICE( - CelFunctionDecl.newFunctionDeclaration( - "slice", - CelOverloadDecl.newMemberOverload( - "list_slice", - "Returns a new sub-list using the indices provided", - ListType.create(TypeParamType.create("T")), - ListType.create(TypeParamType.create("T")), - SimpleType.INT, - SimpleType.INT)), - CelFunctionBinding.from( - "list_slice", - ImmutableList.of(Collection.class, Long.class, Long.class), - (args) -> { - Collection target = (Collection) args[0]; - long from = (Long) args[1]; - long to = (Long) args[2]; - return CelListsExtensions.slice(target, from, to); - })), - FLATTEN( - CelFunctionDecl.newFunctionDeclaration( - "flatten", - CelOverloadDecl.newMemberOverload( - "list_flatten", - "Flattens a list by a single level", - ListType.create(TypeParamType.create("T")), - ListType.create(ListType.create(TypeParamType.create("T")))), - CelOverloadDecl.newMemberOverload( - "list_flatten_list_int", - "Flattens a list to the specified level. A negative depth value flattens the list" - + " recursively to its deepest level.", - ListType.create(SimpleType.DYN), - ListType.create(SimpleType.DYN), - SimpleType.INT)), - CelFunctionBinding.from("list_flatten", Collection.class, list -> flatten(list, 1)), - CelFunctionBinding.from( - "list_flatten_list_int", Collection.class, Long.class, CelListsExtensions::flatten)), - RANGE( - CelFunctionDecl.newFunctionDeclaration( - "lists.range", - CelOverloadDecl.newGlobalOverload( - "lists_range", - "Returns a list of integers from 0 to n-1.", - ListType.create(SimpleType.INT), - SimpleType.INT)), - CelFunctionBinding.from("lists_range", Long.class, CelListsExtensions::genRange)), - DISTINCT( - CelFunctionDecl.newFunctionDeclaration( - "distinct", - CelOverloadDecl.newMemberOverload( - "list_distinct", - "Returns the distinct elements of a list", - ListType.create(TypeParamType.create("T")), - ListType.create(TypeParamType.create("T"))))), - REVERSE( - CelFunctionDecl.newFunctionDeclaration( - "reverse", - CelOverloadDecl.newMemberOverload( - "list_reverse", - "Returns the elements of a list in reverse order", - ListType.create(TypeParamType.create("T")), - ListType.create(TypeParamType.create("T")))), - CelFunctionBinding.from("list_reverse", Collection.class, CelListsExtensions::reverse)), - SORT( - CelFunctionDecl.newFunctionDeclaration( - "sort", - CelOverloadDecl.newMemberOverload( - "list_sort", - "Sorts a list with comparable elements.", - ListType.create(TypeParamType.create("T")), - ListType.create(TypeParamType.create("T")))), - CelFunctionBinding.from("list_sort", Collection.class, CelListsExtensions::sort)), - SORT_BY( - createSortByFunctionDecl(comparableSortKeyTypes()), - createSortByFunctionBindings(comparableSortKeyTypes())); + SLICE(CelListsCompilerLibrary.Function.SLICE, CelListsRuntimeLibrary.Function.SLICE), + FLATTEN(CelListsCompilerLibrary.Function.FLATTEN, CelListsRuntimeLibrary.Function.FLATTEN), + RANGE(CelListsCompilerLibrary.Function.RANGE, CelListsRuntimeLibrary.Function.RANGE), + DISTINCT(CelListsCompilerLibrary.Function.DISTINCT, CelListsRuntimeLibrary.Function.DISTINCT), + REVERSE(CelListsCompilerLibrary.Function.REVERSE, CelListsRuntimeLibrary.Function.REVERSE), + SORT(CelListsCompilerLibrary.Function.SORT, CelListsRuntimeLibrary.Function.SORT), + SORT_BY(CelListsCompilerLibrary.Function.SORT_BY, CelListsRuntimeLibrary.Function.SORT_BY); + + private final CelListsCompilerLibrary.Function compilerFunction; + private final CelListsRuntimeLibrary.Function runtimeFunction; - private static ImmutableList comparableSortKeyTypes() { - return ImmutableList.of( - SimpleType.INT, - SimpleType.UINT, - SimpleType.DOUBLE, - SimpleType.BOOL, - SimpleType.STRING, - SimpleType.BYTES, - SimpleType.DURATION, - SimpleType.TIMESTAMP); + String getFunction() { + return compilerFunction.getFunction(); } - private static CelFunctionDecl createSortByFunctionDecl(ImmutableList keyTypes) { - ImmutableList.Builder overloads = ImmutableList.builder(); - for (CelType type : keyTypes) { - String typeName = Ascii.toLowerCase(type.kind().name()); - overloads.add( - CelOverloadDecl.newMemberOverload( - String.format("list_%s_sortByAssociatedKeys", typeName), - "Sorts a list by an associated list of keys. Used by the 'sortBy' macro", - ListType.create(TypeParamType.create("T")), - ListType.create(TypeParamType.create("T")), - ListType.create(type))); - } - return CelFunctionDecl.newFunctionDeclaration("@sortByAssociatedKeys", overloads.build()); + Function( + CelListsCompilerLibrary.Function compilerFunction, + CelListsRuntimeLibrary.Function runtimeFunction) { + this.compilerFunction = compilerFunction; + this.runtimeFunction = runtimeFunction; } + } - private static CelFunctionBinding[] createSortByFunctionBindings( - ImmutableList keyTypes) { - return keyTypes.stream() - .map( - type -> { - String typeName = Ascii.toLowerCase(type.kind().name()); - return CelFunctionBinding.from( - String.format("list_%s_sortByAssociatedKeys", typeName), - Collection.class, - Collection.class, - CelListsExtensions::sortByAssociatedKeys); - }) - .toArray(CelFunctionBinding[]::new); + private static ImmutableSet getFunctionsForVersion(int version) { + switch (version) { + case 0: + return ImmutableSet.of(Function.SLICE); + case 1: + return ImmutableSet.of(Function.SLICE, Function.FLATTEN); + case 2: + case Integer.MAX_VALUE: + return ImmutableSet.copyOf(Function.values()); + default: + throw new IllegalArgumentException("Unsupported 'lists' extension version " + version); } + } - private final CelFunctionDecl functionDecl; - private final ImmutableSet functionBindings; + private static final class Library implements CelExtensionLibrary { + private final ImmutableSet versions; - String getFunction() { - return functionDecl.name(); + Library() { + versions = + CelListsCompilerLibrary.library().versions().stream() + .map(CelListsExtensions::new) + .collect(toImmutableSet()); } - Function(CelFunctionDecl functionDecl, CelFunctionBinding... functionBindings) { - this.functionDecl = functionDecl; - this.functionBindings = - functionBindings.length > 0 - ? CelFunctionBinding.fromOverloads(functionDecl.name(), functionBindings) - : ImmutableSet.of(); + @Override + public String name() { + return CelListsCompilerLibrary.library().name(); } - } - private static final CelExtensionLibrary LIBRARY = - new CelExtensionLibrary() { - private final CelListsExtensions version0 = - new CelListsExtensions(0, ImmutableSet.of(Function.SLICE)); - private final CelListsExtensions version1 = - new CelListsExtensions( - 1, - ImmutableSet.builder() - .addAll(version0.functions) - .add(Function.FLATTEN) - .build()); - private final CelListsExtensions version2 = - new CelListsExtensions( - 2, - ImmutableSet.builder() - .addAll(version1.functions) - .add( - Function.RANGE, - Function.DISTINCT, - Function.REVERSE, - Function.SORT, - Function.SORT_BY) - .build()); - - @Override - public String name() { - return "lists"; - } + @Override + public ImmutableSet versions() { + return versions; + } + } - @Override - public ImmutableSet versions() { - return ImmutableSet.of(version0, version1, version2); - } - }; + private static final Library LIBRARY = new Library(); static CelExtensionLibrary library() { return LIBRARY; } - private final int version; + private final CelListsCompilerLibrary compilerLibrary; private final ImmutableSet functions; - CelListsExtensions(Set functions) { - this(-1, functions); + CelListsExtensions() { + this(CelListsCompilerLibrary.lists()); } - private CelListsExtensions(int version, Set functions) { - this.version = version; + CelListsExtensions(Set functions) { + this.compilerLibrary = + new CelListsCompilerLibrary( + functions.stream().map(f -> f.compilerFunction).collect(toImmutableSet())); this.functions = ImmutableSet.copyOf(functions); } + private CelListsExtensions(CelListsCompilerLibrary compilerLibrary) { + this.compilerLibrary = compilerLibrary; + this.functions = getFunctionsForVersion(compilerLibrary.version()); + } + @Override public int version() { - return version; + return compilerLibrary.version(); } @Override public ImmutableSet functions() { - return functions.stream().map(f -> f.functionDecl).collect(toImmutableSet()); + return compilerLibrary.functions(); } @Override public ImmutableSet macros() { - if (functions.contains(Function.SORT_BY)) { - return ImmutableSet.of( - CelMacro.newReceiverMacro("sortBy", 2, CelListsExtensions::sortByMacro)); - } - return ImmutableSet.of(); + return compilerLibrary.macros(); } @Override public void setParserOptions(CelParserBuilder parserBuilder) { - parserBuilder.addMacros(macros()); + compilerLibrary.setParserOptions(parserBuilder); } @Override public void setCheckerOptions(CelCheckerBuilder checkerBuilder) { - functions.forEach(function -> checkerBuilder.addFunctionDeclarations(function.functionDecl)); + compilerLibrary.setCheckerOptions(checkerBuilder); } @Override @@ -284,285 +149,13 @@ public void setRuntimeOptions(CelRuntimeBuilder runtimeBuilder) { throw new UnsupportedOperationException("Unsupported"); } - @SuppressWarnings("unchecked") @Override public void setRuntimeOptions( CelRuntimeBuilder runtimeBuilder, RuntimeEquality runtimeEquality, CelOptions celOptions) { - functions.forEach(function -> runtimeBuilder.addFunctionBindings(function.functionBindings)); - - runtimeBuilder.addFunctionBindings( - CelFunctionBinding.fromOverloads( - "distinct", - CelFunctionBinding.from( - "list_distinct", Collection.class, (list) -> distinct(list, runtimeEquality)))); - } - - private static ImmutableList slice(Collection list, long from, long to) { - Preconditions.checkArgument(from >= 0 && to >= 0, "Negative indexes not supported"); - Preconditions.checkArgument(to >= from, "Start index must be less than or equal to end index"); - Preconditions.checkArgument(to <= list.size(), "List is length %s", list.size()); - if (list instanceof List) { - List subList = ((List) list).subList((int) from, (int) to); - if (subList instanceof ImmutableList) { - return (ImmutableList) subList; - } - return ImmutableList.copyOf(subList); - } else { - ImmutableList.Builder builder = ImmutableList.builder(); - long index = 0; - for (Iterator iterator = list.iterator(); iterator.hasNext(); index++) { - Object element = iterator.next(); - if (index >= to) { - break; - } - if (index >= from) { - builder.add(element); - } - } - return builder.build(); - } - } - - @SuppressWarnings("unchecked") - private static ImmutableList flatten(Collection list, long depth) { - Preconditions.checkArgument(depth >= 0, "Level must be non-negative"); - ImmutableList.Builder builder = ImmutableList.builder(); - for (Object element : list) { - if (!(element instanceof Collection) || depth == 0) { - builder.add(element); - } else { - Collection listItem = (Collection) element; - builder.addAll(flatten(listItem, depth - 1)); - } - } - - return builder.build(); - } - - public static ImmutableList genRange(long end) { - checkArgument(end >= 0, "lists.range: size must be non-negative, got %s", end); - checkArgument(end <= 1_000_000, "lists.range: size %s exceeds maximum allowed (1000000)", end); - - ImmutableList.Builder builder = ImmutableList.builderWithExpectedSize((int) end); - for (long i = 0; i < end; i++) { - builder.add(i); - } - return builder.build(); - } - - private static class RuntimeEqualityObjectWrapper { - private final Object object; - private final int hashCode; - private final RuntimeEquality runtimeEquality; - - RuntimeEqualityObjectWrapper(Object object, RuntimeEquality runtimeEquality) { - this.object = object; - this.runtimeEquality = runtimeEquality; - this.hashCode = runtimeEquality.hashCode(object); - } - - @Override - public int hashCode() { - return hashCode; - } - - @Override - public boolean equals(Object obj) { - if (!(obj instanceof RuntimeEqualityObjectWrapper)) { - return false; - } - return runtimeEquality.objectEquals(object, ((RuntimeEqualityObjectWrapper) obj).object); - } - } - - private static ImmutableList distinct( - Collection list, RuntimeEquality runtimeEquality) { - int size = list.size(); - ImmutableList.Builder builder = ImmutableList.builderWithExpectedSize(size); - Set distinctValues = Sets.newHashSetWithExpectedSize(size); - for (Object element : list) { - if (distinctValues.add(new RuntimeEqualityObjectWrapper(element, runtimeEquality))) { - builder.add(element); - } - } - return builder.build(); - } - - private static List reverse(Collection list) { - if (list instanceof List) { - return Lists.reverse((List) list); - } else { - ImmutableList.Builder builder = ImmutableList.builderWithExpectedSize(list.size()); - Object[] objects = list.toArray(); - for (int i = objects.length - 1; i >= 0; i--) { - builder.add(objects[i]); - } - return builder.build(); - } - } - - private static ImmutableList sort(Collection objects) { - if (objects.isEmpty()) { - return ImmutableList.of(); - } - if (objects.size() == 1) { - Object single = objects.iterator().next(); - OBJECT_COMPARATOR.compare(single, single); - return ImmutableList.of(single); - } - return ImmutableList.sortedCopyOf(OBJECT_COMPARATOR, objects); - } - - private static class CelObjectComparator implements Comparator { - - CelObjectComparator() {} - - @SuppressWarnings({"unchecked"}) - @Override - public int compare(Object o1, Object o2) { - if (o1 instanceof Number && o2 instanceof Number) { - return ComparisonFunctions.numericCompare((Number) o1, (Number) o2); - } - if (o1 instanceof CelByteString && o2 instanceof CelByteString) { - return CelByteString.unsignedLexicographicalComparator() - .compare((CelByteString) o1, (CelByteString) o2); - } - - if (!(o1 instanceof Comparable) || !(o2 instanceof Comparable)) { - throw new IllegalArgumentException("List elements must be comparable"); - } - if (o1.getClass() != o2.getClass()) { - throw new IllegalArgumentException("List elements must have the same type"); - } - return ((Comparable) o1).compareTo(o2); - } - } - - /** - * Expands the {@code list.sortBy(var, expr)} receiver macro into a binding expression that sorts - * the target list using keys evaluated by mapping {@code expr} over each element. - * - *

For example, given: - * - *

{@code
-   * myList.sortBy(item, -item.field)
-   * }
- * - *

The macro expands into: - * - *

{@code
-   * cel.bind(@__sortBy_input__, myList,
-   *     @__sortBy_input__.@sortByAssociatedKeys(
-   *         @__sortBy_input__.map(item, -item.field)
-   *     )
-   * )
-   * }
- * - *

Where: - * - *

    - *
  • {@code @__sortBy_input__.map(item, -item.field)} evaluates the sort key for each element. - *
  • {@code @sortByAssociatedKeys} stably sorts the input list elements based on their - * corresponding sort keys. - *
- */ - private static Optional sortByMacro( - CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { - checkNotNull(exprFactory); - checkNotNull(target); - checkArgument(arguments.size() == 2); - CelExpr varIdent = checkNotNull(arguments.get(0)); - if (varIdent.exprKind().getKind() != CelExpr.ExprKind.Kind.IDENT) { - return Optional.of( - exprFactory.reportError( - CelIssue.formatError( - exprFactory.getSourceLocation(varIdent), - "sortBy(var, ...) variable name must be a simple identifier"))); - } - - String varName = varIdent.ident().name(); - CelExpr sortKeyExpr = checkNotNull(arguments.get(1)); - - // Build map comprehension: @__sortBy_input__.map(varName, sortKeyExpr) - CelExpr targetIdent = exprFactory.newIdentifier(SORT_BY_INPUT_VAR); - CelExpr mapStep = - exprFactory.newGlobalCall( - Operator.ADD.getFunction(), - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName()), - exprFactory.newList(sortKeyExpr)); - CelExpr mapCompr = - exprFactory.fold( - varName, - targetIdent, - exprFactory.getAccumulatorVarName(), - exprFactory.newList(), - exprFactory.newBoolLiteral(true), - mapStep, - exprFactory.newIdentifier(exprFactory.getAccumulatorVarName())); - - // Build call: @__sortBy_input__.@sortByAssociatedKeys(mapCompr) - CelExpr callExpr = - exprFactory.newReceiverCall( - Function.SORT_BY.getFunction(), exprFactory.newIdentifier(SORT_BY_INPUT_VAR), mapCompr); - - // Build bind: cel.bind(@__sortBy_input__, target, callExpr) - CelExpr bindExpr = - exprFactory.fold( - UNUSED_ITER_VAR, - exprFactory.newList(), - SORT_BY_INPUT_VAR, - target, - exprFactory.newBoolLiteral(false), - exprFactory.newIdentifier(SORT_BY_INPUT_VAR), - callExpr); - - return Optional.of(bindExpr); - } - - /** - * Sorts elements of {@code list} based on the natural order of corresponding elements in {@code - * keys}. - * - *

Both {@code list} and {@code keys} must have the exact same size. The sorting is stable - * (i.e., preserves the relative order of elements with equal keys). - * - * @param list The input list to sort - * @param keys The associated keys evaluated for each element in {@code list} - * @return A new {@link ImmutableList} containing the elements of {@code list} sorted by {@code - * keys} - */ - private static ImmutableList sortByAssociatedKeys( - Collection list, Collection keys) { - checkArgument( - list.size() == keys.size(), - "@sortByAssociatedKeys() expected a list of the same size as the associated keys" - + " list, but got %s in list and %s in keys", - list.size(), - keys.size()); - - int listSize = list.size(); - if (listSize == 0) { - return ImmutableList.of(); - } - - Object[] listArray = list.toArray(); - Object[] keysArray = keys.toArray(); - if (listSize == 1) { - OBJECT_COMPARATOR.compare(keysArray[0], keysArray[0]); - return ImmutableList.of(listArray[0]); - } - - Integer[] indices = new Integer[listSize]; - for (int i = 0; i < listSize; i++) { - indices[i] = i; - } - - Arrays.sort(indices, (i1, i2) -> OBJECT_COMPARATOR.compare(keysArray[i1], keysArray[i2])); - - ImmutableList.Builder builder = ImmutableList.builderWithExpectedSize(listSize); - for (int index : indices) { - builder.add(listArray[index]); - } - return builder.build(); + CelListsRuntimeLibrary listsRuntime = + new CelListsRuntimeLibrary( + runtimeEquality, + functions.stream().map(f -> f.runtimeFunction).collect(toImmutableSet())); + runtimeBuilder.addFunctionBindings(listsRuntime.newFunctionBindings()); } } diff --git a/extensions/src/main/java/dev/cel/extensions/CelListsRuntimeLibrary.java b/extensions/src/main/java/dev/cel/extensions/CelListsRuntimeLibrary.java new file mode 100644 index 000000000..ccd2a46d1 --- /dev/null +++ b/extensions/src/main/java/dev/cel/extensions/CelListsRuntimeLibrary.java @@ -0,0 +1,405 @@ +// 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.base.Preconditions; +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableSet; +import com.google.common.collect.Lists; +import com.google.common.collect.Sets; +import com.google.errorprone.annotations.Immutable; +import dev.cel.common.CelOptions; +import dev.cel.common.internal.ComparisonFunctions; +import dev.cel.common.values.CelByteString; +import dev.cel.runtime.CelFunctionBinding; +import dev.cel.runtime.CelLiteRuntimeBuilder; +import dev.cel.runtime.CelLiteRuntimeLibrary; +import dev.cel.runtime.RuntimeEquality; +import dev.cel.runtime.RuntimeHelpers; +import java.util.Arrays; +import java.util.Collection; +import java.util.Comparator; +import java.util.Iterator; +import java.util.List; +import java.util.Set; + +/** Runtime implementation of CEL List extension functions. */ +@Immutable +public final class CelListsRuntimeLibrary implements CelLiteRuntimeLibrary { + + private static final CelObjectComparator OBJECT_COMPARATOR = new CelObjectComparator(); + private static final ImmutableList SORT_BY_KEY_TYPE_NAMES = + ImmutableList.of("int", "uint", "double", "bool", "string", "bytes", "duration", "timestamp"); + + /** Enumeration of functions for List runtime extension. */ + public enum Function { + SLICE("slice"), + FLATTEN("flatten"), + RANGE("lists.range"), + DISTINCT("distinct"), + REVERSE("reverse"), + SORT("sort"), + SORT_BY("@sortByAssociatedKeys"); + + private final String functionName; + + public String getFunction() { + return functionName; + } + + Function(String functionName) { + this.functionName = functionName; + } + } + + private static ImmutableSet getFunctionsForVersion(int version) { + switch (version) { + case 0: + return ImmutableSet.of(Function.SLICE); + case 1: + return ImmutableSet.of(Function.SLICE, Function.FLATTEN); + case 2: + case Integer.MAX_VALUE: + return ImmutableSet.copyOf(Function.values()); + default: + throw new IllegalArgumentException("Unsupported 'lists' extension version " + version); + } + } + + /** + * Returns the latest version of the 'lists' runtime functions using {@link CelOptions#DEFAULT}. + */ + public static CelListsRuntimeLibrary lists() { + return lists(CelOptions.DEFAULT); + } + + /** + * Returns the specified version of the 'lists' runtime functions using {@link + * CelOptions#DEFAULT}. + */ + public static CelListsRuntimeLibrary lists(int version) { + return lists(CelOptions.DEFAULT, version); + } + + /** + * Returns the 'lists' runtime functions with only the specified functions using {@link + * CelOptions#DEFAULT}. + */ + public static CelListsRuntimeLibrary lists(Function... functions) { + return lists(CelOptions.DEFAULT, functions); + } + + /** + * Returns the 'lists' runtime functions with only the specified functions using {@link + * CelOptions#DEFAULT}. + */ + public static CelListsRuntimeLibrary lists(Set functions) { + return lists(CelOptions.DEFAULT, functions); + } + + /** Returns the latest version of the 'lists' runtime functions. */ + public static CelListsRuntimeLibrary lists(CelOptions celOptions) { + return lists(celOptions, Integer.MAX_VALUE); + } + + /** Returns the specified version of the 'lists' runtime functions. */ + public static CelListsRuntimeLibrary lists(CelOptions celOptions, int version) { + return lists(celOptions, getFunctionsForVersion(version)); + } + + /** Returns the 'lists' runtime functions with only the specified functions. */ + public static CelListsRuntimeLibrary lists(CelOptions celOptions, Function... functions) { + return lists(celOptions, ImmutableSet.copyOf(functions)); + } + + /** Returns the 'lists' runtime functions with only the specified functions. */ + public static CelListsRuntimeLibrary lists(CelOptions celOptions, Set functions) { + RuntimeEquality runtimeEquality = RuntimeEquality.create(RuntimeHelpers.create(), celOptions); + return new CelListsRuntimeLibrary(runtimeEquality, functions); + } + + private final RuntimeEquality runtimeEquality; + private final ImmutableSet functions; + + CelListsRuntimeLibrary(RuntimeEquality runtimeEquality, int version) { + this(runtimeEquality, getFunctionsForVersion(version)); + } + + CelListsRuntimeLibrary(RuntimeEquality runtimeEquality, Set functions) { + this.runtimeEquality = runtimeEquality; + this.functions = ImmutableSet.copyOf(functions); + } + + @Override + public void setRuntimeOptions(CelLiteRuntimeBuilder runtimeBuilder) { + runtimeBuilder.addFunctionBindings(newFunctionBindings()); + } + + @SuppressWarnings("unchecked") + public ImmutableSet newFunctionBindings() { + ImmutableSet.Builder bindingBuilder = ImmutableSet.builder(); + for (Function function : functions) { + switch (function) { + case SLICE: + bindingBuilder.addAll( + CelFunctionBinding.fromOverloads( + function.getFunction(), + CelFunctionBinding.from( + "list_slice", + ImmutableList.of(Collection.class, Long.class, Long.class), + (args) -> { + Collection target = (Collection) args[0]; + long from = (Long) args[1]; + long to = (Long) args[2]; + return slice(target, from, to); + }))); + break; + case FLATTEN: + bindingBuilder.addAll( + CelFunctionBinding.fromOverloads( + function.getFunction(), + CelFunctionBinding.from( + "list_flatten", Collection.class, list -> flatten(list, 1)), + CelFunctionBinding.from( + "list_flatten_list_int", + Collection.class, + Long.class, + CelListsRuntimeLibrary::flatten))); + break; + case RANGE: + bindingBuilder.addAll( + CelFunctionBinding.fromOverloads( + function.getFunction(), + CelFunctionBinding.from( + "lists_range", Long.class, CelListsRuntimeLibrary::genRange))); + break; + case DISTINCT: + bindingBuilder.addAll( + CelFunctionBinding.fromOverloads( + function.getFunction(), + CelFunctionBinding.from( + "list_distinct", + Collection.class, + (list) -> distinct(list, runtimeEquality)))); + break; + case REVERSE: + bindingBuilder.addAll( + CelFunctionBinding.fromOverloads( + function.getFunction(), + CelFunctionBinding.from( + "list_reverse", Collection.class, CelListsRuntimeLibrary::reverse))); + break; + case SORT: + bindingBuilder.addAll( + CelFunctionBinding.fromOverloads( + function.getFunction(), + CelFunctionBinding.from( + "list_sort", Collection.class, CelListsRuntimeLibrary::sort))); + break; + case SORT_BY: + for (String typeName : SORT_BY_KEY_TYPE_NAMES) { + bindingBuilder.addAll( + CelFunctionBinding.fromOverloads( + function.getFunction(), + CelFunctionBinding.from( + String.format("list_%s_sortByAssociatedKeys", typeName), + Collection.class, + Collection.class, + CelListsRuntimeLibrary::sortByAssociatedKeys))); + } + break; + } + } + return bindingBuilder.build(); + } + + private static ImmutableList slice(Collection list, long from, long to) { + Preconditions.checkArgument(from >= 0 && to >= 0, "Negative indexes not supported"); + Preconditions.checkArgument(to >= from, "Start index must be less than or equal to end index"); + Preconditions.checkArgument(to <= list.size(), "List is length %s", list.size()); + if (list instanceof List) { + List subList = ((List) list).subList((int) from, (int) to); + if (subList instanceof ImmutableList) { + return (ImmutableList) subList; + } + return ImmutableList.copyOf(subList); + } else { + ImmutableList.Builder builder = ImmutableList.builder(); + long index = 0; + for (Iterator iterator = list.iterator(); iterator.hasNext(); index++) { + Object element = iterator.next(); + if (index >= to) { + break; + } + if (index >= from) { + builder.add(element); + } + } + return builder.build(); + } + } + + @SuppressWarnings("unchecked") + private static ImmutableList flatten(Collection list, long depth) { + Preconditions.checkArgument(depth >= 0, "Level must be non-negative"); + ImmutableList.Builder builder = ImmutableList.builder(); + for (Object element : list) { + if (!(element instanceof Collection) || depth == 0) { + builder.add(element); + } else { + Collection listItem = (Collection) element; + builder.addAll(flatten(listItem, depth - 1)); + } + } + + return builder.build(); + } + + public static ImmutableList genRange(long end) { + checkArgument(end >= 0, "lists.range: size must be non-negative, got %s", end); + checkArgument(end <= 1_000_000, "lists.range: size %s exceeds maximum allowed (1000000)", end); + + ImmutableList.Builder builder = ImmutableList.builderWithExpectedSize((int) end); + for (long i = 0; i < end; i++) { + builder.add(i); + } + return builder.build(); + } + + private static class RuntimeEqualityObjectWrapper { + private final Object object; + private final int hashCode; + private final RuntimeEquality runtimeEquality; + + RuntimeEqualityObjectWrapper(Object object, RuntimeEquality runtimeEquality) { + this.object = object; + this.runtimeEquality = runtimeEquality; + this.hashCode = runtimeEquality.hashCode(object); + } + + @Override + public int hashCode() { + return hashCode; + } + + @Override + public boolean equals(Object obj) { + if (!(obj instanceof RuntimeEqualityObjectWrapper)) { + return false; + } + return runtimeEquality.objectEquals(object, ((RuntimeEqualityObjectWrapper) obj).object); + } + } + + private static ImmutableList distinct( + Collection list, RuntimeEquality runtimeEquality) { + int size = list.size(); + ImmutableList.Builder builder = ImmutableList.builderWithExpectedSize(size); + Set distinctValues = Sets.newHashSetWithExpectedSize(size); + for (Object element : list) { + if (distinctValues.add(new RuntimeEqualityObjectWrapper(element, runtimeEquality))) { + builder.add(element); + } + } + return builder.build(); + } + + private static List reverse(Collection list) { + if (list instanceof List) { + return Lists.reverse((List) list); + } else { + ImmutableList.Builder builder = ImmutableList.builderWithExpectedSize(list.size()); + Object[] objects = list.toArray(); + for (int i = objects.length - 1; i >= 0; i--) { + builder.add(objects[i]); + } + return builder.build(); + } + } + + private static ImmutableList sort(Collection objects) { + if (objects.isEmpty()) { + return ImmutableList.of(); + } + if (objects.size() == 1) { + Object single = objects.iterator().next(); + OBJECT_COMPARATOR.compare(single, single); + return ImmutableList.of(single); + } + return ImmutableList.sortedCopyOf(OBJECT_COMPARATOR, objects); + } + + private static class CelObjectComparator implements Comparator { + + CelObjectComparator() {} + + @SuppressWarnings({"unchecked"}) + @Override + public int compare(Object o1, Object o2) { + if (o1 instanceof Number && o2 instanceof Number) { + return ComparisonFunctions.numericCompare((Number) o1, (Number) o2); + } + if (o1 instanceof CelByteString && o2 instanceof CelByteString) { + return CelByteString.unsignedLexicographicalComparator() + .compare((CelByteString) o1, (CelByteString) o2); + } + + if (!(o1 instanceof Comparable) || !(o2 instanceof Comparable)) { + throw new IllegalArgumentException("List elements must be comparable"); + } + if (o1.getClass() != o2.getClass()) { + throw new IllegalArgumentException("List elements must have the same type"); + } + return ((Comparable) o1).compareTo(o2); + } + } + + private static ImmutableList sortByAssociatedKeys( + Collection list, Collection keys) { + checkArgument( + list.size() == keys.size(), + "@sortByAssociatedKeys() expected a list of the same size as the associated keys" + + " list, but got %s in list and %s in keys", + list.size(), + keys.size()); + + int listSize = list.size(); + if (listSize == 0) { + return ImmutableList.of(); + } + + Object[] listArray = list.toArray(); + Object[] keysArray = keys.toArray(); + if (listSize == 1) { + OBJECT_COMPARATOR.compare(keysArray[0], keysArray[0]); + return ImmutableList.of(listArray[0]); + } + + Integer[] indices = new Integer[listSize]; + for (int i = 0; i < listSize; i++) { + indices[i] = i; + } + + Arrays.sort(indices, (i1, i2) -> OBJECT_COMPARATOR.compare(keysArray[i1], keysArray[i2])); + + ImmutableList.Builder builder = ImmutableList.builderWithExpectedSize(listSize); + for (int index : indices) { + builder.add(listArray[index]); + } + return builder.build(); + } +} diff --git a/extensions/src/test/java/dev/cel/extensions/BUILD.bazel b/extensions/src/test/java/dev/cel/extensions/BUILD.bazel index 8812bd7cc..4002d2879 100644 --- a/extensions/src/test/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/test/java/dev/cel/extensions/BUILD.bazel @@ -37,6 +37,9 @@ java_library( "//extensions:encoders_compiler_library", "//extensions:encoders_runtime_library", "//extensions:extension_library", + "//extensions:lists", + "//extensions:lists_compiler_library", + "//extensions:lists_runtime_library", "//extensions:math", "//extensions:math_compiler_library", "//extensions:math_runtime_library", diff --git a/extensions/src/test/java/dev/cel/extensions/CelListsExtensionsTest.java b/extensions/src/test/java/dev/cel/extensions/CelListsExtensionsTest.java index 73400ccb6..23e701cd0 100644 --- a/extensions/src/test/java/dev/cel/extensions/CelListsExtensionsTest.java +++ b/extensions/src/test/java/dev/cel/extensions/CelListsExtensionsTest.java @@ -16,20 +16,30 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.Assert.assertThrows; +import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; +import com.google.common.collect.ImmutableSet; import com.google.common.collect.ImmutableSortedMultiset; import com.google.common.collect.ImmutableSortedSet; 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.CelContainer; import dev.cel.common.CelValidationException; import dev.cel.common.CelValidationResult; import dev.cel.common.types.SimpleType; +import dev.cel.compiler.CelCompiler; +import dev.cel.compiler.CelCompilerFactory; import dev.cel.expr.conformance.test.SimpleTest; import dev.cel.parser.CelStandardMacro; import dev.cel.runtime.CelEvaluationException; +import dev.cel.runtime.CelLiteRuntime; +import dev.cel.runtime.CelLiteRuntimeFactory; +import dev.cel.runtime.CelRuntime; +import dev.cel.runtime.CelRuntimeBuilder; +import dev.cel.runtime.CelRuntimeFactory; import dev.cel.testing.CelRuntimeFlavor; import dev.cel.validator.CelValidator; import dev.cel.validator.CelValidatorFactory; @@ -381,4 +391,137 @@ public void sortBy_withHomogeneousLiteralValidator_success() throws Exception { assertThat(result.hasError()).isFalse(); assertThat(evalResult).isEqualTo("bar"); } + + @Test + public void separateLibraryAndRuntime_allFunctions_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .setStandardMacros(CelStandardMacro.STANDARD_MACROS) + .addLibraries(CelListsCompilerLibrary.lists()) + .build(); + CelLiteRuntime celLiteRuntime = + CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .addLibraries(CelListsRuntimeLibrary.lists()) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("[1, 2, 3, 4].slice(1, 3)").getAst(); + Object result = celLiteRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableList.of(2L, 3L)); + } + + @Test + public void separateLibraryAndRuntime_versioned_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelListsCompilerLibrary.lists(0)) + .build(); + CelRuntime celRuntime = + CelRuntimeFactory.standardCelRuntimeBuilder() + .addFunctionBindings(CelListsRuntimeLibrary.lists(0).newFunctionBindings()) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("[1, 2, 3, 4].slice(1, 3)").getAst(); + Object result = celRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableList.of(2L, 3L)); + } + + @Test + public void separateLibraryAndRuntime_subsetOfFunctions_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelListsCompilerLibrary.lists(CelListsCompilerLibrary.Function.SLICE)) + .build(); + CelRuntime celRuntime = + CelRuntimeFactory.standardCelRuntimeBuilder() + .addFunctionBindings( + CelListsRuntimeLibrary.lists(CelListsRuntimeLibrary.Function.SLICE) + .newFunctionBindings()) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("[1, 2, 3, 4].slice(1, 3)").getAst(); + Object result = celRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableList.of(2L, 3L)); + assertThrows( + CelValidationException.class, () -> celCompiler.compile("[[1], [2]].flatten()").getAst()); + } + + @Test + public void separateLibraryAndRuntime_setOfFunctions_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries( + CelListsCompilerLibrary.lists( + ImmutableSet.of(CelListsCompilerLibrary.Function.SLICE))) + .build(); + CelRuntime celRuntime = + CelRuntimeFactory.standardCelRuntimeBuilder() + .addFunctionBindings( + CelListsRuntimeLibrary.lists(ImmutableSet.of(CelListsRuntimeLibrary.Function.SLICE)) + .newFunctionBindings()) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("[1, 2, 3, 4].slice(1, 3)").getAst(); + Object result = celRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableList.of(2L, 3L)); + } + + @Test + public void separateLibraryAndRuntime_unsupportedVersion_throws() { + assertThrows(IllegalArgumentException.class, () -> CelListsCompilerLibrary.lists(99)); + assertThrows(IllegalArgumentException.class, () -> CelListsRuntimeLibrary.lists(99)); + } + + @Test + public void lists_subsetOfFunctions_success() throws Exception { + Cel cel = + CelFactory.standardCelBuilder() + .addCompilerLibraries(CelExtensions.lists(CelListsExtensions.Function.SLICE)) + .addRuntimeLibraries(CelExtensions.lists(CelListsExtensions.Function.SLICE)) + .build(); + + CelAbstractSyntaxTree ast = cel.compile("[1, 2, 3, 4].slice(1, 3)").getAst(); + Object result = cel.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableList.of(2L, 3L)); + assertThrows(CelValidationException.class, () -> cel.compile("[[1], [2]].flatten()").getAst()); + } + + @Test + public void lists_setOfFunctions_success() throws Exception { + Cel cel = + CelFactory.standardCelBuilder() + .addCompilerLibraries( + CelExtensions.lists(ImmutableSet.of(CelListsExtensions.Function.SLICE))) + .addRuntimeLibraries( + CelExtensions.lists(ImmutableSet.of(CelListsExtensions.Function.SLICE))) + .build(); + + CelAbstractSyntaxTree ast = cel.compile("[1, 2, 3, 4].slice(1, 3)").getAst(); + Object result = cel.createProgram(ast).eval(); + + assertThat(result).isEqualTo(ImmutableList.of(2L, 3L)); + } + + @Test + public void lists_noArgConstructor_success() { + CelListsExtensions extensions = new CelListsExtensions(); + + assertThat(extensions.version()).isEqualTo(CelListsCompilerLibrary.lists().version()); + assertThat(extensions.functions()).isNotEmpty(); + assertThat(extensions.macros()).isNotEmpty(); + } + + @Test + public void setRuntimeOptions_withoutEquality_throws() { + CelListsExtensions extensions = CelExtensions.lists(); + CelRuntimeBuilder runtimeBuilder = CelRuntimeFactory.standardCelRuntimeBuilder(); + + assertThrows( + UnsupportedOperationException.class, () -> extensions.setRuntimeOptions(runtimeBuilder)); + } } + diff --git a/runtime/src/test/java/dev/cel/runtime/BUILD.bazel b/runtime/src/test/java/dev/cel/runtime/BUILD.bazel index 5b2517b7c..c6562ef6d 100644 --- a/runtime/src/test/java/dev/cel/runtime/BUILD.bazel +++ b/runtime/src/test/java/dev/cel/runtime/BUILD.bazel @@ -195,6 +195,7 @@ cel_android_local_test( "//common/values:cel_value_provider_android", "//common/values:proto_message_lite_value_provider_android", "//extensions:encoders_runtime_library_android", + "//extensions:lists_runtime_library_android", "//extensions:math_runtime_library_android", "//extensions:sets_runtime_library_android", "//extensions:strings_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 426988efe..63e50543a 100644 --- a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java +++ b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java @@ -46,6 +46,7 @@ import dev.cel.expr.conformance.proto3.TestAllTypes; import dev.cel.expr.conformance.proto3.TestAllTypesCelLiteDescriptor; import dev.cel.extensions.CelEncoderRuntimeLibrary; +import dev.cel.extensions.CelListsRuntimeLibrary; import dev.cel.extensions.CelMathRuntimeLibrary; import dev.cel.extensions.CelSetsRuntimeLibrary; import dev.cel.extensions.CelStringRuntimeLibrary; @@ -778,6 +779,18 @@ public void eval_encodersExtension() throws Exception { assertThat(runtime.createProgram(ast).eval()).isEqualTo("aGVsbG8="); } + @Test + public void eval_listsExtension() throws Exception { + CelLiteRuntime runtime = + CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .addLibraries(CelListsRuntimeLibrary.lists()) + .build(); + // Expr: [1, 2, 3].slice(0, 2) + CelAbstractSyntaxTree ast = readCheckedExpr("compiled_lists_slice"); + + assertThat(runtime.createProgram(ast).eval()).isEqualTo(ImmutableList.of(1L, 2L)); + } + private enum CelOptionsTestCase { UNSIGNED_LONG_DISABLED(newBaseTestOptions().enableUnsignedLongs(false).build()), UNWRAP_WKT_DISABLED(newBaseTestOptions().unwrapWellKnownTypesOnFunctionDispatch(false).build()), 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 3f1d569a5..448ee48e0 100644 --- a/testing/src/main/java/dev/cel/testing/compiled/BUILD.bazel +++ b/testing/src/main/java/dev/cel/testing/compiled/BUILD.bazel @@ -50,6 +50,7 @@ java_library( ":compiled_extensions", ":compiled_hello_world", ":compiled_list_literal", + ":compiled_lists_slice", ":compiled_math_greatest", ":compiled_one_plus_two", ":compiled_primitive_variables", @@ -126,6 +127,12 @@ compile_cel( expression = "sets.contains([1, 2], [2])", ) +compile_cel( + name = "compiled_lists_slice", + environment = "//testing/environment:all_extensions", + expression = "[1, 2, 3].slice(0, 2)", +) + compile_cel( name = "compiled_string_lower_ascii", environment = "//testing/environment:all_extensions",