diff --git a/extensions/BUILD.bazel b/extensions/BUILD.bazel index 8ef633901..cfe6e8fca 100644 --- a/extensions/BUILD.bazel +++ b/extensions/BUILD.bazel @@ -141,3 +141,25 @@ cel_android_library( name = "lists_runtime_library_android", exports = ["//extensions/src/main/java/dev/cel/extensions:lists_runtime_library_android"], ) + +java_library( + name = "regex", + exports = ["//extensions/src/main/java/dev/cel/extensions:regex"], +) + +java_library( + name = "regex_compiler_library", + visibility = ["//:internal"], + exports = ["//extensions/src/main/java/dev/cel/extensions:regex_compiler_library"], +) + +java_library( + name = "regex_runtime_library", + visibility = ["//:internal"], + exports = ["//extensions/src/main/java/dev/cel/extensions:regex_runtime_library"], +) + +cel_android_library( + name = "regex_runtime_library_android", + exports = ["//extensions/src/main/java/dev/cel/extensions:regex_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 1f48552a4..d664397c0 100644 --- a/extensions/src/main/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/main/java/dev/cel/extensions/BUILD.bazel @@ -483,20 +483,66 @@ cel_android_library( java_library( name = "regex", srcs = ["CelRegexExtensions.java"], + tags = [ + ], deps = [ + ":extension_library", + ":regex_compiler_library", + ":regex_runtime_library", "//checker:checker_builder", - "//common:compiler_common", - "//common/types", + "//common:cel_function_decl", "//compiler:compiler_builder", - "//extensions:extension_library", "//runtime", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven//:com_google_guava_guava", + ], +) + +java_library( + name = "regex_compiler_library", + srcs = ["CelRegexCompilerLibrary.java"], + tags = [ + ], + deps = [ + ":extension_library", + "//checker:checker_builder", + "//common:cel_function_decl", + "//common:cel_overload_decl", + "//common/types", + "//compiler:compiler_builder", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven//:com_google_guava_guava", + ], +) + +java_library( + name = "regex_runtime_library", + srcs = ["CelRegexRuntimeLibrary.java"], + tags = [ + ], + deps = [ "//runtime:function_binding", + "//runtime:lite_runtime", "@maven//:com_google_errorprone_error_prone_annotations", "@maven//:com_google_guava_guava", "@maven//:com_google_re2j_re2j", ], ) +cel_android_library( + name = "regex_runtime_library_android", + srcs = ["CelRegexRuntimeLibrary.java"], + tags = [ + ], + deps = [ + "//runtime:function_binding_android", + "//runtime:lite_runtime_android", + "@maven//:com_google_errorprone_error_prone_annotations", + "@maven//:com_google_re2j_re2j", + "@maven_android//:com_google_guava_guava", + ], +) + java_library( name = "comprehensions", srcs = ["CelComprehensionsExtensions.java"], diff --git a/extensions/src/main/java/dev/cel/extensions/CelExtensions.java b/extensions/src/main/java/dev/cel/extensions/CelExtensions.java index 8ca84ffa9..fc9a5874c 100644 --- a/extensions/src/main/java/dev/cel/extensions/CelExtensions.java +++ b/extensions/src/main/java/dev/cel/extensions/CelExtensions.java @@ -341,6 +341,39 @@ public static CelRegexExtensions regex() { return REGEX_EXTENSIONS; } + /** + * Extended functions for Regular Expressions. + * + *

Refer to README.md for available functions. + */ + public static CelRegexExtensions regex(int version) { + return CelRegexExtensions.library().version(version); + } + + /** + * Extended functions for Regular Expressions. + * + *

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

This will include only the specific functions denoted by {@link + * CelRegexExtensions.Function}. + */ + public static CelRegexExtensions regex(CelRegexExtensions.Function... functions) { + return regex(ImmutableSet.copyOf(functions)); + } + + /** + * Extended functions for Regular Expressions. + * + *

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

This will include only the specific functions denoted by {@link + * CelRegexExtensions.Function}. + */ + public static CelRegexExtensions regex(Set functions) { + return new CelRegexExtensions(functions); + } + /** * Extended functions for Two Variable Comprehensions Expressions. * diff --git a/extensions/src/main/java/dev/cel/extensions/CelRegexCompilerLibrary.java b/extensions/src/main/java/dev/cel/extensions/CelRegexCompilerLibrary.java new file mode 100644 index 000000000..a5137200e --- /dev/null +++ b/extensions/src/main/java/dev/cel/extensions/CelRegexCompilerLibrary.java @@ -0,0 +1,160 @@ +// Copyright 2025 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package dev.cel.extensions; + +import static com.google.common.collect.ImmutableSet.toImmutableSet; + +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.CelOverloadDecl; +import dev.cel.common.types.ListType; +import dev.cel.common.types.OptionalType; +import dev.cel.common.types.SimpleType; +import dev.cel.compiler.CelCompilerLibrary; +import java.util.Set; + +/** Internal implementation of CEL regex compile-time extensions. */ +@Immutable +public final class CelRegexCompilerLibrary + implements CelCompilerLibrary, CelExtensionLibrary.FeatureSet { + + private static final String REGEX_REPLACE_FUNCTION = "regex.replace"; + private static final String REGEX_EXTRACT_FUNCTION = "regex.extract"; + private static final String REGEX_EXTRACT_ALL_FUNCTION = "regex.extractAll"; + + /** Supported functions for the Regex compiler extension. */ + public enum Function { + REPLACE( + CelFunctionDecl.newFunctionDeclaration( + REGEX_REPLACE_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "regex_replaceAll_string_string_string", + "Replaces all the matched values using the given replace string.", + SimpleType.STRING, + SimpleType.STRING, + SimpleType.STRING, + SimpleType.STRING), + CelOverloadDecl.newGlobalOverload( + "regex_replaceCount_string_string_string_int", + "Replaces the given number of matched values using the given replace string.", + SimpleType.STRING, + SimpleType.STRING, + SimpleType.STRING, + SimpleType.STRING, + SimpleType.INT))), + EXTRACT( + CelFunctionDecl.newFunctionDeclaration( + REGEX_EXTRACT_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "regex_extract_string_string", + "Returns the first substring that matches the regex.", + OptionalType.create(SimpleType.STRING), + SimpleType.STRING, + SimpleType.STRING))), + EXTRACTALL( + CelFunctionDecl.newFunctionDeclaration( + REGEX_EXTRACT_ALL_FUNCTION, + CelOverloadDecl.newGlobalOverload( + "regex_extractAll_string_string", + "Returns an array of all substrings that match the regex.", + ListType.create(SimpleType.STRING), + SimpleType.STRING, + SimpleType.STRING))); + + private final CelFunctionDecl functionDecl; + + 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 CelRegexCompilerLibrary version0 = + new CelRegexCompilerLibrary(0, ImmutableSet.copyOf(Function.values())); + + @Override + public String name() { + return "regex"; + } + + @Override + public ImmutableSet versions() { + return ImmutableSet.of(version0); + } + } + + private static final Library LIBRARY = new Library(); + + public static CelExtensionLibrary library() { + return LIBRARY; + } + + /** Returns the latest version of the 'regex' compiler extension. */ + public static CelRegexCompilerLibrary regex() { + return library().latest(); + } + + /** Returns the specified version of the 'regex' compiler extension. */ + public static CelRegexCompilerLibrary regex(int version) { + return library().version(version); + } + + /** Returns the 'regex' compiler extension with only the specified functions. */ + public static CelRegexCompilerLibrary regex(Function... functions) { + return regex(ImmutableSet.copyOf(functions)); + } + + /** Returns the 'regex' compiler extension with only the specified functions. */ + public static CelRegexCompilerLibrary regex(Set functions) { + return new CelRegexCompilerLibrary(functions); + } + + private final ImmutableSet functions; + private final int version; + + CelRegexCompilerLibrary(Set functions) { + this(-1, functions); + } + + private CelRegexCompilerLibrary(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 void setCheckerOptions(CelCheckerBuilder checkerBuilder) { + functions.forEach(function -> checkerBuilder.addFunctionDeclarations(function.functionDecl)); + } +} diff --git a/extensions/src/main/java/dev/cel/extensions/CelRegexExtensions.java b/extensions/src/main/java/dev/cel/extensions/CelRegexExtensions.java index 564422cd4..06c230151 100644 --- a/extensions/src/main/java/dev/cel/extensions/CelRegexExtensions.java +++ b/extensions/src/main/java/dev/cel/extensions/CelRegexExtensions.java @@ -16,23 +16,13 @@ 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 com.google.re2j.Matcher; -import com.google.re2j.Pattern; -import com.google.re2j.PatternSyntaxException; import dev.cel.checker.CelCheckerBuilder; import dev.cel.common.CelFunctionDecl; -import dev.cel.common.CelOverloadDecl; -import dev.cel.common.types.ListType; -import dev.cel.common.types.OptionalType; -import dev.cel.common.types.SimpleType; import dev.cel.compiler.CelCompilerLibrary; -import dev.cel.runtime.CelFunctionBinding; import dev.cel.runtime.CelRuntimeBuilder; import dev.cel.runtime.CelRuntimeLibrary; -import java.util.Optional; import java.util.Set; /** Internal implementation of CEL regex extensions. */ @@ -40,271 +30,93 @@ public final class CelRegexExtensions implements CelCompilerLibrary, CelRuntimeLibrary, CelExtensionLibrary.FeatureSet { - private static final String REGEX_REPLACE_FUNCTION = "regex.replace"; - private static final String REGEX_EXTRACT_FUNCTION = "regex.extract"; - private static final String REGEX_EXTRACT_ALL_FUNCTION = "regex.extractAll"; - - enum Function { - REPLACE( - CelFunctionDecl.newFunctionDeclaration( - REGEX_REPLACE_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "regex_replaceAll_string_string_string", - "Replaces all the matched values using the given replace string.", - SimpleType.STRING, - SimpleType.STRING, - SimpleType.STRING, - SimpleType.STRING), - CelOverloadDecl.newGlobalOverload( - "regex_replaceCount_string_string_string_int", - "Replaces the given number of matched values using the given replace string.", - SimpleType.STRING, - SimpleType.STRING, - SimpleType.STRING, - SimpleType.STRING, - SimpleType.INT)), - ImmutableSet.of( - CelFunctionBinding.from( - "regex_replaceAll_string_string_string", - ImmutableList.of(String.class, String.class, String.class), - (args) -> { - String target = (String) args[0]; - String pattern = (String) args[1]; - String replaceStr = (String) args[2]; - return CelRegexExtensions.replace(target, pattern, replaceStr); - }), - CelFunctionBinding.from( - "regex_replaceCount_string_string_string_int", - ImmutableList.of(String.class, String.class, String.class, Long.class), - (args) -> { - String target = (String) args[0]; - String pattern = (String) args[1]; - String replaceStr = (String) args[2]; - long count = (long) args[3]; - return CelRegexExtensions.replaceN(target, pattern, replaceStr, count); - }))), - EXTRACT( - CelFunctionDecl.newFunctionDeclaration( - REGEX_EXTRACT_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "regex_extract_string_string", - "Returns the first substring that matches the regex.", - OptionalType.create(SimpleType.STRING), - SimpleType.STRING, - SimpleType.STRING)), - ImmutableSet.of( - CelFunctionBinding.from( - "regex_extract_string_string", - String.class, - String.class, - CelRegexExtensions::extract))), + /** Denotes the regex extension function. */ + public enum Function { + REPLACE(CelRegexCompilerLibrary.Function.REPLACE, CelRegexRuntimeLibrary.Function.REPLACE), + EXTRACT(CelRegexCompilerLibrary.Function.EXTRACT, CelRegexRuntimeLibrary.Function.EXTRACT), EXTRACTALL( - CelFunctionDecl.newFunctionDeclaration( - REGEX_EXTRACT_ALL_FUNCTION, - CelOverloadDecl.newGlobalOverload( - "regex_extractAll_string_string", - "Returns an array of all substrings that match the regex.", - ListType.create(SimpleType.STRING), - SimpleType.STRING, - SimpleType.STRING)), - ImmutableSet.of( - CelFunctionBinding.from( - "regex_extractAll_string_string", - String.class, - String.class, - CelRegexExtensions::extractAll))); + CelRegexCompilerLibrary.Function.EXTRACTALL, CelRegexRuntimeLibrary.Function.EXTRACTALL); - private final CelFunctionDecl functionDecl; - private final ImmutableSet functionBindings; + private final CelRegexCompilerLibrary.Function compilerFunction; + private final CelRegexRuntimeLibrary.Function runtimeFunction; String getFunction() { - return functionDecl.name(); + return compilerFunction.getFunction(); } - Function(CelFunctionDecl functionDecl, ImmutableSet functionBindings) { - this.functionDecl = functionDecl; - this.functionBindings = - CelFunctionBinding.fromOverloads(functionDecl.name(), functionBindings); + Function( + CelRegexCompilerLibrary.Function compilerFunction, + CelRegexRuntimeLibrary.Function runtimeFunction) { + this.compilerFunction = compilerFunction; + this.runtimeFunction = runtimeFunction; } } - private static final CelExtensionLibrary LIBRARY = - new CelExtensionLibrary() { - private final CelRegexExtensions version0 = new CelRegexExtensions(); + private static final class Library implements CelExtensionLibrary { + private final ImmutableSet versions; + + Library() { + versions = + CelRegexCompilerLibrary.library().versions().stream() + .map(CelRegexExtensions::new) + .collect(toImmutableSet()); + } + + @Override + public String name() { + return CelRegexCompilerLibrary.library().name(); + } - @Override - public String name() { - return "regex"; - } + @Override + public ImmutableSet versions() { + return versions; + } + } - @Override - public ImmutableSet versions() { - return ImmutableSet.of(version0); - } - }; + private static final Library LIBRARY = new Library(); static CelExtensionLibrary library() { return LIBRARY; } - private final ImmutableSet functions; + private final CelRegexCompilerLibrary compilerLibrary; + private final CelRegexRuntimeLibrary regexRuntime; CelRegexExtensions() { - this.functions = ImmutableSet.copyOf(Function.values()); + this(CelRegexCompilerLibrary.regex()); } CelRegexExtensions(Set functions) { - this.functions = ImmutableSet.copyOf(functions); + this.compilerLibrary = + new CelRegexCompilerLibrary( + functions.stream().map(f -> f.compilerFunction).collect(toImmutableSet())); + this.regexRuntime = + new CelRegexRuntimeLibrary( + functions.stream().map(f -> f.runtimeFunction).collect(toImmutableSet())); + } + + private CelRegexExtensions(CelRegexCompilerLibrary compilerLibrary) { + this.compilerLibrary = compilerLibrary; + this.regexRuntime = CelRegexRuntimeLibrary.regex(compilerLibrary.version()); } @Override public int version() { - return 0; + return compilerLibrary.version(); } @Override public ImmutableSet functions() { - return functions.stream().map(f -> f.functionDecl).collect(toImmutableSet()); + return compilerLibrary.functions(); } @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 -> runtimeBuilder.addFunctionBindings(function.functionBindings)); - } - - private static Pattern compileRegexPattern(String regex) { - try { - return Pattern.compile(regex); - } catch (PatternSyntaxException e) { - throw new IllegalArgumentException("Failed to compile regex: " + regex, e); - } - } - - private static String replace(String target, String regex, String replaceStr) { - return replaceN(target, regex, replaceStr, -1); - } - - private static String replaceN( - String target, String regex, String replaceStr, long replaceCount) { - if (replaceCount == 0) { - return target; - } - // For all negative replaceCount, do a replaceAll - if (replaceCount < 0) { - replaceCount = -1; - } - - Pattern pattern = compileRegexPattern(regex); - Matcher matcher = pattern.matcher(target); - StringBuffer sb = new StringBuffer(); - int counter = 0; - - while (matcher.find()) { - if (replaceCount != -1 && counter >= replaceCount) { - break; - } - - String processedReplacement = replaceStrValidator(matcher, replaceStr); - matcher.appendReplacement(sb, Matcher.quoteReplacement(processedReplacement)); - counter++; - } - matcher.appendTail(sb); - - return sb.toString(); - } - - private static String replaceStrValidator(Matcher matcher, String replacement) { - StringBuilder sb = new StringBuilder(); - for (int i = 0; i < replacement.length(); i++) { - char c = replacement.charAt(i); - - if (c != '\\') { - sb.append(c); - continue; - } - - if (i + 1 >= replacement.length()) { - throw new IllegalArgumentException("Invalid replacement string: \\ not allowed at end"); - } - - char nextChar = replacement.charAt(++i); - - if (Character.isDigit(nextChar)) { - int groupNum = Character.digit(nextChar, 10); - int groupCount = matcher.groupCount(); - - if (groupNum > groupCount) { - throw new IllegalArgumentException( - "Replacement string references group " - + groupNum - + " but regex has only " - + groupCount - + " group(s)"); - } - - String groupValue = matcher.group(groupNum); - if (groupValue != null) { - sb.append(groupValue); - } - } else if (nextChar == '\\') { - sb.append('\\'); - } else { - throw new IllegalArgumentException( - "Invalid replacement string: \\ must be followed by a digit"); - } - } - return sb.toString(); - } - - private static Optional extract(String target, String regex) { - Pattern pattern = compileRegexPattern(regex); - Matcher matcher = pattern.matcher(target); - - if (!matcher.find()) { - return Optional.empty(); - } - - int groupCount = matcher.groupCount(); - if (groupCount > 1) { - throw new IllegalArgumentException( - "Regular expression has more than one capturing group: " + regex); - } - - String result = (groupCount == 1) ? matcher.group(1) : matcher.group(0); - - return Optional.ofNullable(result); - } - - private static ImmutableList extractAll(String target, String regex) { - Pattern pattern = compileRegexPattern(regex); - Matcher matcher = pattern.matcher(target); - - if (matcher.groupCount() > 1) { - throw new IllegalArgumentException( - "Regular expression has more than one capturing group: " + regex); - } - - ImmutableList.Builder builder = ImmutableList.builder(); - boolean hasOneGroup = matcher.groupCount() == 1; - - while (matcher.find()) { - if (hasOneGroup) { - String group = matcher.group(1); - // Add the captured group's content only if it's not null - if (group != null) { - builder.add(group); - } - } else { - // No capturing groups - builder.add(matcher.group(0)); - } - } - - return builder.build(); + runtimeBuilder.addFunctionBindings(regexRuntime.newFunctionBindings()); } } diff --git a/extensions/src/main/java/dev/cel/extensions/CelRegexRuntimeLibrary.java b/extensions/src/main/java/dev/cel/extensions/CelRegexRuntimeLibrary.java new file mode 100644 index 000000000..276295096 --- /dev/null +++ b/extensions/src/main/java/dev/cel/extensions/CelRegexRuntimeLibrary.java @@ -0,0 +1,267 @@ +// Copyright 2025 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package dev.cel.extensions; + +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableSet; +import com.google.errorprone.annotations.Immutable; +import com.google.re2j.Matcher; +import com.google.re2j.Pattern; +import com.google.re2j.PatternSyntaxException; +import dev.cel.runtime.CelFunctionBinding; +import dev.cel.runtime.CelLiteRuntimeBuilder; +import dev.cel.runtime.CelLiteRuntimeLibrary; +import java.util.Optional; +import java.util.Set; + +/** Runtime implementation of CEL regex extension functions. */ +@Immutable +public final class CelRegexRuntimeLibrary implements CelLiteRuntimeLibrary { + + /** Enumeration of functions for the Regex runtime extension. */ + public enum Function { + REPLACE( + "regex.replace", + ImmutableSet.of( + CelFunctionBinding.from( + "regex_replaceAll_string_string_string", + ImmutableList.of(String.class, String.class, String.class), + (args) -> { + String target = (String) args[0]; + String pattern = (String) args[1]; + String replaceStr = (String) args[2]; + return CelRegexRuntimeLibrary.replace(target, pattern, replaceStr); + }), + CelFunctionBinding.from( + "regex_replaceCount_string_string_string_int", + ImmutableList.of(String.class, String.class, String.class, Long.class), + (args) -> { + String target = (String) args[0]; + String pattern = (String) args[1]; + String replaceStr = (String) args[2]; + long count = (long) args[3]; + return CelRegexRuntimeLibrary.replaceN(target, pattern, replaceStr, count); + }))), + EXTRACT( + "regex.extract", + ImmutableSet.of( + CelFunctionBinding.from( + "regex_extract_string_string", + String.class, + String.class, + CelRegexRuntimeLibrary::extract))), + EXTRACTALL( + "regex.extractAll", + ImmutableSet.of( + CelFunctionBinding.from( + "regex_extractAll_string_string", + String.class, + String.class, + CelRegexRuntimeLibrary::extractAll))); + + private final String functionName; + private final ImmutableSet functionBindings; + + String getFunction() { + return functionName; + } + + Function(String functionName, ImmutableSet functionBindings) { + this.functionName = functionName; + this.functionBindings = functionBindings; + } + } + + private static final CelRegexRuntimeLibrary VERSION_0 = + new CelRegexRuntimeLibrary(ImmutableSet.copyOf(Function.values())); + + /** Returns the latest version of the 'regex' runtime functions. */ + public static CelRegexRuntimeLibrary regex() { + return VERSION_0; + } + + /** Returns the specified version of the 'regex' runtime functions. */ + public static CelRegexRuntimeLibrary regex(int version) { + switch (version) { + case 0: + case Integer.MAX_VALUE: + return VERSION_0; + default: + throw new IllegalArgumentException("Unsupported 'regex' extension version " + version); + } + } + + /** Returns the 'regex' runtime functions with only the specified functions. */ + public static CelRegexRuntimeLibrary regex(Function... functions) { + return regex(ImmutableSet.copyOf(functions)); + } + + /** Returns the 'regex' runtime functions with only the specified functions. */ + public static CelRegexRuntimeLibrary regex(Set functions) { + return new CelRegexRuntimeLibrary(functions); + } + + private final ImmutableSet functions; + + CelRegexRuntimeLibrary(Set functions) { + this.functions = ImmutableSet.copyOf(functions); + } + + @Override + public void setRuntimeOptions(CelLiteRuntimeBuilder runtimeBuilder) { + runtimeBuilder.addFunctionBindings(newFunctionBindings()); + } + + /** Creates the {@link CelFunctionBinding}s for the configured regex 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 Pattern compileRegexPattern(String regex) { + try { + return Pattern.compile(regex); + } catch (PatternSyntaxException e) { + throw new IllegalArgumentException("Failed to compile regex: " + regex, e); + } + } + + private static String replace(String target, String regex, String replaceStr) { + return replaceN(target, regex, replaceStr, -1); + } + + private static String replaceN( + String target, String regex, String replaceStr, long replaceCount) { + if (replaceCount == 0) { + return target; + } + // For all negative replaceCount, do a replaceAll + if (replaceCount < 0) { + replaceCount = -1; + } + + Pattern pattern = compileRegexPattern(regex); + Matcher matcher = pattern.matcher(target); + StringBuffer sb = new StringBuffer(); + int counter = 0; + + while (matcher.find()) { + if (replaceCount != -1 && counter >= replaceCount) { + break; + } + + String processedReplacement = replaceStrValidator(matcher, replaceStr); + matcher.appendReplacement(sb, Matcher.quoteReplacement(processedReplacement)); + counter++; + } + matcher.appendTail(sb); + + return sb.toString(); + } + + private static String replaceStrValidator(Matcher matcher, String replacement) { + StringBuilder sb = new StringBuilder(); + for (int i = 0; i < replacement.length(); i++) { + char c = replacement.charAt(i); + + if (c != '\\') { + sb.append(c); + continue; + } + + if (i + 1 >= replacement.length()) { + throw new IllegalArgumentException("Invalid replacement string: \\ not allowed at end"); + } + + char nextChar = replacement.charAt(++i); + + if (Character.isDigit(nextChar)) { + int groupNum = Character.digit(nextChar, 10); + int groupCount = matcher.groupCount(); + + if (groupNum > groupCount) { + throw new IllegalArgumentException( + "Replacement string references group " + + groupNum + + " but regex has only " + + groupCount + + " group(s)"); + } + + String groupValue = matcher.group(groupNum); + if (groupValue != null) { + sb.append(groupValue); + } + } else if (nextChar == '\\') { + sb.append('\\'); + } else { + throw new IllegalArgumentException( + "Invalid replacement string: \\ must be followed by a digit"); + } + } + return sb.toString(); + } + + private static Optional extract(String target, String regex) { + Pattern pattern = compileRegexPattern(regex); + Matcher matcher = pattern.matcher(target); + + if (!matcher.find()) { + return Optional.empty(); + } + + int groupCount = matcher.groupCount(); + if (groupCount > 1) { + throw new IllegalArgumentException( + "Regular expression has more than one capturing group: " + regex); + } + + String result = (groupCount == 1) ? matcher.group(1) : matcher.group(0); + + return Optional.ofNullable(result); + } + + private static ImmutableList extractAll(String target, String regex) { + Pattern pattern = compileRegexPattern(regex); + Matcher matcher = pattern.matcher(target); + + if (matcher.groupCount() > 1) { + throw new IllegalArgumentException( + "Regular expression has more than one capturing group: " + regex); + } + + ImmutableList.Builder builder = ImmutableList.builder(); + boolean hasOneGroup = matcher.groupCount() == 1; + + while (matcher.find()) { + if (hasOneGroup) { + String group = matcher.group(1); + if (group != null) { + builder.add(group); + } + } else { + builder.add(matcher.group(0)); + } + } + + 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 4002d2879..0e3bca0fd 100644 --- a/extensions/src/test/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/test/java/dev/cel/extensions/BUILD.bazel @@ -45,6 +45,9 @@ java_library( "//extensions:math_runtime_library", "//extensions:native", "//extensions:optional_library", + "//extensions:regex", + "//extensions:regex_compiler_library", + "//extensions:regex_runtime_library", "//extensions:sets", "//extensions:sets_compiler_library", "//extensions:sets_runtime_library", diff --git a/extensions/src/test/java/dev/cel/extensions/CelRegexExtensionsTest.java b/extensions/src/test/java/dev/cel/extensions/CelRegexExtensionsTest.java index 97d0cc90c..8819d16fd 100644 --- a/extensions/src/test/java/dev/cel/extensions/CelRegexExtensionsTest.java +++ b/extensions/src/test/java/dev/cel/extensions/CelRegexExtensionsTest.java @@ -17,13 +17,23 @@ import static org.junit.Assert.assertThrows; import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableSet; import com.google.testing.junit.testparameterinjector.TestParameter; import com.google.testing.junit.testparameterinjector.TestParameterInjector; import com.google.testing.junit.testparameterinjector.TestParameters; import dev.cel.bundle.Cel; +import dev.cel.bundle.CelFactory; +import dev.cel.common.CelAbstractSyntaxTree; import dev.cel.common.CelFunctionDecl; import dev.cel.common.CelOptions; +import dev.cel.common.CelValidationException; +import dev.cel.compiler.CelCompiler; +import dev.cel.compiler.CelCompilerFactory; import dev.cel.runtime.CelEvaluationException; +import dev.cel.runtime.CelLiteRuntime; +import dev.cel.runtime.CelLiteRuntimeFactory; +import dev.cel.runtime.CelRuntime; +import dev.cel.runtime.CelRuntimeFactory; import java.util.Optional; import org.junit.Test; import org.junit.runner.RunWith; @@ -264,5 +274,150 @@ public void extractAll_multipleCaptureGroups_throwsException(String target, Stri .contains("Regular expression has more than one capturing group:"); } + @Test + public void separateLibraryAndRuntime_allFunctions_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelRegexCompilerLibrary.regex()) + .build(); + CelLiteRuntime celLiteRuntime = + CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .addLibraries(CelRegexRuntimeLibrary.regex()) + .build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("regex.replace('hello world', 'world', 'cel')").getAst(); + Object result = celLiteRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo("hello cel"); + } + + @Test + public void separateLibraryAndRuntime_versioned_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelRegexCompilerLibrary.regex(0)) + .build(); + CelRuntime celRuntime = + CelRuntimeFactory.standardCelRuntimeBuilder() + .addFunctionBindings(CelRegexRuntimeLibrary.regex(0).newFunctionBindings()) + .build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("regex.replace('hello world', 'world', 'cel')").getAst(); + Object result = celRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo("hello cel"); + } + + @Test + public void separateLibraryAndRuntime_subsetOfFunctions_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries(CelRegexCompilerLibrary.regex(CelRegexCompilerLibrary.Function.REPLACE)) + .build(); + CelRuntime celRuntime = + CelRuntimeFactory.standardCelRuntimeBuilder() + .addFunctionBindings( + CelRegexRuntimeLibrary.regex(CelRegexRuntimeLibrary.Function.REPLACE) + .newFunctionBindings()) + .build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("regex.replace('hello world', 'world', 'cel')").getAst(); + Object result = celRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo("hello cel"); + assertThrows( + CelValidationException.class, + () -> celCompiler.compile("regex.extract('hello world', 'world')").getAst()); + } + + @Test + public void separateLibraryAndRuntime_setOfFunctions_success() throws Exception { + CelCompiler celCompiler = + CelCompilerFactory.standardCelCompilerBuilder() + .addLibraries( + CelRegexCompilerLibrary.regex( + ImmutableSet.of(CelRegexCompilerLibrary.Function.REPLACE))) + .build(); + CelRuntime celRuntime = + CelRuntimeFactory.standardCelRuntimeBuilder() + .addFunctionBindings( + CelRegexRuntimeLibrary.regex( + ImmutableSet.of(CelRegexRuntimeLibrary.Function.REPLACE)) + .newFunctionBindings()) + .build(); + + CelAbstractSyntaxTree ast = + celCompiler.compile("regex.replace('hello world', 'world', 'cel')").getAst(); + Object result = celRuntime.createProgram(ast).eval(); + + assertThat(result).isEqualTo("hello cel"); + } + + @Test + public void separateLibraryAndRuntime_unsupportedVersion_throws() { + assertThrows(IllegalArgumentException.class, () -> CelRegexCompilerLibrary.regex(99)); + assertThrows(IllegalArgumentException.class, () -> CelRegexRuntimeLibrary.regex(99)); + } + + @Test + public void regex_subsetOfFunctions_success() throws Exception { + Cel cel = + CelFactory.standardCelBuilder() + .addCompilerLibraries(CelExtensions.regex(CelRegexExtensions.Function.REPLACE)) + .addRuntimeLibraries(CelExtensions.regex(CelRegexExtensions.Function.REPLACE)) + .build(); + + CelAbstractSyntaxTree ast = + cel.compile("regex.replace('hello world', 'world', 'cel')").getAst(); + Object result = cel.createProgram(ast).eval(); + + assertThat(result).isEqualTo("hello cel"); + assertThrows( + CelValidationException.class, + () -> cel.compile("regex.extract('hello world', 'world')").getAst()); + } + + @Test + public void regex_setOfFunctions_success() throws Exception { + Cel cel = + CelFactory.standardCelBuilder() + .addCompilerLibraries( + CelExtensions.regex(ImmutableSet.of(CelRegexExtensions.Function.REPLACE))) + .addRuntimeLibraries( + CelExtensions.regex(ImmutableSet.of(CelRegexExtensions.Function.REPLACE))) + .build(); + + CelAbstractSyntaxTree ast = + cel.compile("regex.replace('hello world', 'world', 'cel')").getAst(); + Object result = cel.createProgram(ast).eval(); + assertThat(result).isEqualTo("hello cel"); + } + + @Test + public void regex_versioned_success() throws Exception { + Cel cel = + CelFactory.standardCelBuilder() + .addCompilerLibraries(CelExtensions.regex(0)) + .addRuntimeLibraries(CelExtensions.regex(0)) + .build(); + + CelAbstractSyntaxTree ast = + cel.compile("regex.replace('hello world', 'world', 'cel')").getAst(); + Object result = cel.createProgram(ast).eval(); + + assertThat(result).isEqualTo("hello cel"); + } + + @Test + public void regex_noArgConstructor_success() { + CelRegexExtensions extensions = new CelRegexExtensions(); + + assertThat(extensions.version()).isEqualTo(CelRegexCompilerLibrary.regex().version()); + assertThat(extensions.functions()).isNotEmpty(); + assertThat(extensions.macros()).isEmpty(); + } } diff --git a/runtime/src/test/java/dev/cel/runtime/BUILD.bazel b/runtime/src/test/java/dev/cel/runtime/BUILD.bazel index c6562ef6d..54f936cb6 100644 --- a/runtime/src/test/java/dev/cel/runtime/BUILD.bazel +++ b/runtime/src/test/java/dev/cel/runtime/BUILD.bazel @@ -197,6 +197,7 @@ cel_android_local_test( "//extensions:encoders_runtime_library_android", "//extensions:lists_runtime_library_android", "//extensions:math_runtime_library_android", + "//extensions:regex_runtime_library_android", "//extensions:sets_runtime_library_android", "//extensions:strings_runtime_library_android", "//runtime:evaluation_exception", diff --git a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java index 63e50543a..58c2f2bf9 100644 --- a/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java +++ b/runtime/src/test/java/dev/cel/runtime/CelLiteRuntimeAndroidTest.java @@ -48,6 +48,7 @@ import dev.cel.extensions.CelEncoderRuntimeLibrary; import dev.cel.extensions.CelListsRuntimeLibrary; import dev.cel.extensions.CelMathRuntimeLibrary; +import dev.cel.extensions.CelRegexRuntimeLibrary; import dev.cel.extensions.CelSetsRuntimeLibrary; import dev.cel.extensions.CelStringRuntimeLibrary; import dev.cel.runtime.standard.EqualsOperator; @@ -791,6 +792,18 @@ public void eval_listsExtension() throws Exception { assertThat(runtime.createProgram(ast).eval()).isEqualTo(ImmutableList.of(1L, 2L)); } + @Test + public void eval_regexExtension() throws Exception { + CelLiteRuntime runtime = + CelLiteRuntimeFactory.newLiteRuntimeBuilder() + .addLibraries(CelRegexRuntimeLibrary.regex()) + .build(); + // Expr: regex.replace('hello world', 'world', 'cel') + CelAbstractSyntaxTree ast = readCheckedExpr("compiled_regex_replace"); + + assertThat(runtime.createProgram(ast).eval()).isEqualTo("hello cel"); + } + 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 448ee48e0..92786afc9 100644 --- a/testing/src/main/java/dev/cel/testing/compiled/BUILD.bazel +++ b/testing/src/main/java/dev/cel/testing/compiled/BUILD.bazel @@ -67,6 +67,7 @@ java_library( ":compiled_proto3_select_repeated_fields", ":compiled_proto3_select_wrappers", ":compiled_proto_message", + ":compiled_regex_replace", ":compiled_sets_contains", ":compiled_string_lower_ascii", ], @@ -133,6 +134,12 @@ compile_cel( expression = "[1, 2, 3].slice(0, 2)", ) +compile_cel( + name = "compiled_regex_replace", + environment = "//testing/environment:all_extensions", + expression = "regex.replace('hello world', 'world', 'cel')", +) + compile_cel( name = "compiled_string_lower_ascii", environment = "//testing/environment:all_extensions", diff --git a/testing/src/test/resources/environment/all_extensions.yaml b/testing/src/test/resources/environment/all_extensions.yaml index 623a0fb5b..034a051b1 100644 --- a/testing/src/test/resources/environment/all_extensions.yaml +++ b/testing/src/test/resources/environment/all_extensions.yaml @@ -20,5 +20,6 @@ extensions: - name: "math" - name: "optional" - name: "protos" +- name: "regex" - name: "sets" - name: "strings"