From e901ffa82906fcd1e1e36609a7289cb35ab4f783 Mon Sep 17 00:00:00 2001 From: Dmitri Plotnikov Date: Mon, 28 Sep 2026 20:30:11 -0700 Subject: [PATCH] Split "math" extension into compiler and runtime libraries to allow its usage in Lite runtime PiperOrigin-RevId: 990012647 --- extensions/BUILD.bazel | 10 + .../main/java/dev/cel/extensions/BUILD.bazel | 41 +- .../extensions/CelMathCompilerLibrary.java | 681 +++++++++++ .../dev/cel/extensions/CelMathExtensions.java | 1057 +---------------- .../cel/extensions/CelMathRuntimeLibrary.java | 655 ++++++++++ .../test/java/dev/cel/extensions/BUILD.bazel | 8 +- .../cel/extensions/CelMathExtensionsTest.java | 59 + 7 files changed, 1500 insertions(+), 1011 deletions(-) create mode 100644 extensions/src/main/java/dev/cel/extensions/CelMathCompilerLibrary.java create mode 100644 extensions/src/main/java/dev/cel/extensions/CelMathRuntimeLibrary.java diff --git a/extensions/BUILD.bazel b/extensions/BUILD.bazel index f9c2aee45..36a013dad 100644 --- a/extensions/BUILD.bazel +++ b/extensions/BUILD.bazel @@ -37,6 +37,16 @@ java_library( exports = ["//extensions/src/main/java/dev/cel/extensions:math"], ) +java_library( + name = "math_compiler_library", + exports = ["//extensions/src/main/java/dev/cel/extensions:math_compiler_library"], +) + +java_library( + name = "math_runtime_library", + exports = ["//extensions/src/main/java/dev/cel/extensions:math_runtime_library"], +) + java_library( name = "optional_library", exports = ["//extensions/src/main/java/dev/cel/extensions:optional_library"], diff --git a/extensions/src/main/java/dev/cel/extensions/BUILD.bazel b/extensions/src/main/java/dev/cel/extensions/BUILD.bazel index 4c7d676dd..157516986 100644 --- a/extensions/src/main/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/main/java/dev/cel/extensions/BUILD.bazel @@ -131,17 +131,50 @@ java_library( ], deps = [ ":extension_library", + ":math_compiler_library", + ":math_runtime_library", "//checker:checker_builder", - "//common:compiler_common", + "//common:cel_function_decl", + "//compiler:compiler_builder", + "//parser:macro", + "//parser:parser_builder", + "//runtime", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven//:com_google_guava_guava", + ], +) + +java_library( + name = "math_compiler_library", + srcs = ["CelMathCompilerLibrary.java"], + tags = [ + ], + deps = [ + ":extension_library", + "//checker:checker_builder", + "//common:cel_function_decl", + "//common:cel_issue", + "//common:cel_overload_decl", "//common/ast", - "//common/exceptions:numeric_overflow", - "//common/internal:comparison_functions", "//common/types", "//compiler:compiler_builder", "//parser:macro", "//parser:parser_builder", - "//runtime", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven//:com_google_guava_guava", + ], +) + +java_library( + name = "math_runtime_library", + srcs = ["CelMathRuntimeLibrary.java"], + tags = [ + ], + deps = [ + "//common/exceptions:numeric_overflow", + "//common/internal:comparison_functions", "//runtime:function_binding", + "//runtime:lite_runtime", "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", ], diff --git a/extensions/src/main/java/dev/cel/extensions/CelMathCompilerLibrary.java b/extensions/src/main/java/dev/cel/extensions/CelMathCompilerLibrary.java new file mode 100644 index 000000000..37c9ffe48 --- /dev/null +++ b/extensions/src/main/java/dev/cel/extensions/CelMathCompilerLibrary.java @@ -0,0 +1,681 @@ +// Copyright 2023 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.collect.ImmutableSet.toImmutableSet; + +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.ast.CelConstant; +import dev.cel.common.ast.CelExpr; +import dev.cel.common.ast.CelExpr.ExprKind.Kind; +import dev.cel.common.types.ListType; +import dev.cel.common.types.SimpleType; +import dev.cel.compiler.CelCompilerLibrary; +import dev.cel.parser.CelMacro; +import dev.cel.parser.CelMacroExprFactory; +import dev.cel.parser.CelParserBuilder; +import java.util.List; +import java.util.Optional; +import java.util.Set; + +/** Internal implementation of CEL Math compile-time extensions. */ +@Immutable +public final class CelMathCompilerLibrary + implements CelCompilerLibrary, CelExtensionLibrary.FeatureSet { + + private static final String MATH_NAMESPACE = "math"; + + static final String MATH_MAX_FUNCTION = "math.@max"; + private static final String MATH_MAX_OVERLOAD_DOC = + "Returns the greatest valued number present in the arguments."; + static final String MATH_MIN_FUNCTION = "math.@min"; + private static final String MATH_MIN_OVERLOAD_DOC = + "Returns the least valued number present in the arguments."; + + // Rounding Functions + static final String MATH_CEIL_FUNCTION = "math.ceil"; + static final String MATH_FLOOR_FUNCTION = "math.floor"; + static final String MATH_ROUND_FUNCTION = "math.round"; + static final String MATH_TRUNC_FUNCTION = "math.trunc"; + + // Floating Point Functions + static final String MATH_ISFINITE_FUNCTION = "math.isFinite"; + static final String MATH_ISNAN_FUNCTION = "math.isNaN"; + static final String MATH_ISINF_FUNCTION = "math.isInf"; + + // Signedness Functions + static final String MATH_ABS_FUNCTION = "math.abs"; + static final String MATH_SIGN_FUNCTION = "math.sign"; + + // Bitwise Functions + static final String MATH_BIT_AND_FUNCTION = "math.bitAnd"; + static final String MATH_BIT_OR_FUNCTION = "math.bitOr"; + static final String MATH_BIT_XOR_FUNCTION = "math.bitXor"; + static final String MATH_BIT_NOT_FUNCTION = "math.bitNot"; + static final String MATH_BIT_LEFT_SHIFT_FUNCTION = "math.bitShiftLeft"; + static final String MATH_BIT_RIGHT_SHIFT_FUNCTION = "math.bitShiftRight"; + + static final String MATH_SQRT_FUNCTION = "math.sqrt"; + + /** Enumeration of functions for Math compile-time extension. */ + public enum Function { + MAX( + CelFunctionDecl.newFunctionDeclaration( + MATH_MAX_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_@max_double", MATH_MAX_OVERLOAD_DOC, SimpleType.DOUBLE, SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_@max_int", MATH_MAX_OVERLOAD_DOC, SimpleType.INT, SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_@max_uint", MATH_MAX_OVERLOAD_DOC, SimpleType.UINT, SimpleType.UINT), + CelOverloadDecl.newGlobalOverload( + "math_@max_double_double", + MATH_MAX_OVERLOAD_DOC, + SimpleType.DOUBLE, + SimpleType.DOUBLE, + SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_@max_int_int", + MATH_MAX_OVERLOAD_DOC, + SimpleType.INT, + SimpleType.INT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_@max_uint_uint", + MATH_MAX_OVERLOAD_DOC, + SimpleType.UINT, + SimpleType.UINT, + SimpleType.UINT), + CelOverloadDecl.newGlobalOverload( + "math_@max_int_uint", + MATH_MAX_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.INT, + SimpleType.UINT), + CelOverloadDecl.newGlobalOverload( + "math_@max_int_double", + MATH_MAX_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.INT, + SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_@max_double_int", + MATH_MAX_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.DOUBLE, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_@max_double_uint", + MATH_MAX_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.DOUBLE, + SimpleType.UINT), + CelOverloadDecl.newGlobalOverload( + "math_@max_uint_int", + MATH_MAX_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.UINT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_@max_uint_double", + MATH_MAX_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.UINT, + SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_@max_list_dyn", // Implementation supports double, int and uint as list + // literals. Anything else will error during macro expansion. + MATH_MAX_OVERLOAD_DOC, + SimpleType.DYN, + ListType.create(SimpleType.DYN)))), + MIN( + CelFunctionDecl.newFunctionDeclaration( + MATH_MIN_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_@min_double", MATH_MIN_OVERLOAD_DOC, SimpleType.DOUBLE, SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_@min_int", MATH_MIN_OVERLOAD_DOC, SimpleType.INT, SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_@min_uint", MATH_MIN_OVERLOAD_DOC, SimpleType.UINT, SimpleType.UINT), + CelOverloadDecl.newGlobalOverload( + "math_@min_double_double", + MATH_MIN_OVERLOAD_DOC, + SimpleType.DOUBLE, + SimpleType.DOUBLE, + SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_@min_int_int", + MATH_MIN_OVERLOAD_DOC, + SimpleType.INT, + SimpleType.INT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_@min_uint_uint", + MATH_MIN_OVERLOAD_DOC, + SimpleType.UINT, + SimpleType.UINT, + SimpleType.UINT), + CelOverloadDecl.newGlobalOverload( + "math_@min_int_uint", + MATH_MIN_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.INT, + SimpleType.UINT), + CelOverloadDecl.newGlobalOverload( + "math_@min_int_double", + MATH_MIN_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.INT, + SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_@min_double_int", + MATH_MIN_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.DOUBLE, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_@min_double_uint", + MATH_MIN_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.DOUBLE, + SimpleType.UINT), + CelOverloadDecl.newGlobalOverload( + "math_@min_uint_int", + MATH_MIN_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.UINT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_@min_uint_double", + MATH_MIN_OVERLOAD_DOC, + SimpleType.DYN, + SimpleType.UINT, + SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_@min_list_dyn", // Implementation supports double, int and uint as list + // literals. Anything else will error during macro expansion. + MATH_MIN_OVERLOAD_DOC, + SimpleType.DYN, + ListType.create(SimpleType.DYN)))), + CEIL( + CelFunctionDecl.newFunctionDeclaration( + MATH_CEIL_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_ceil_double", + "Compute the ceiling of a double value.", + SimpleType.DOUBLE, + SimpleType.DOUBLE))), + FLOOR( + CelFunctionDecl.newFunctionDeclaration( + MATH_FLOOR_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_floor_double", + "Compute the floor of a double value.", + SimpleType.DOUBLE, + SimpleType.DOUBLE))), + ROUND( + CelFunctionDecl.newFunctionDeclaration( + MATH_ROUND_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_round_double", + "Rounds the double value to the nearest whole number with ties rounding away from" + + " zero.", + SimpleType.DOUBLE, + SimpleType.DOUBLE))), + TRUNC( + CelFunctionDecl.newFunctionDeclaration( + MATH_TRUNC_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_trunc_double", + "Truncates the fractional portion of the double value.", + SimpleType.DOUBLE, + SimpleType.DOUBLE))), + ISFINITE( + CelFunctionDecl.newFunctionDeclaration( + MATH_ISFINITE_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_isFinite_double", + "Returns true if the value is a finite number.", + SimpleType.BOOL, + SimpleType.DOUBLE))), + ISNAN( + CelFunctionDecl.newFunctionDeclaration( + MATH_ISNAN_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_isNaN_double", + "Returns true if the input double value is NaN, false otherwise.", + SimpleType.BOOL, + SimpleType.DOUBLE))), + ISINF( + CelFunctionDecl.newFunctionDeclaration( + MATH_ISINF_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_isInf_double", + "Returns true if the input double value is -Inf or +Inf.", + SimpleType.BOOL, + SimpleType.DOUBLE))), + ABS( + CelFunctionDecl.newFunctionDeclaration( + MATH_ABS_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_abs_double", + "Compute the absolute value of a double value.", + SimpleType.DOUBLE, + SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_abs_int", + "Compute the absolute value of an int value.", + SimpleType.INT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_abs_uint", + "Compute the absolute value of a uint value.", + SimpleType.UINT, + SimpleType.UINT))), + SIGN( + CelFunctionDecl.newFunctionDeclaration( + MATH_SIGN_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_sign_double", + "Returns the sign of the input numeric type, either -1, 0, 1 cast as double.", + SimpleType.DOUBLE, + SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_sign_uint", + "Returns the sign of the input numeric type, either -1, 0, 1 case as uint.", + SimpleType.UINT, + SimpleType.UINT), + CelOverloadDecl.newGlobalOverload( + "math_sign_int", + "Returns the sign of the input numeric type, either -1, 0, 1.", + SimpleType.INT, + SimpleType.INT))), + BITAND( + CelFunctionDecl.newFunctionDeclaration( + MATH_BIT_AND_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_bitAnd_int_int", + "Performs a bitwise-AND operation over two int values.", + SimpleType.INT, + SimpleType.INT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_bitAnd_uint_uint", + "Performs a bitwise-AND operation over two uint values.", + SimpleType.UINT, + SimpleType.UINT, + SimpleType.UINT))), + BITOR( + CelFunctionDecl.newFunctionDeclaration( + MATH_BIT_OR_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_bitOr_int_int", + "Performs a bitwise-OR operation over two int values.", + SimpleType.INT, + SimpleType.INT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_bitOr_uint_uint", + "Performs a bitwise-OR operation over two uint values.", + SimpleType.UINT, + SimpleType.UINT, + SimpleType.UINT))), + BITXOR( + CelFunctionDecl.newFunctionDeclaration( + MATH_BIT_XOR_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_bitXor_int_int", + "Performs a bitwise-XOR operation over two int values.", + SimpleType.INT, + SimpleType.INT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_bitXor_uint_uint", + "Performs a bitwise-XOR operation over two uint values.", + SimpleType.UINT, + SimpleType.UINT, + SimpleType.UINT))), + BITNOT( + CelFunctionDecl.newFunctionDeclaration( + MATH_BIT_NOT_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_bitNot_int_int", + "Performs a bitwise-NOT operation over two int values.", + SimpleType.INT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_bitNot_uint_uint", + "Performs a bitwise-NOT operation over two uint values.", + SimpleType.UINT, + SimpleType.UINT))), + BITSHIFTLEFT( + CelFunctionDecl.newFunctionDeclaration( + MATH_BIT_LEFT_SHIFT_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_bitShiftLeft_int_int", + "Performs a bitwise-SHIFTLEFT operation over two int values.", + SimpleType.INT, + SimpleType.INT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_bitShiftLeft_uint_int", + "Performs a bitwise-SHIFTLEFT operation over two uint values.", + SimpleType.UINT, + SimpleType.UINT, + SimpleType.INT))), + BITSHIFTRIGHT( + CelFunctionDecl.newFunctionDeclaration( + MATH_BIT_RIGHT_SHIFT_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_bitShiftRight_int_int", + "Performs a bitwise-SHIFTRIGHT operation over two int values.", + SimpleType.INT, + SimpleType.INT, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_bitShiftRight_uint_int", + "Performs a bitwise-SHIFTRIGHT operation over two uint values.", + SimpleType.UINT, + SimpleType.UINT, + SimpleType.INT))), + SQRT( + CelFunctionDecl.newFunctionDeclaration( + MATH_SQRT_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "math_sqrt_double", + "Computes square root of the double value.", + SimpleType.DOUBLE, + SimpleType.DOUBLE), + CelOverloadDecl.newGlobalOverload( + "math_sqrt_int", + "Computes square root of the int value.", + SimpleType.DOUBLE, + SimpleType.INT), + CelOverloadDecl.newGlobalOverload( + "math_sqrt_uint", + "Computes square root of the unsigned value.", + SimpleType.DOUBLE, + SimpleType.UINT))); + + private final CelFunctionDecl functionDecl; + + public String getFunction() { + return functionDecl.name(); + } + + public CelFunctionDecl getFunctionDecl() { + return functionDecl; + } + + Function(CelFunctionDecl functionDecl) { + this.functionDecl = functionDecl; + } + } + + private static final class Library implements CelExtensionLibrary { + private final CelMathCompilerLibrary version0; + private final CelMathCompilerLibrary version1; + private final CelMathCompilerLibrary version2; + + Library() { + version0 = new CelMathCompilerLibrary(0, ImmutableSet.of(Function.MIN, Function.MAX)); + + version1 = + new CelMathCompilerLibrary( + 1, + ImmutableSet.builder() + .addAll(version0.functions) + .add( + Function.CEIL, + Function.FLOOR, + Function.ROUND, + Function.TRUNC, + Function.ISINF, + Function.ISNAN, + Function.ISFINITE, + Function.ABS, + Function.SIGN, + Function.BITAND, + Function.BITOR, + Function.BITXOR, + Function.BITNOT, + Function.BITSHIFTLEFT, + Function.BITSHIFTRIGHT) + .build()); + + version2 = + new CelMathCompilerLibrary( + 2, + ImmutableSet.builder() + .addAll(version1.functions) + .add(Function.SQRT) + .build()); + } + + @Override + public String name() { + return "math"; + } + + @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 'math' compiler extension. */ + public static CelMathCompilerLibrary math() { + return library().latest(); + } + + /** Returns the specified version of the 'math' compiler extension. */ + public static CelMathCompilerLibrary math(int version) { + return library().version(version); + } + + /** Returns the 'math' compiler extension with only the specified functions. */ + public static CelMathCompilerLibrary math(Function... functions) { + return math(ImmutableSet.copyOf(functions)); + } + + /** Returns the 'math' compiler extension with only the specified functions. */ + public static CelMathCompilerLibrary math(Set functions) { + return new CelMathCompilerLibrary(functions); + } + + private final ImmutableSet functions; + private final int version; + + CelMathCompilerLibrary(Set functions) { + this(-1, functions); + } + + private CelMathCompilerLibrary(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() { + return ImmutableSet.of( + CelMacro.newReceiverVarArgMacro("greatest", CelMathCompilerLibrary::expandGreatestMacro), + CelMacro.newReceiverVarArgMacro("least", CelMathCompilerLibrary::expandLeastMacro)); + } + + @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 expandGreatestMacro( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { + if (!isTargetInNamespace(target)) { + // Return empty to indicate that we're not interested in expanding this macro, and + // that the parser should default to a function call on the receiver. + return Optional.empty(); + } + + switch (arguments.size()) { + case 0: + return newError(exprFactory, "math.greatest() requires at least one argument", target); + case 1: + Optional invalidArg = + checkInvalidArgumentSingleArg(exprFactory, "math.greatest()", arguments.get(0)); + if (invalidArg.isPresent()) { + return invalidArg; + } + + return Optional.of(exprFactory.newGlobalCall(MATH_MAX_FUNCTION, arguments.get(0))); + case 2: + invalidArg = checkInvalidArgument(exprFactory, "math.greatest()", arguments); + if (invalidArg.isPresent()) { + return invalidArg; + } + + return Optional.of(exprFactory.newGlobalCall(MATH_MAX_FUNCTION, arguments)); + default: + invalidArg = checkInvalidArgument(exprFactory, "math.greatest()", arguments); + if (invalidArg.isPresent()) { + return invalidArg; + } + + return Optional.of( + exprFactory.newGlobalCall(MATH_MAX_FUNCTION, exprFactory.newList(arguments))); + } + } + + private static Optional expandLeastMacro( + CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { + if (!isTargetInNamespace(target)) { + // Return empty to indicate that we're not interested in expanding this macro, and + // that the parser should default to a function call on the receiver. + return Optional.empty(); + } + + switch (arguments.size()) { + case 0: + return newError(exprFactory, "math.least() requires at least one argument", target); + case 1: + Optional invalidArg = + checkInvalidArgumentSingleArg(exprFactory, "math.least()", arguments.get(0)); + if (invalidArg.isPresent()) { + return invalidArg; + } + + return Optional.of(exprFactory.newGlobalCall(MATH_MIN_FUNCTION, arguments.get(0))); + case 2: + invalidArg = checkInvalidArgument(exprFactory, "math.least()", arguments); + if (invalidArg.isPresent()) { + return invalidArg; + } + + return Optional.of(exprFactory.newGlobalCall(MATH_MIN_FUNCTION, arguments)); + default: + invalidArg = checkInvalidArgument(exprFactory, "math.least()", arguments); + if (invalidArg.isPresent()) { + return invalidArg; + } + + return Optional.of( + exprFactory.newGlobalCall(MATH_MIN_FUNCTION, exprFactory.newList(arguments))); + } + } + + private static boolean isTargetInNamespace(CelExpr target) { + return target.exprKind().getKind().equals(Kind.IDENT) + && target.ident().name().equals(MATH_NAMESPACE); + } + + private static Optional checkInvalidArgument( + CelMacroExprFactory exprFactory, String functionName, List arguments) { + + for (CelExpr arg : arguments) { + if (!isArgumentValidType(arg)) { + return newError( + exprFactory, + String.format("%s simple literal arguments must be numeric", functionName), + arg); + } + } + return Optional.empty(); + } + + private static Optional checkInvalidArgumentSingleArg( + CelMacroExprFactory exprFactory, String functionName, CelExpr argument) { + if (argument.exprKind().getKind() == Kind.LIST) { + if (argument.list().elements().isEmpty()) { + return newError( + exprFactory, String.format("%s invalid single argument value", functionName), argument); + } + + return checkInvalidArgument(exprFactory, functionName, argument.list().elements()); + } + if (isArgumentValidType(argument)) { + return Optional.empty(); + } + + return newError( + exprFactory, String.format("%s invalid single argument value", functionName), argument); + } + + private static boolean isArgumentValidType(CelExpr argument) { + if (argument.exprKind().getKind() == Kind.CONSTANT) { + CelConstant constant = argument.constant(); + return constant.getKind() == CelConstant.Kind.INT64_VALUE + || constant.getKind() == CelConstant.Kind.UINT64_VALUE + || constant.getKind() == CelConstant.Kind.DOUBLE_VALUE; + } else if (argument.exprKind().getKind().equals(Kind.LIST) + || argument.exprKind().getKind().equals(Kind.STRUCT) + || argument.exprKind().getKind().equals(Kind.MAP)) { + return false; + } + + return true; + } + + private static Optional newError( + CelMacroExprFactory exprFactory, String errorMessage, CelExpr argument) { + return Optional.of( + exprFactory.reportError( + CelIssue.formatError(exprFactory.getSourceLocation(argument), errorMessage))); + } +} diff --git a/extensions/src/main/java/dev/cel/extensions/CelMathExtensions.java b/extensions/src/main/java/dev/cel/extensions/CelMathExtensions.java index 63108aa0c..f57cd772e 100644 --- a/extensions/src/main/java/dev/cel/extensions/CelMathExtensions.java +++ b/extensions/src/main/java/dev/cel/extensions/CelMathExtensions.java @@ -14,39 +14,18 @@ package dev.cel.extensions; -import static com.google.common.collect.Comparators.max; -import static com.google.common.collect.Comparators.min; import static com.google.common.collect.ImmutableSet.toImmutableSet; -import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableSet; -import com.google.common.collect.ImmutableTable; -import com.google.common.math.DoubleMath; -import com.google.common.primitives.UnsignedLong; 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.ast.CelConstant; -import dev.cel.common.ast.CelExpr; -import dev.cel.common.ast.CelExpr.ExprKind.Kind; -import dev.cel.common.exceptions.CelNumericOverflowException; -import dev.cel.common.internal.ComparisonFunctions; -import dev.cel.common.types.ListType; -import dev.cel.common.types.SimpleType; 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.CelRuntimeBuilder; import dev.cel.runtime.CelRuntimeLibrary; -import java.math.RoundingMode; -import java.util.List; -import java.util.Optional; import java.util.Set; -import java.util.function.BiFunction; /** * Internal implementation of Math Extensions @@ -54,673 +33,67 @@ *

Note: For equal numbers with different types, the result is always the first argument e.g.: * math.greatest(1u, 1.0) -> 1u */ -@SuppressWarnings({"rawtypes", "unchecked"}) // Use of raw Comparables. @Immutable public final class CelMathExtensions implements CelCompilerLibrary, CelRuntimeLibrary, CelExtensionLibrary.FeatureSet { - private static final String MATH_NAMESPACE = "math"; - - private static final String MATH_MAX_FUNCTION = "math.@max"; - private static final String MATH_MAX_OVERLOAD_DOC = - "Returns the greatest valued number present in the arguments."; - private static final String MATH_MIN_FUNCTION = "math.@min"; - private static final String MATH_MIN_OVERLOAD_DOC = - "Returns the least valued number present in the arguments."; - - // Rounding Functions - private static final String MATH_CEIL_FUNCTION = "math.ceil"; - private static final String MATH_FLOOR_FUNCTION = "math.floor"; - private static final String MATH_ROUND_FUNCTION = "math.round"; - private static final String MATH_TRUNC_FUNCTION = "math.trunc"; - - // Floating Point Functions - private static final String MATH_ISFINITE_FUNCTION = "math.isFinite"; - private static final String MATH_ISNAN_FUNCTION = "math.isNaN"; - private static final String MATH_ISINF_FUNCTION = "math.isInf"; - - // Signedness Functions - private static final String MATH_ABS_FUNCTION = "math.abs"; - private static final String MATH_SIGN_FUNCTION = "math.sign"; - - // Bitwise Functions - private static final String MATH_BIT_AND_FUNCTION = "math.bitAnd"; - private static final String MATH_BIT_OR_FUNCTION = "math.bitOr"; - private static final String MATH_BIT_XOR_FUNCTION = "math.bitXor"; - private static final String MATH_BIT_NOT_FUNCTION = "math.bitNot"; - private static final String MATH_BIT_LEFT_SHIFT_FUNCTION = "math.bitShiftLeft"; - private static final String MATH_BIT_RIGHT_SHIFT_FUNCTION = "math.bitShiftRight"; - - private static final String MATH_SQRT_FUNCTION = "math.sqrt"; - - private static final int MAX_BIT_SHIFT = 63; - - /** - * Returns the proper comparison function to use for a math function call involving different - * argument types. - * - *

Example: (uint, int) -> {@link ComparisonFunctions#compareUintInt(UnsignedLong, long)} - */ - private static final ImmutableTable> - CLASSES_TO_COMPARATORS = newComparatorTable(); - - private static ImmutableTable> - newComparatorTable() { - ImmutableTable.Builder> builder = - new ImmutableTable.Builder<>(); - builder.put( - Long.class, - Double.class, - (x, y) -> ComparisonFunctions.compareIntDouble((Long) x, (Double) y)); - builder.put( - Double.class, - Long.class, - (x, y) -> ComparisonFunctions.compareDoubleInt((Double) x, (Long) y)); - builder.put( - Double.class, - UnsignedLong.class, - (x, y) -> ComparisonFunctions.compareDoubleUint((Double) x, (UnsignedLong) y)); - builder.put( - UnsignedLong.class, - Double.class, - (x, y) -> ComparisonFunctions.compareUintDouble((UnsignedLong) x, (Double) y)); - builder.put( - Long.class, - UnsignedLong.class, - (x, y) -> ComparisonFunctions.compareIntUint((Long) x, (UnsignedLong) y)); - builder.put( - UnsignedLong.class, - Long.class, - (x, y) -> ComparisonFunctions.compareUintInt((UnsignedLong) x, (Long) y)); - return builder.buildOrThrow(); - } - /** Enumeration of functions for Math extension. */ public enum Function { - MAX( - CelFunctionDecl.newFunctionDeclaration( - MATH_MAX_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_@max_double", MATH_MAX_OVERLOAD_DOC, SimpleType.DOUBLE, SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_@max_int", MATH_MAX_OVERLOAD_DOC, SimpleType.INT, SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_@max_uint", MATH_MAX_OVERLOAD_DOC, SimpleType.UINT, SimpleType.UINT), - CelOverloadDecl.newGlobalOverload( - "math_@max_double_double", - MATH_MAX_OVERLOAD_DOC, - SimpleType.DOUBLE, - SimpleType.DOUBLE, - SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_@max_int_int", - MATH_MAX_OVERLOAD_DOC, - SimpleType.INT, - SimpleType.INT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_@max_uint_uint", - MATH_MAX_OVERLOAD_DOC, - SimpleType.UINT, - SimpleType.UINT, - SimpleType.UINT), - CelOverloadDecl.newGlobalOverload( - "math_@max_int_uint", - MATH_MAX_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.INT, - SimpleType.UINT), - CelOverloadDecl.newGlobalOverload( - "math_@max_int_double", - MATH_MAX_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.INT, - SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_@max_double_int", - MATH_MAX_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.DOUBLE, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_@max_double_uint", - MATH_MAX_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.DOUBLE, - SimpleType.UINT), - CelOverloadDecl.newGlobalOverload( - "math_@max_uint_int", - MATH_MAX_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.UINT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_@max_uint_double", - MATH_MAX_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.UINT, - SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_@max_list_dyn", // Implementation supports double, int and uint as list - // literals. Anything else will error during macro expansion. - MATH_MAX_OVERLOAD_DOC, - SimpleType.DYN, - ListType.create(SimpleType.DYN))), - ImmutableSet.builder() - .add(CelFunctionBinding.from("math_@max_double", Double.class, x -> x)) - .add(CelFunctionBinding.from("math_@max_int", Long.class, x -> x)) - .add( - CelFunctionBinding.from( - "math_@max_double_double", - Double.class, - Double.class, - CelMathExtensions::maxPair)) - .add( - CelFunctionBinding.from( - "math_@max_int_int", Long.class, Long.class, CelMathExtensions::maxPair)) - .add( - CelFunctionBinding.from( - "math_@max_int_double", Long.class, Double.class, CelMathExtensions::maxPair)) - .add( - CelFunctionBinding.from( - "math_@max_double_int", Double.class, Long.class, CelMathExtensions::maxPair)) - .add( - CelFunctionBinding.from( - "math_@max_list_dyn", List.class, CelMathExtensions::maxList)) - .add(CelFunctionBinding.from("math_@max_uint", UnsignedLong.class, x -> x)) - .add( - CelFunctionBinding.from( - "math_@max_uint_uint", - UnsignedLong.class, - UnsignedLong.class, - CelMathExtensions::maxPair)) - .add( - CelFunctionBinding.from( - "math_@max_double_uint", - Double.class, - UnsignedLong.class, - CelMathExtensions::maxPair)) - .add( - CelFunctionBinding.from( - "math_@max_uint_int", - UnsignedLong.class, - Long.class, - CelMathExtensions::maxPair)) - .add( - CelFunctionBinding.from( - "math_@max_uint_double", - UnsignedLong.class, - Double.class, - CelMathExtensions::maxPair)) - .add( - CelFunctionBinding.from( - "math_@max_int_uint", - Long.class, - UnsignedLong.class, - CelMathExtensions::maxPair)) - .build()), - MIN( - CelFunctionDecl.newFunctionDeclaration( - MATH_MIN_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_@min_double", MATH_MIN_OVERLOAD_DOC, SimpleType.DOUBLE, SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_@min_int", MATH_MIN_OVERLOAD_DOC, SimpleType.INT, SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_@min_uint", MATH_MIN_OVERLOAD_DOC, SimpleType.UINT, SimpleType.UINT), - CelOverloadDecl.newGlobalOverload( - "math_@min_double_double", - MATH_MIN_OVERLOAD_DOC, - SimpleType.DOUBLE, - SimpleType.DOUBLE, - SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_@min_int_int", - MATH_MIN_OVERLOAD_DOC, - SimpleType.INT, - SimpleType.INT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_@min_uint_uint", - MATH_MIN_OVERLOAD_DOC, - SimpleType.UINT, - SimpleType.UINT, - SimpleType.UINT), - CelOverloadDecl.newGlobalOverload( - "math_@min_int_uint", - MATH_MIN_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.INT, - SimpleType.UINT), - CelOverloadDecl.newGlobalOverload( - "math_@min_int_double", - MATH_MIN_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.INT, - SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_@min_double_int", - MATH_MIN_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.DOUBLE, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_@min_double_uint", - MATH_MIN_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.DOUBLE, - SimpleType.UINT), - CelOverloadDecl.newGlobalOverload( - "math_@min_uint_int", - MATH_MIN_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.UINT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_@min_uint_double", - MATH_MIN_OVERLOAD_DOC, - SimpleType.DYN, - SimpleType.UINT, - SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_@min_list_dyn", // Implementation supports double, int and uint as list - // literals. Anything else will error during macro expansion. - MATH_MIN_OVERLOAD_DOC, - SimpleType.DYN, - ListType.create(SimpleType.DYN))), - ImmutableSet.builder() - .add(CelFunctionBinding.from("math_@min_double", Double.class, x -> x)) - .add(CelFunctionBinding.from("math_@min_int", Long.class, x -> x)) - .add( - CelFunctionBinding.from( - "math_@min_double_double", - Double.class, - Double.class, - CelMathExtensions::minPair)) - .add( - CelFunctionBinding.from( - "math_@min_int_int", Long.class, Long.class, CelMathExtensions::minPair)) - .add( - CelFunctionBinding.from( - "math_@min_int_double", Long.class, Double.class, CelMathExtensions::minPair)) - .add( - CelFunctionBinding.from( - "math_@min_double_int", Double.class, Long.class, CelMathExtensions::minPair)) - .add( - CelFunctionBinding.from( - "math_@min_list_dyn", List.class, CelMathExtensions::minList)) - .add(CelFunctionBinding.from("math_@min_uint", UnsignedLong.class, x -> x)) - .add( - CelFunctionBinding.from( - "math_@min_uint_uint", - UnsignedLong.class, - UnsignedLong.class, - CelMathExtensions::minPair)) - .add( - CelFunctionBinding.from( - "math_@min_double_uint", - Double.class, - UnsignedLong.class, - CelMathExtensions::minPair)) - .add( - CelFunctionBinding.from( - "math_@min_uint_int", - UnsignedLong.class, - Long.class, - CelMathExtensions::minPair)) - .add( - CelFunctionBinding.from( - "math_@min_uint_double", - UnsignedLong.class, - Double.class, - CelMathExtensions::minPair)) - .add( - CelFunctionBinding.from( - "math_@min_int_uint", - Long.class, - UnsignedLong.class, - CelMathExtensions::minPair)) - .build()), - CEIL( - CelFunctionDecl.newFunctionDeclaration( - MATH_CEIL_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_ceil_double", - "Compute the ceiling of a double value.", - SimpleType.DOUBLE, - SimpleType.DOUBLE)), - ImmutableSet.of(CelFunctionBinding.from("math_ceil_double", Double.class, Math::ceil))), - FLOOR( - CelFunctionDecl.newFunctionDeclaration( - MATH_FLOOR_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_floor_double", - "Compute the floor of a double value.", - SimpleType.DOUBLE, - SimpleType.DOUBLE)), - ImmutableSet.of(CelFunctionBinding.from("math_floor_double", Double.class, Math::floor))), - ROUND( - CelFunctionDecl.newFunctionDeclaration( - MATH_ROUND_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_round_double", - "Rounds the double value to the nearest whole number with ties rounding away from" - + " zero.", - SimpleType.DOUBLE, - SimpleType.DOUBLE)), - ImmutableSet.of( - CelFunctionBinding.from("math_round_double", Double.class, CelMathExtensions::round))), - TRUNC( - CelFunctionDecl.newFunctionDeclaration( - MATH_TRUNC_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_trunc_double", - "Truncates the fractional portion of the double value.", - SimpleType.DOUBLE, - SimpleType.DOUBLE)), - ImmutableSet.of( - CelFunctionBinding.from("math_trunc_double", Double.class, CelMathExtensions::trunc))), - ISFINITE( - CelFunctionDecl.newFunctionDeclaration( - MATH_ISFINITE_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_isFinite_double", - "Returns true if the value is a finite number.", - SimpleType.BOOL, - SimpleType.DOUBLE)), - ImmutableSet.of( - CelFunctionBinding.from("math_isFinite_double", Double.class, Double::isFinite))), - ISNAN( - CelFunctionDecl.newFunctionDeclaration( - MATH_ISNAN_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_isNaN_double", - "Returns true if the input double value is NaN, false otherwise.", - SimpleType.BOOL, - SimpleType.DOUBLE)), - ImmutableSet.of( - CelFunctionBinding.from("math_isNaN_double", Double.class, CelMathExtensions::isNaN))), - ISINF( - CelFunctionDecl.newFunctionDeclaration( - MATH_ISINF_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_isInf_double", - "Returns true if the input double value is -Inf or +Inf.", - SimpleType.BOOL, - SimpleType.DOUBLE)), - ImmutableSet.of( - CelFunctionBinding.from( - "math_isInf_double", Double.class, CelMathExtensions::isInfinite))), - ABS( - CelFunctionDecl.newFunctionDeclaration( - MATH_ABS_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_abs_double", - "Compute the absolute value of a double value.", - SimpleType.DOUBLE, - SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_abs_int", - "Compute the absolute value of an int value.", - SimpleType.INT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_abs_uint", - "Compute the absolute value of a uint value.", - SimpleType.UINT, - SimpleType.UINT)), - ImmutableSet.of( - CelFunctionBinding.from("math_abs_double", Double.class, Math::abs), - CelFunctionBinding.from("math_abs_int", Long.class, CelMathExtensions::absExact), - CelFunctionBinding.from("math_abs_uint", UnsignedLong.class, x -> x))), - SIGN( - CelFunctionDecl.newFunctionDeclaration( - MATH_SIGN_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_sign_double", - "Returns the sign of the input numeric type, either -1, 0, 1 cast as double.", - SimpleType.DOUBLE, - SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_sign_uint", - "Returns the sign of the input numeric type, either -1, 0, 1 case as uint.", - SimpleType.UINT, - SimpleType.UINT), - CelOverloadDecl.newGlobalOverload( - "math_sign_int", - "Returns the sign of the input numeric type, either -1, 0, 1.", - SimpleType.INT, - SimpleType.INT)), - ImmutableSet.of( - CelFunctionBinding.from("math_sign_double", Double.class, CelMathExtensions::sign), - CelFunctionBinding.from("math_sign_int", Long.class, CelMathExtensions::sign), - CelFunctionBinding.from( - "math_sign_uint", UnsignedLong.class, CelMathExtensions::sign))), - BITAND( - CelFunctionDecl.newFunctionDeclaration( - MATH_BIT_AND_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_bitAnd_int_int", - "Performs a bitwise-AND operation over two int values.", - SimpleType.INT, - SimpleType.INT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_bitAnd_uint_uint", - "Performs a bitwise-AND operation over two uint values.", - SimpleType.UINT, - SimpleType.UINT, - SimpleType.UINT)), - ImmutableSet.of( - CelFunctionBinding.from( - "math_bitAnd_int_int", Long.class, Long.class, CelMathExtensions::intBitAnd), - CelFunctionBinding.from( - "math_bitAnd_uint_uint", - UnsignedLong.class, - UnsignedLong.class, - CelMathExtensions::uintBitAnd))), - BITOR( - CelFunctionDecl.newFunctionDeclaration( - MATH_BIT_OR_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_bitOr_int_int", - "Performs a bitwise-OR operation over two int values.", - SimpleType.INT, - SimpleType.INT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_bitOr_uint_uint", - "Performs a bitwise-OR operation over two uint values.", - SimpleType.UINT, - SimpleType.UINT, - SimpleType.UINT)), - ImmutableSet.of( - CelFunctionBinding.from( - "math_bitOr_int_int", Long.class, Long.class, CelMathExtensions::intBitOr), - CelFunctionBinding.from( - "math_bitOr_uint_uint", - UnsignedLong.class, - UnsignedLong.class, - CelMathExtensions::uintBitOr))), - BITXOR( - CelFunctionDecl.newFunctionDeclaration( - MATH_BIT_XOR_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_bitXor_int_int", - "Performs a bitwise-XOR operation over two int values.", - SimpleType.INT, - SimpleType.INT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_bitXor_uint_uint", - "Performs a bitwise-XOR operation over two uint values.", - SimpleType.UINT, - SimpleType.UINT, - SimpleType.UINT)), - ImmutableSet.of( - CelFunctionBinding.from( - "math_bitXor_int_int", Long.class, Long.class, CelMathExtensions::intBitXor), - CelFunctionBinding.from( - "math_bitXor_uint_uint", - UnsignedLong.class, - UnsignedLong.class, - CelMathExtensions::uintBitXor))), - BITNOT( - CelFunctionDecl.newFunctionDeclaration( - MATH_BIT_NOT_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_bitNot_int_int", - "Performs a bitwise-NOT operation over two int values.", - SimpleType.INT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_bitNot_uint_uint", - "Performs a bitwise-NOT operation over two uint values.", - SimpleType.UINT, - SimpleType.UINT)), - ImmutableSet.of( - CelFunctionBinding.from( - "math_bitNot_int_int", Long.class, CelMathExtensions::intBitNot), - CelFunctionBinding.from( - "math_bitNot_uint_uint", UnsignedLong.class, CelMathExtensions::uintBitNot))), + MAX(CelMathCompilerLibrary.Function.MAX, CelMathRuntimeLibrary.Function.MAX), + MIN(CelMathCompilerLibrary.Function.MIN, CelMathRuntimeLibrary.Function.MIN), + CEIL(CelMathCompilerLibrary.Function.CEIL, CelMathRuntimeLibrary.Function.CEIL), + FLOOR(CelMathCompilerLibrary.Function.FLOOR, CelMathRuntimeLibrary.Function.FLOOR), + ROUND(CelMathCompilerLibrary.Function.ROUND, CelMathRuntimeLibrary.Function.ROUND), + TRUNC(CelMathCompilerLibrary.Function.TRUNC, CelMathRuntimeLibrary.Function.TRUNC), + ISFINITE(CelMathCompilerLibrary.Function.ISFINITE, CelMathRuntimeLibrary.Function.ISFINITE), + ISNAN(CelMathCompilerLibrary.Function.ISNAN, CelMathRuntimeLibrary.Function.ISNAN), + ISINF(CelMathCompilerLibrary.Function.ISINF, CelMathRuntimeLibrary.Function.ISINF), + ABS(CelMathCompilerLibrary.Function.ABS, CelMathRuntimeLibrary.Function.ABS), + SIGN(CelMathCompilerLibrary.Function.SIGN, CelMathRuntimeLibrary.Function.SIGN), + BITAND(CelMathCompilerLibrary.Function.BITAND, CelMathRuntimeLibrary.Function.BITAND), + BITOR(CelMathCompilerLibrary.Function.BITOR, CelMathRuntimeLibrary.Function.BITOR), + BITXOR(CelMathCompilerLibrary.Function.BITXOR, CelMathRuntimeLibrary.Function.BITXOR), + BITNOT(CelMathCompilerLibrary.Function.BITNOT, CelMathRuntimeLibrary.Function.BITNOT), BITSHIFTLEFT( - CelFunctionDecl.newFunctionDeclaration( - MATH_BIT_LEFT_SHIFT_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_bitShiftLeft_int_int", - "Performs a bitwise-SHIFTLEFT operation over two int values.", - SimpleType.INT, - SimpleType.INT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_bitShiftLeft_uint_int", - "Performs a bitwise-SHIFTLEFT operation over two uint values.", - SimpleType.UINT, - SimpleType.UINT, - SimpleType.INT)), - ImmutableSet.of( - CelFunctionBinding.from( - "math_bitShiftLeft_int_int", - Long.class, - Long.class, - CelMathExtensions::intBitShiftLeft), - CelFunctionBinding.from( - "math_bitShiftLeft_uint_int", - UnsignedLong.class, - Long.class, - CelMathExtensions::uintBitShiftLeft))), + CelMathCompilerLibrary.Function.BITSHIFTLEFT, CelMathRuntimeLibrary.Function.BITSHIFTLEFT), BITSHIFTRIGHT( - CelFunctionDecl.newFunctionDeclaration( - MATH_BIT_RIGHT_SHIFT_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_bitShiftRight_int_int", - "Performs a bitwise-SHIFTRIGHT operation over two int values.", - SimpleType.INT, - SimpleType.INT, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_bitShiftRight_uint_int", - "Performs a bitwise-SHIFTRIGHT operation over two uint values.", - SimpleType.UINT, - SimpleType.UINT, - SimpleType.INT)), - ImmutableSet.of( - CelFunctionBinding.from( - "math_bitShiftRight_int_int", - Long.class, - Long.class, - CelMathExtensions::intBitShiftRight), - CelFunctionBinding.from( - "math_bitShiftRight_uint_int", - UnsignedLong.class, - Long.class, - CelMathExtensions::uintBitShiftRight))), - SQRT( - CelFunctionDecl.newFunctionDeclaration( - MATH_SQRT_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "math_sqrt_double", - "Computes square root of the double value.", - SimpleType.DOUBLE, - SimpleType.DOUBLE), - CelOverloadDecl.newGlobalOverload( - "math_sqrt_int", - "Computes square root of the int value.", - SimpleType.DOUBLE, - SimpleType.INT), - CelOverloadDecl.newGlobalOverload( - "math_sqrt_uint", - "Computes square root of the unsigned value.", - SimpleType.DOUBLE, - SimpleType.UINT)), - ImmutableSet.of( - CelFunctionBinding.from( - "math_sqrt_double", Double.class, CelMathExtensions::sqrtDouble), - CelFunctionBinding.from("math_sqrt_int", Long.class, CelMathExtensions::sqrtInt), - CelFunctionBinding.from( - "math_sqrt_uint", UnsignedLong.class, CelMathExtensions::sqrtUint))); + CelMathCompilerLibrary.Function.BITSHIFTRIGHT, + CelMathRuntimeLibrary.Function.BITSHIFTRIGHT), + SQRT(CelMathCompilerLibrary.Function.SQRT, CelMathRuntimeLibrary.Function.SQRT); - private final CelFunctionDecl functionDecl; - private final ImmutableSet functionBindings; + private final CelMathCompilerLibrary.Function compilerFunction; + private final CelMathRuntimeLibrary.Function runtimeFunction; String getFunction() { - return functionDecl.name(); + return compilerFunction.getFunction(); } - Function(CelFunctionDecl functionDecl, ImmutableSet bindings) { - this.functionDecl = functionDecl; - this.functionBindings = bindings; + Function( + CelMathCompilerLibrary.Function compilerFunction, + CelMathRuntimeLibrary.Function runtimeFunction) { + this.compilerFunction = compilerFunction; + this.runtimeFunction = runtimeFunction; } } private static final class Library implements CelExtensionLibrary { - private final CelMathExtensions version0; - private final CelMathExtensions version1; - private final CelMathExtensions version2; + private final ImmutableSet versions; Library() { - version0 = new CelMathExtensions(0, ImmutableSet.of(Function.MIN, Function.MAX)); - - version1 = - new CelMathExtensions( - 1, - ImmutableSet.builder() - .addAll(version0.functions) - .add( - Function.CEIL, - Function.FLOOR, - Function.ROUND, - Function.TRUNC, - Function.ISINF, - Function.ISNAN, - Function.ISFINITE, - Function.ABS, - Function.SIGN, - Function.BITAND, - Function.BITOR, - Function.BITXOR, - Function.BITNOT, - Function.BITSHIFTLEFT, - Function.BITSHIFTRIGHT) - .build()); - - version2 = - new CelMathExtensions( - 2, - ImmutableSet.builder() - .addAll(version1.functions) - .add(Function.SQRT) - .build()); + versions = + CelMathCompilerLibrary.library().versions().stream() + .map(CelMathExtensions::new) + .collect(toImmutableSet()); } @Override public String name() { - return "math"; + return CelMathCompilerLibrary.library().name(); } @Override public ImmutableSet versions() { - return ImmutableSet.of(version0, version1, version2); + return versions; } } @@ -730,378 +103,50 @@ static CelExtensionLibrary library() { return LIBRARY; } - private final ImmutableSet functions; - private final int version; + private final CelMathCompilerLibrary compilerLibrary; + private final CelMathRuntimeLibrary mathRuntime; CelMathExtensions(Set functions) { - this(-1, functions); + this.compilerLibrary = + new CelMathCompilerLibrary( + functions.stream().map(f -> f.compilerFunction).collect(toImmutableSet())); + this.mathRuntime = + new CelMathRuntimeLibrary( + functions.stream().map(f -> f.runtimeFunction).collect(toImmutableSet())); } - private CelMathExtensions(int version, Set functions) { - this.version = version; - this.functions = ImmutableSet.copyOf(functions); + private CelMathExtensions(CelMathCompilerLibrary compilerLibrary) { + this.compilerLibrary = compilerLibrary; + this.mathRuntime = CelMathRuntimeLibrary.math(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() { - return ImmutableSet.of( - CelMacro.newReceiverVarArgMacro("greatest", CelMathExtensions::expandGreatestMacro), - CelMacro.newReceiverVarArgMacro("least", CelMathExtensions::expandLeastMacro)); + 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 public void setRuntimeOptions(CelRuntimeBuilder runtimeBuilder) { - functions.forEach( - function -> { - ImmutableSet combined = function.functionBindings; - if (!combined.isEmpty()) { - runtimeBuilder.addFunctionBindings( - CelFunctionBinding.fromOverloads(function.functionDecl.name(), combined)); - } - }); - } - - private static Optional expandGreatestMacro( - CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { - if (!isTargetInNamespace(target)) { - // Return empty to indicate that we're not interested in expanding this macro, and - // that the parser should default to a function call on the receiver. - return Optional.empty(); - } - - switch (arguments.size()) { - case 0: - return newError(exprFactory, "math.greatest() requires at least one argument", target); - case 1: - Optional invalidArg = - checkInvalidArgumentSingleArg(exprFactory, "math.greatest()", arguments.get(0)); - if (invalidArg.isPresent()) { - return invalidArg; - } - - return Optional.of(exprFactory.newGlobalCall(MATH_MAX_FUNCTION, arguments.get(0))); - case 2: - invalidArg = checkInvalidArgument(exprFactory, "math.greatest()", arguments); - if (invalidArg.isPresent()) { - return invalidArg; - } - - return Optional.of(exprFactory.newGlobalCall(MATH_MAX_FUNCTION, arguments)); - default: - invalidArg = checkInvalidArgument(exprFactory, "math.greatest()", arguments); - if (invalidArg.isPresent()) { - return invalidArg; - } - - return Optional.of( - exprFactory.newGlobalCall(MATH_MAX_FUNCTION, exprFactory.newList(arguments))); - } - } - - private static Comparable maxPair(Comparable x, Comparable y) { - if (x.getClass().equals(y.getClass())) { - return max(x, y); - } - - return CLASSES_TO_COMPARATORS.get(x.getClass(), y.getClass()).apply(x, y) >= 0 ? x : y; - } - - private static Comparable maxList(List list) { - if (list.isEmpty()) { - throw new IllegalStateException("math.@max(list) argument must not be empty"); - } - - Comparable max = list.get(0); - for (int i = 1; i < list.size(); i++) { - max = maxPair(max, list.get(i)); - } - - return max; - } - - private static Comparable minPair(Comparable x, Comparable y) { - if (x.getClass().equals(y.getClass())) { - return min(x, y); - } - - return CLASSES_TO_COMPARATORS.get(x.getClass(), y.getClass()).apply(x, y) <= 0 ? x : y; - } - - private static long absExact(long x) { - if (x == Long.MIN_VALUE) { - // The only case where standard Math.abs overflows silently - throw new CelNumericOverflowException("integer overflow"); - } - return Math.abs(x); - } - - private static boolean isNaN(double x) { - return Double.isNaN(x); - } - - private static Double trunc(Double x) { - if (isNaN(x) || isInfinite(x)) { - return x; - } - return (double) x.longValue(); - } - - private static boolean isInfinite(double x) { - return Double.isInfinite(x); - } - - private static double round(double x) { - if (isNaN(x) || isInfinite(x)) { - return x; - } - return DoubleMath.roundToLong(x, RoundingMode.HALF_UP); - } - - private static Number sign(Number x) { - if (x instanceof Double) { - double val = x.doubleValue(); - if (isNaN(val)) { - return val; - } - if (val == 0) { - return 0.0; - } - return val > 0 ? 1.0 : -1.0; - } - - if (x instanceof Long) { - long val = x.longValue(); - if (val == 0) { - return 0L; - } - return val > 0 ? 1L : -1L; - } - - if (x instanceof UnsignedLong) { - UnsignedLong val = (UnsignedLong) x; - if (val.equals(UnsignedLong.ZERO)) { - return val; - } - return UnsignedLong.ONE; - } - - throw new IllegalArgumentException("Unsupported type: " + x.getClass()); - } - - private static Long intBitAnd(long x, long y) { - return x & y; - } - - private static UnsignedLong uintBitAnd(UnsignedLong x, UnsignedLong y) { - return UnsignedLong.fromLongBits(x.longValue() & y.longValue()); - } - - private static Long intBitOr(long x, long y) { - return x | y; - } - - private static UnsignedLong uintBitOr(UnsignedLong x, UnsignedLong y) { - return UnsignedLong.fromLongBits(x.longValue() | y.longValue()); - } - - private static Long intBitXor(long x, long y) { - return x ^ y; - } - - private static UnsignedLong uintBitXor(UnsignedLong x, UnsignedLong y) { - return UnsignedLong.fromLongBits(x.longValue() ^ y.longValue()); - } - - private static Long intBitNot(long x) { - return ~x; - } - - private static UnsignedLong uintBitNot(UnsignedLong x) { - return UnsignedLong.fromLongBits(~x.longValue()); - } - - private static Long intBitShiftLeft(long value, long shiftAmount) { - if (shiftAmount < 0) { - throw new IllegalArgumentException("math.bitShiftLeft() negative offset:" + shiftAmount); - } - - if (shiftAmount > MAX_BIT_SHIFT) { - return 0L; - } - return value << shiftAmount; - } - - private static UnsignedLong uintBitShiftLeft(UnsignedLong value, long shiftAmount) { - if (shiftAmount < 0) { - throw new IllegalArgumentException("math.bitShiftLeft() negative offset:" + shiftAmount); - } - - if (shiftAmount > MAX_BIT_SHIFT) { - return UnsignedLong.ZERO; - } - return UnsignedLong.fromLongBits(value.longValue() << shiftAmount); - } - - private static Long intBitShiftRight(long value, long shiftAmount) { - if (shiftAmount < 0) { - throw new IllegalArgumentException("math.bitShiftRight() negative offset:" + shiftAmount); - } - - if (shiftAmount > MAX_BIT_SHIFT) { - return 0L; - } - return value >>> shiftAmount; - } - - private static UnsignedLong uintBitShiftRight(UnsignedLong value, long shiftAmount) { - if (shiftAmount < 0) { - throw new IllegalArgumentException("math.bitShiftRight() negative offset:" + shiftAmount); - } - - if (shiftAmount > MAX_BIT_SHIFT) { - return UnsignedLong.ZERO; - } - return UnsignedLong.fromLongBits(value.longValue() >>> shiftAmount); - } - - private static Double sqrtDouble(double x) { - return Math.sqrt(x); - } - - private static Double sqrtInt(Long x) { - return sqrtDouble(x.doubleValue()); - } - - private static Double sqrtUint(UnsignedLong x) { - return sqrtDouble(x.doubleValue()); - } - - private static Comparable minList(List list) { - if (list.isEmpty()) { - throw new IllegalStateException("math.@min(list) argument must not be empty"); - } - - Comparable min = list.get(0); - for (int i = 1; i < list.size(); i++) { - min = minPair(min, list.get(i)); - } - - return min; - } - - private static Optional expandLeastMacro( - CelMacroExprFactory exprFactory, CelExpr target, ImmutableList arguments) { - if (!isTargetInNamespace(target)) { - // Return empty to indicate that we're not interested in expanding this macro, and - // that the parser should default to a function call on the receiver. - return Optional.empty(); - } - - switch (arguments.size()) { - case 0: - return newError(exprFactory, "math.least() requires at least one argument", target); - case 1: - Optional invalidArg = - checkInvalidArgumentSingleArg(exprFactory, "math.least()", arguments.get(0)); - if (invalidArg.isPresent()) { - return invalidArg; - } - - return Optional.of(exprFactory.newGlobalCall(MATH_MIN_FUNCTION, arguments.get(0))); - case 2: - invalidArg = checkInvalidArgument(exprFactory, "math.least()", arguments); - if (invalidArg.isPresent()) { - return invalidArg; - } - - return Optional.of(exprFactory.newGlobalCall(MATH_MIN_FUNCTION, arguments)); - default: - invalidArg = checkInvalidArgument(exprFactory, "math.least()", arguments); - if (invalidArg.isPresent()) { - return invalidArg; - } - - return Optional.of( - exprFactory.newGlobalCall(MATH_MIN_FUNCTION, exprFactory.newList(arguments))); - } - } - - private static boolean isTargetInNamespace(CelExpr target) { - return target.exprKind().getKind().equals(Kind.IDENT) - && target.ident().name().equals(MATH_NAMESPACE); - } - - private static Optional checkInvalidArgument( - CelMacroExprFactory exprFactory, String functionName, List arguments) { - - for (CelExpr arg : arguments) { - if (!isArgumentValidType(arg)) { - return newError( - exprFactory, - String.format("%s simple literal arguments must be numeric", functionName), - arg); - } - } - return Optional.empty(); - } - - private static Optional checkInvalidArgumentSingleArg( - CelMacroExprFactory exprFactory, String functionName, CelExpr argument) { - if (argument.exprKind().getKind() == Kind.LIST) { - if (argument.list().elements().isEmpty()) { - return newError( - exprFactory, String.format("%s invalid single argument value", functionName), argument); - } - - return checkInvalidArgument(exprFactory, functionName, argument.list().elements()); - } - if (isArgumentValidType(argument)) { - return Optional.empty(); - } - - return newError( - exprFactory, String.format("%s invalid single argument value", functionName), argument); - } - - private static boolean isArgumentValidType(CelExpr argument) { - if (argument.exprKind().getKind() == Kind.CONSTANT) { - CelConstant constant = argument.constant(); - return constant.getKind() == CelConstant.Kind.INT64_VALUE - || constant.getKind() == CelConstant.Kind.UINT64_VALUE - || constant.getKind() == CelConstant.Kind.DOUBLE_VALUE; - } else if (argument.exprKind().getKind().equals(Kind.LIST) - || argument.exprKind().getKind().equals(Kind.STRUCT) - || argument.exprKind().getKind().equals(Kind.MAP)) { - return false; - } - - return true; - } - - private static Optional newError( - CelMacroExprFactory exprFactory, String errorMessage, CelExpr argument) { - return Optional.of( - exprFactory.reportError( - CelIssue.formatError(exprFactory.getSourceLocation(argument), errorMessage))); + runtimeBuilder.addFunctionBindings(mathRuntime.newFunctionBindings()); } } diff --git a/extensions/src/main/java/dev/cel/extensions/CelMathRuntimeLibrary.java b/extensions/src/main/java/dev/cel/extensions/CelMathRuntimeLibrary.java new file mode 100644 index 000000000..f9e32ced1 --- /dev/null +++ b/extensions/src/main/java/dev/cel/extensions/CelMathRuntimeLibrary.java @@ -0,0 +1,655 @@ +// Copyright 2023 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.collect.Comparators.max; +import static com.google.common.collect.Comparators.min; + +import com.google.common.collect.ImmutableSet; +import com.google.common.collect.ImmutableTable; +import com.google.common.math.DoubleMath; +import com.google.common.primitives.UnsignedLong; +import com.google.errorprone.annotations.Immutable; +import dev.cel.common.exceptions.CelNumericOverflowException; +import dev.cel.common.internal.ComparisonFunctions; +import dev.cel.runtime.CelFunctionBinding; +import dev.cel.runtime.CelLiteRuntimeBuilder; +import dev.cel.runtime.CelLiteRuntimeLibrary; +import java.math.RoundingMode; +import java.util.List; +import java.util.Set; +import java.util.function.BiFunction; + +/** + * Runtime implementation of CEL Math extension functions. + * + *

Note: For equal numbers with different types, the result is always the first argument e.g.: + * math.greatest(1u, 1.0) -> 1u + */ +@SuppressWarnings({"rawtypes", "unchecked"}) // Use of raw Comparables. +@Immutable +public final class CelMathRuntimeLibrary implements CelLiteRuntimeLibrary { + + private static final String MATH_MAX_FUNCTION = "math.@max"; + private static final String MATH_MIN_FUNCTION = "math.@min"; + + // Rounding Functions + private static final String MATH_CEIL_FUNCTION = "math.ceil"; + private static final String MATH_FLOOR_FUNCTION = "math.floor"; + private static final String MATH_ROUND_FUNCTION = "math.round"; + private static final String MATH_TRUNC_FUNCTION = "math.trunc"; + + // Floating Point Functions + private static final String MATH_ISFINITE_FUNCTION = "math.isFinite"; + private static final String MATH_ISNAN_FUNCTION = "math.isNaN"; + private static final String MATH_ISINF_FUNCTION = "math.isInf"; + + // Signedness Functions + private static final String MATH_ABS_FUNCTION = "math.abs"; + private static final String MATH_SIGN_FUNCTION = "math.sign"; + + // Bitwise Functions + private static final String MATH_BIT_AND_FUNCTION = "math.bitAnd"; + private static final String MATH_BIT_OR_FUNCTION = "math.bitOr"; + private static final String MATH_BIT_XOR_FUNCTION = "math.bitXor"; + private static final String MATH_BIT_NOT_FUNCTION = "math.bitNot"; + private static final String MATH_BIT_LEFT_SHIFT_FUNCTION = "math.bitShiftLeft"; + private static final String MATH_BIT_RIGHT_SHIFT_FUNCTION = "math.bitShiftRight"; + + private static final String MATH_SQRT_FUNCTION = "math.sqrt"; + + private static final int MAX_BIT_SHIFT = 63; + + /** + * Returns the proper comparison function to use for a math function call involving different + * argument types. + * + *

Example: (uint, int) -> {@link ComparisonFunctions#compareUintInt(UnsignedLong, long)} + */ + private static final ImmutableTable> + CLASSES_TO_COMPARATORS = newComparatorTable(); + + private static ImmutableTable> + newComparatorTable() { + ImmutableTable.Builder> builder = + new ImmutableTable.Builder<>(); + builder.put( + Long.class, + Double.class, + (x, y) -> ComparisonFunctions.compareIntDouble((Long) x, (Double) y)); + builder.put( + Double.class, + Long.class, + (x, y) -> ComparisonFunctions.compareDoubleInt((Double) x, (Long) y)); + builder.put( + Double.class, + UnsignedLong.class, + (x, y) -> ComparisonFunctions.compareDoubleUint((Double) x, (UnsignedLong) y)); + builder.put( + UnsignedLong.class, + Double.class, + (x, y) -> ComparisonFunctions.compareUintDouble((UnsignedLong) x, (Double) y)); + builder.put( + Long.class, + UnsignedLong.class, + (x, y) -> ComparisonFunctions.compareIntUint((Long) x, (UnsignedLong) y)); + builder.put( + UnsignedLong.class, + Long.class, + (x, y) -> ComparisonFunctions.compareUintInt((UnsignedLong) x, (Long) y)); + return builder.buildOrThrow(); + } + + /** Enumeration of runtime function bindings for the Math extension. */ + public enum Function { + MAX( + MATH_MAX_FUNCTION, + ImmutableSet.builder() + .add(CelFunctionBinding.from("math_@max_double", Double.class, x -> x)) + .add(CelFunctionBinding.from("math_@max_int", Long.class, x -> x)) + .add( + CelFunctionBinding.from( + "math_@max_double_double", + Double.class, + Double.class, + CelMathRuntimeLibrary::maxPair)) + .add( + CelFunctionBinding.from( + "math_@max_int_int", Long.class, Long.class, CelMathRuntimeLibrary::maxPair)) + .add( + CelFunctionBinding.from( + "math_@max_int_double", + Long.class, + Double.class, + CelMathRuntimeLibrary::maxPair)) + .add( + CelFunctionBinding.from( + "math_@max_double_int", + Double.class, + Long.class, + CelMathRuntimeLibrary::maxPair)) + .add( + CelFunctionBinding.from( + "math_@max_list_dyn", List.class, CelMathRuntimeLibrary::maxList)) + .add(CelFunctionBinding.from("math_@max_uint", UnsignedLong.class, x -> x)) + .add( + CelFunctionBinding.from( + "math_@max_uint_uint", + UnsignedLong.class, + UnsignedLong.class, + CelMathRuntimeLibrary::maxPair)) + .add( + CelFunctionBinding.from( + "math_@max_double_uint", + Double.class, + UnsignedLong.class, + CelMathRuntimeLibrary::maxPair)) + .add( + CelFunctionBinding.from( + "math_@max_uint_int", + UnsignedLong.class, + Long.class, + CelMathRuntimeLibrary::maxPair)) + .add( + CelFunctionBinding.from( + "math_@max_uint_double", + UnsignedLong.class, + Double.class, + CelMathRuntimeLibrary::maxPair)) + .add( + CelFunctionBinding.from( + "math_@max_int_uint", + Long.class, + UnsignedLong.class, + CelMathRuntimeLibrary::maxPair)) + .build()), + MIN( + MATH_MIN_FUNCTION, + ImmutableSet.builder() + .add(CelFunctionBinding.from("math_@min_double", Double.class, x -> x)) + .add(CelFunctionBinding.from("math_@min_int", Long.class, x -> x)) + .add( + CelFunctionBinding.from( + "math_@min_double_double", + Double.class, + Double.class, + CelMathRuntimeLibrary::minPair)) + .add( + CelFunctionBinding.from( + "math_@min_int_int", Long.class, Long.class, CelMathRuntimeLibrary::minPair)) + .add( + CelFunctionBinding.from( + "math_@min_int_double", + Long.class, + Double.class, + CelMathRuntimeLibrary::minPair)) + .add( + CelFunctionBinding.from( + "math_@min_double_int", + Double.class, + Long.class, + CelMathRuntimeLibrary::minPair)) + .add( + CelFunctionBinding.from( + "math_@min_list_dyn", List.class, CelMathRuntimeLibrary::minList)) + .add(CelFunctionBinding.from("math_@min_uint", UnsignedLong.class, x -> x)) + .add( + CelFunctionBinding.from( + "math_@min_uint_uint", + UnsignedLong.class, + UnsignedLong.class, + CelMathRuntimeLibrary::minPair)) + .add( + CelFunctionBinding.from( + "math_@min_double_uint", + Double.class, + UnsignedLong.class, + CelMathRuntimeLibrary::minPair)) + .add( + CelFunctionBinding.from( + "math_@min_uint_int", + UnsignedLong.class, + Long.class, + CelMathRuntimeLibrary::minPair)) + .add( + CelFunctionBinding.from( + "math_@min_uint_double", + UnsignedLong.class, + Double.class, + CelMathRuntimeLibrary::minPair)) + .add( + CelFunctionBinding.from( + "math_@min_int_uint", + Long.class, + UnsignedLong.class, + CelMathRuntimeLibrary::minPair)) + .build()), + CEIL( + MATH_CEIL_FUNCTION, + ImmutableSet.of(CelFunctionBinding.from("math_ceil_double", Double.class, Math::ceil))), + FLOOR( + MATH_FLOOR_FUNCTION, + ImmutableSet.of(CelFunctionBinding.from("math_floor_double", Double.class, Math::floor))), + ROUND( + MATH_ROUND_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_round_double", Double.class, CelMathRuntimeLibrary::round))), + TRUNC( + MATH_TRUNC_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_trunc_double", Double.class, CelMathRuntimeLibrary::trunc))), + ISFINITE( + MATH_ISFINITE_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from("math_isFinite_double", Double.class, Double::isFinite))), + ISNAN( + MATH_ISNAN_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_isNaN_double", Double.class, CelMathRuntimeLibrary::isNaN))), + ISINF( + MATH_ISINF_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_isInf_double", Double.class, CelMathRuntimeLibrary::isInfinite))), + ABS( + MATH_ABS_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from("math_abs_double", Double.class, Math::abs), + CelFunctionBinding.from("math_abs_int", Long.class, CelMathRuntimeLibrary::absExact), + CelFunctionBinding.from("math_abs_uint", UnsignedLong.class, x -> x))), + SIGN( + MATH_SIGN_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from("math_sign_double", Double.class, CelMathRuntimeLibrary::sign), + CelFunctionBinding.from("math_sign_int", Long.class, CelMathRuntimeLibrary::sign), + CelFunctionBinding.from( + "math_sign_uint", UnsignedLong.class, CelMathRuntimeLibrary::sign))), + BITAND( + MATH_BIT_AND_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_bitAnd_int_int", Long.class, Long.class, CelMathRuntimeLibrary::intBitAnd), + CelFunctionBinding.from( + "math_bitAnd_uint_uint", + UnsignedLong.class, + UnsignedLong.class, + CelMathRuntimeLibrary::uintBitAnd))), + BITOR( + MATH_BIT_OR_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_bitOr_int_int", Long.class, Long.class, CelMathRuntimeLibrary::intBitOr), + CelFunctionBinding.from( + "math_bitOr_uint_uint", + UnsignedLong.class, + UnsignedLong.class, + CelMathRuntimeLibrary::uintBitOr))), + BITXOR( + MATH_BIT_XOR_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_bitXor_int_int", Long.class, Long.class, CelMathRuntimeLibrary::intBitXor), + CelFunctionBinding.from( + "math_bitXor_uint_uint", + UnsignedLong.class, + UnsignedLong.class, + CelMathRuntimeLibrary::uintBitXor))), + BITNOT( + MATH_BIT_NOT_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_bitNot_int_int", Long.class, CelMathRuntimeLibrary::intBitNot), + CelFunctionBinding.from( + "math_bitNot_uint_uint", UnsignedLong.class, CelMathRuntimeLibrary::uintBitNot))), + BITSHIFTLEFT( + MATH_BIT_LEFT_SHIFT_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_bitShiftLeft_int_int", + Long.class, + Long.class, + CelMathRuntimeLibrary::intBitShiftLeft), + CelFunctionBinding.from( + "math_bitShiftLeft_uint_int", + UnsignedLong.class, + Long.class, + CelMathRuntimeLibrary::uintBitShiftLeft))), + BITSHIFTRIGHT( + MATH_BIT_RIGHT_SHIFT_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_bitShiftRight_int_int", + Long.class, + Long.class, + CelMathRuntimeLibrary::intBitShiftRight), + CelFunctionBinding.from( + "math_bitShiftRight_uint_int", + UnsignedLong.class, + Long.class, + CelMathRuntimeLibrary::uintBitShiftRight))), + SQRT( + MATH_SQRT_FUNCTION, + ImmutableSet.of( + CelFunctionBinding.from( + "math_sqrt_double", Double.class, CelMathRuntimeLibrary::sqrtDouble), + CelFunctionBinding.from("math_sqrt_int", Long.class, CelMathRuntimeLibrary::sqrtInt), + CelFunctionBinding.from( + "math_sqrt_uint", UnsignedLong.class, CelMathRuntimeLibrary::sqrtUint))); + + private final String functionName; + private final ImmutableSet functionBindings; + + public String getFunction() { + return functionName; + } + + public ImmutableSet getFunctionBindings() { + return functionBindings; + } + + Function(String functionName, ImmutableSet bindings) { + this.functionName = functionName; + this.functionBindings = bindings; + } + } + + private static final CelMathRuntimeLibrary VERSION_0 = + new CelMathRuntimeLibrary(0, ImmutableSet.of(Function.MIN, Function.MAX)); + + private static final CelMathRuntimeLibrary VERSION_1 = + new CelMathRuntimeLibrary( + 1, + ImmutableSet.builder() + .addAll(VERSION_0.functions) + .add( + Function.CEIL, + Function.FLOOR, + Function.ROUND, + Function.TRUNC, + Function.ISINF, + Function.ISNAN, + Function.ISFINITE, + Function.ABS, + Function.SIGN, + Function.BITAND, + Function.BITOR, + Function.BITXOR, + Function.BITNOT, + Function.BITSHIFTLEFT, + Function.BITSHIFTRIGHT) + .build()); + + private static final CelMathRuntimeLibrary VERSION_2 = + new CelMathRuntimeLibrary( + 2, + ImmutableSet.builder().addAll(VERSION_1.functions).add(Function.SQRT).build()); + + /** Returns the latest version of the 'math' runtime functions. */ + public static CelMathRuntimeLibrary math() { + return VERSION_2; + } + + /** Returns the specified version of the 'math' runtime functions. */ + public static CelMathRuntimeLibrary math(int version) { + switch (version) { + case 0: + return VERSION_0; + case 1: + return VERSION_1; + case 2: + case Integer.MAX_VALUE: + return VERSION_2; + default: + throw new IllegalArgumentException("Unsupported 'math' extension version " + version); + } + } + + /** Returns the 'math' runtime functions with only the specified functions. */ + public static CelMathRuntimeLibrary math(Function... functions) { + return math(ImmutableSet.copyOf(functions)); + } + + /** Returns the 'math' runtime functions with only the specified functions. */ + public static CelMathRuntimeLibrary math(Set functions) { + return new CelMathRuntimeLibrary(functions); + } + + private final int version; + private final ImmutableSet functions; + + CelMathRuntimeLibrary(Set functions) { + this(-1, functions); + } + + private CelMathRuntimeLibrary(int version, Set functions) { + this.version = version; + this.functions = ImmutableSet.copyOf(functions); + } + + public int version() { + return version; + } + + @Override + public void setRuntimeOptions(CelLiteRuntimeBuilder runtimeBuilder) { + runtimeBuilder.addFunctionBindings(newFunctionBindings()); + } + + /** Creates the {@link CelFunctionBinding}s for the configured math functions. */ + public ImmutableSet newFunctionBindings() { + ImmutableSet.Builder builder = ImmutableSet.builder(); + for (Function function : functions) { + if (!function.functionBindings.isEmpty()) { + builder.addAll( + CelFunctionBinding.fromOverloads(function.functionName, function.functionBindings)); + } + } + return builder.build(); + } + + private static Comparable maxPair(Comparable x, Comparable y) { + if (x.getClass().equals(y.getClass())) { + return max(x, y); + } + + return CLASSES_TO_COMPARATORS.get(x.getClass(), y.getClass()).apply(x, y) >= 0 ? x : y; + } + + private static Comparable maxList(List list) { + if (list.isEmpty()) { + throw new IllegalStateException("math.@max(list) argument must not be empty"); + } + + Comparable max = list.get(0); + for (int i = 1; i < list.size(); i++) { + max = maxPair(max, list.get(i)); + } + + return max; + } + + private static Comparable minPair(Comparable x, Comparable y) { + if (x.getClass().equals(y.getClass())) { + return min(x, y); + } + + return CLASSES_TO_COMPARATORS.get(x.getClass(), y.getClass()).apply(x, y) <= 0 ? x : y; + } + + private static long absExact(long x) { + if (x == Long.MIN_VALUE) { + // The only case where standard Math.abs overflows silently + throw new CelNumericOverflowException("integer overflow"); + } + return Math.abs(x); + } + + private static boolean isNaN(double x) { + return Double.isNaN(x); + } + + private static Double trunc(Double x) { + if (isNaN(x) || isInfinite(x)) { + return x; + } + return (double) x.longValue(); + } + + private static boolean isInfinite(double x) { + return Double.isInfinite(x); + } + + private static double round(double x) { + if (isNaN(x) || isInfinite(x)) { + return x; + } + return DoubleMath.roundToLong(x, RoundingMode.HALF_UP); + } + + private static Number sign(Number x) { + if (x instanceof Double) { + double val = x.doubleValue(); + if (isNaN(val)) { + return val; + } + if (val == 0) { + return 0.0; + } + return val > 0 ? 1.0 : -1.0; + } + + if (x instanceof Long) { + long val = x.longValue(); + if (val == 0) { + return 0L; + } + return val > 0 ? 1L : -1L; + } + + if (x instanceof UnsignedLong) { + UnsignedLong val = (UnsignedLong) x; + if (val.equals(UnsignedLong.ZERO)) { + return val; + } + return UnsignedLong.ONE; + } + + throw new IllegalArgumentException("Unsupported type: " + x.getClass()); + } + + private static Long intBitAnd(long x, long y) { + return x & y; + } + + private static UnsignedLong uintBitAnd(UnsignedLong x, UnsignedLong y) { + return UnsignedLong.fromLongBits(x.longValue() & y.longValue()); + } + + private static Long intBitOr(long x, long y) { + return x | y; + } + + private static UnsignedLong uintBitOr(UnsignedLong x, UnsignedLong y) { + return UnsignedLong.fromLongBits(x.longValue() | y.longValue()); + } + + private static Long intBitXor(long x, long y) { + return x ^ y; + } + + private static UnsignedLong uintBitXor(UnsignedLong x, UnsignedLong y) { + return UnsignedLong.fromLongBits(x.longValue() ^ y.longValue()); + } + + private static Long intBitNot(long x) { + return ~x; + } + + private static UnsignedLong uintBitNot(UnsignedLong x) { + return UnsignedLong.fromLongBits(~x.longValue()); + } + + private static Long intBitShiftLeft(long value, long shiftAmount) { + if (shiftAmount < 0) { + throw new IllegalArgumentException("math.bitShiftLeft() negative offset:" + shiftAmount); + } + + if (shiftAmount > MAX_BIT_SHIFT) { + return 0L; + } + return value << shiftAmount; + } + + private static UnsignedLong uintBitShiftLeft(UnsignedLong value, long shiftAmount) { + if (shiftAmount < 0) { + throw new IllegalArgumentException("math.bitShiftLeft() negative offset:" + shiftAmount); + } + + if (shiftAmount > MAX_BIT_SHIFT) { + return UnsignedLong.ZERO; + } + return UnsignedLong.fromLongBits(value.longValue() << shiftAmount); + } + + private static Long intBitShiftRight(long value, long shiftAmount) { + if (shiftAmount < 0) { + throw new IllegalArgumentException("math.bitShiftRight() negative offset:" + shiftAmount); + } + + if (shiftAmount > MAX_BIT_SHIFT) { + return 0L; + } + return value >>> shiftAmount; + } + + private static UnsignedLong uintBitShiftRight(UnsignedLong value, long shiftAmount) { + if (shiftAmount < 0) { + throw new IllegalArgumentException("math.bitShiftRight() negative offset:" + shiftAmount); + } + + if (shiftAmount > MAX_BIT_SHIFT) { + return UnsignedLong.ZERO; + } + return UnsignedLong.fromLongBits(value.longValue() >>> shiftAmount); + } + + private static Double sqrtDouble(double x) { + return Math.sqrt(x); + } + + private static Double sqrtInt(Long x) { + return sqrtDouble(x.doubleValue()); + } + + private static Double sqrtUint(UnsignedLong x) { + return sqrtDouble(x.doubleValue()); + } + + private static Comparable minList(List list) { + if (list.isEmpty()) { + throw new IllegalStateException("math.@min(list) argument must not be empty"); + } + + Comparable min = list.get(0); + for (int i = 1; i < list.size(); i++) { + min = minPair(min, list.get(i)); + } + + return min; + } +} diff --git a/extensions/src/test/java/dev/cel/extensions/BUILD.bazel b/extensions/src/test/java/dev/cel/extensions/BUILD.bazel index f7b996610..b9557e540 100644 --- a/extensions/src/test/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/test/java/dev/cel/extensions/BUILD.bazel @@ -14,7 +14,11 @@ java_library( "//bundle:cel", "//common:cel_ast", "//common:cel_exception", - "//common:compiler_common", + "//common:cel_function_decl", + "//common:cel_overload_decl", + "//common:cel_validation_exception", + "//common:cel_validation_result", + "//common:cel_var_decl", "//common:container", "//common:options", "//common/ast", @@ -33,6 +37,8 @@ java_library( "//extensions:extension_library", "//extensions:lite_extensions", "//extensions:math", + "//extensions:math_compiler_library", + "//extensions:math_runtime_library", "//extensions:native", "//extensions:optional_library", "//extensions:sets", diff --git a/extensions/src/test/java/dev/cel/extensions/CelMathExtensionsTest.java b/extensions/src/test/java/dev/cel/extensions/CelMathExtensionsTest.java index 5b57f1fb2..df8c6f0ab 100644 --- a/extensions/src/test/java/dev/cel/extensions/CelMathExtensionsTest.java +++ b/extensions/src/test/java/dev/cel/extensions/CelMathExtensionsTest.java @@ -35,6 +35,8 @@ import dev.cel.compiler.CelCompilerFactory; import dev.cel.runtime.CelEvaluationException; import dev.cel.runtime.CelFunctionBinding; +import dev.cel.runtime.CelLiteRuntime; +import dev.cel.runtime.CelLiteRuntimeFactory; import dev.cel.runtime.CelRuntime; import dev.cel.runtime.CelRuntimeFactory; import dev.cel.testing.CelRuntimeFlavor; @@ -1109,6 +1111,63 @@ public void sqrt_success(String expr, double expectedResult) throws Exception { assertThat(result).isEqualTo(expectedResult); } + @Test + public void separateLibraryAndRuntime_allFunctions_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelMathCompilerLibrary.math()) + .build(); + CelLiteRuntime celLiteRuntime = + CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .addLibraries(CelMathRuntimeLibrary.math()) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("math.greatest(1, 2.0)").getAst(); + Object result = celLiteRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo(2.0); + } + + @Test + public void separateLibraryAndRuntime_versioned_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelMathCompilerLibrary.math(0)) + .build(); + CelRuntime celRuntime = + CelRuntimeFactory.standardCelRuntimeBuilder() + .addFunctionBindings(CelMathRuntimeLibrary.math(0).newFunctionBindings()) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("math.greatest(1, 2) == 2").getAst(); + boolean result = (boolean) celRuntime.createProgram(ast).eval(); + + assertThat(result).isTrue(); + assertThrows( + CelValidationException.class, () -> celCompiler.compile("math.ceil(1.5)").getAst()); + } + + @Test + public void separateLibraryAndRuntime_subsetOfFunctions_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelMathCompilerLibrary.math(CelMathCompilerLibrary.Function.MAX)) + .build(); + CelRuntime celRuntime = + CelRuntimeFactory.standardCelRuntimeBuilder() + .addFunctionBindings( + CelMathRuntimeLibrary.math(CelMathRuntimeLibrary.Function.MAX) + .newFunctionBindings()) + .build(); + + CelAbstractSyntaxTree ast = celCompiler.compile("math.greatest(1, 2) == 2").getAst(); + boolean result = (boolean) celRuntime.createProgram(ast).eval(); + + assertThat(result).isTrue(); + assertThrows( + CelValidationException.class, () -> celCompiler.compile("math.least(1, 2)").getAst()); + } + private Object eval(Cel cel, String expression, Map variables) throws Exception { CelAbstractSyntaxTree ast; if (isParseOnly) {