diff --git a/MODULE.bazel b/MODULE.bazel index bec7543d6..925fc335a 100644 --- a/MODULE.bazel +++ b/MODULE.bazel @@ -27,7 +27,7 @@ bazel_dep(name = "rules_java", version = "9.7.0") bazel_dep(name = "rules_android", version = "0.7.3") bazel_dep(name = "rules_shell", version = "0.8.0") bazel_dep(name = "googleapis-java", version = "1.1.5") -bazel_dep(name = "cel-spec", version = "0.25.2", repo_name = "cel_spec") +bazel_dep(name = "cel-spec", version = "0.25.3", repo_name = "cel_spec") bazel_dep(name = "rules_go", version = "0.62.0") # Required by cel-spec to satisfy gazelle transitive dependency diff --git a/conformance/src/test/java/dev/cel/conformance/BUILD.bazel b/conformance/src/test/java/dev/cel/conformance/BUILD.bazel index 4abc705c3..1ccab2dd6 100644 --- a/conformance/src/test/java/dev/cel/conformance/BUILD.bazel +++ b/conformance/src/test/java/dev/cel/conformance/BUILD.bazel @@ -89,6 +89,7 @@ _ALL_TESTS = [ "@cel_spec//tests/simple:testdata/fp_math.textproto", "@cel_spec//tests/simple:testdata/integer_math.textproto", "@cel_spec//tests/simple:testdata/lists.textproto", + "@cel_spec//tests/simple:testdata/lists_ext.textproto", "@cel_spec//tests/simple:testdata/logic.textproto", "@cel_spec//tests/simple:testdata/macros.textproto", "@cel_spec//tests/simple:testdata/macros2.textproto", diff --git a/conformance/src/test/java/dev/cel/conformance/ConformanceTest.java b/conformance/src/test/java/dev/cel/conformance/ConformanceTest.java index 82b4cf812..f1d234101 100644 --- a/conformance/src/test/java/dev/cel/conformance/ConformanceTest.java +++ b/conformance/src/test/java/dev/cel/conformance/ConformanceTest.java @@ -82,7 +82,8 @@ public final class ConformanceTest extends Statement { CelExtensions.bindings(), CelExtensions.comprehensions(), CelExtensions.encoders(OPTIONS), - CelExtensions.math(OPTIONS), + CelExtensions.lists(), + CelExtensions.math(), CelExtensions.protos(), CelExtensions.sets(OPTIONS), CelExtensions.strings(), @@ -93,7 +94,8 @@ public final class ConformanceTest extends Statement { ImmutableList.of( CelExtensions.comprehensions(), CelExtensions.encoders(OPTIONS), - CelExtensions.math(OPTIONS), + CelExtensions.lists(), + CelExtensions.math(), CelExtensions.sets(OPTIONS), CelExtensions.strings(), CelOptionalLibrary.INSTANCE); diff --git a/extensions/src/main/java/dev/cel/extensions/CelListsExtensions.java b/extensions/src/main/java/dev/cel/extensions/CelListsExtensions.java index 91e6f8dc8..3b68e4e6a 100644 --- a/extensions/src/main/java/dev/cel/extensions/CelListsExtensions.java +++ b/extensions/src/main/java/dev/cel/extensions/CelListsExtensions.java @@ -262,7 +262,7 @@ public ImmutableSet functions() { @Override public ImmutableSet macros() { - if (version >= 2) { + if (version >= 2 || (version == -1 && functions.contains(Function.SORT_BY))) { return ImmutableSet.of( CelMacro.newReceiverMacro("sortBy", 2, CelListsExtensions::sortByMacro)); } @@ -340,7 +340,10 @@ private static ImmutableList flatten(Collection list, long depth } public static ImmutableList genRange(long end) { - ImmutableList.Builder builder = ImmutableList.builder(); + checkArgument(end >= 0, "lists.range: size must be non-negative, got %s", end); + checkArgument(end <= 1_000_000, "lists.range: size %s exceeds maximum allowed (1000000)", end); + + ImmutableList.Builder builder = ImmutableList.builderWithExpectedSize((int) end); for (long i = 0; i < end; i++) { builder.add(i); } @@ -425,7 +428,7 @@ public int compare(Object o1, Object o2) { .compare((CelByteString) o1, (CelByteString) o2); } - if (!(o1 instanceof Comparable)) { + if (!(o1 instanceof Comparable) || !(o2 instanceof Comparable)) { throw new IllegalArgumentException("List elements must be comparable"); } if (o1.getClass() != o2.getClass()) { diff --git a/extensions/src/test/java/dev/cel/extensions/CelListsExtensionsTest.java b/extensions/src/test/java/dev/cel/extensions/CelListsExtensionsTest.java index 1f893c7ba..031de9c5d 100644 --- a/extensions/src/test/java/dev/cel/extensions/CelListsExtensionsTest.java +++ b/extensions/src/test/java/dev/cel/extensions/CelListsExtensionsTest.java @@ -175,13 +175,23 @@ public void flattenSingleLevel_listIsSingleLevel_throws(String expression) { @Test @TestParameters("{expression: 'lists.range(9) == [0,1,2,3,4,5,6,7,8]'}") @TestParameters("{expression: 'lists.range(0) == []'}") - @TestParameters("{expression: 'lists.range(-1) == []'}") public void range_success(String expression) throws Exception { boolean result = (boolean) eval(cel, expression); assertThat(result).isTrue(); } + @Test + @TestParameters( + "{expression: 'lists.range(-1)', expectedError: 'lists.range: size must be non-negative'}") + @TestParameters("{expression: 'lists.range(1000001)', expectedError: 'exceeds maximum allowed'}") + public void range_throws(String expression, String expectedError) throws Exception { + CelEvaluationException e = + assertThrows(CelEvaluationException.class, () -> eval(cel, expression)); + + assertThat(e).hasCauseThat().hasMessageThat().contains(expectedError); + } + @Test @TestParameters("{expression: '[].distinct()', expected: '[]'}") @TestParameters("{expression: '[].distinct()', expected: '[]'}") @@ -238,7 +248,10 @@ public void reverse_success(String expression, String expected) throws Exception @TestParameters( "{expression: '[\"d\", \"a\", \"b\", \"c\"].sort()', " + "expected: '[\"a\", \"b\", \"c\", \"d\"]'}") - @TestParameters("{expression: '[b\"b\", b\"a\"].sort()', " + "expected: '[b\"a\", b\"b\"]'}") + @TestParameters("{expression: '[b\"b\", b\"a\"].sort()', expected: '[b\"a\", b\"b\"]'}") + @TestParameters( + "{expression: '[b\"d\", b\"a\", b\"aa\"].sort()', " + + "expected: '[b\"a\", b\"aa\", b\"d\"]'}") @TestParameters( "{expression: '[duration(\"2s\"), duration(\"1s\")].sort()', " + "expected: '[duration(\"1s\"), duration(\"2s\")]'}") @@ -272,11 +285,19 @@ public void sort_success_heterogeneousNumbers(String expression, String expected @TestParameters( "{expression: '[SimpleTest{name: \"a\"}].sort()', " + "expectedError: 'List elements must be comparable'}") + @TestParameters( + "{expression: '[[1, 2, 3]].sort()', " + "expectedError: 'List elements must be comparable'}") + @TestParameters( + "{expression: '[{1: 2}].sort()', " + "expectedError: 'List elements must be comparable'}") + @TestParameters( + "{expression: '[1, null].sort()', " + "expectedError: 'List elements must be comparable'}") + @TestParameters( + "{expression: '[null, 1].sort()', " + "expectedError: 'List elements must be comparable'}") public void sort_throws(String expression, String expectedError) throws Exception { - assertThat(assertThrows(CelEvaluationException.class, () -> eval(cel, expression))) - .hasCauseThat() - .hasMessageThat() - .contains(expectedError); + CelEvaluationException e = + assertThrows(CelEvaluationException.class, () -> eval(cel, expression)); + + assertThat(e).hasCauseThat().hasMessageThat().contains(expectedError); } @Test @@ -331,12 +352,16 @@ public void sortBy_success(String expression, String expected) throws Exception @TestParameters( "{expression: '[SimpleTest{name: \"a\"}, SimpleTest{name: \"b\"}].sortBy(e, e)', " + "expectedError: 'found no matching overload for ''@sortByAssociatedKeys'''}") + @TestParameters( + "{expression: '[1, 2].sortBy(e, [e])', " + + "expectedError: 'found no matching overload for ''@sortByAssociatedKeys'''}") public void sortBy_throws_validationException(String expression, String expectedError) throws Exception { CelValidationResult result = cel.compile(expression); - assertThat(assertThrows(CelValidationException.class, () -> result.getAst())) - .hasMessageThat() - .contains(expectedError); + + CelValidationException e = assertThrows(CelValidationException.class, () -> result.getAst()); + + assertThat(e).hasMessageThat().contains(expectedError); } @Test @@ -345,15 +370,16 @@ public void sortBy_withHomogeneousLiteralValidator_success() throws Exception { CelValidatorFactory.standardCelValidatorBuilder(cel) .addAstValidators(HomogeneousLiteralValidator.newInstance()) .build(); - CelAbstractSyntaxTree ast = cel.compile( "[SimpleTest{name: 'baz'}, SimpleTest{name: 'foo'}, SimpleTest{name: 'bar'}]" + ".sortBy(e, e.name)[0].name") .getAst(); + CelValidationResult result = validator.validate(ast); + Object evalResult = cel.createProgram(ast).eval(); assertThat(result.hasError()).isFalse(); - assertThat(cel.createProgram(ast).eval()).isEqualTo("bar"); + assertThat(evalResult).isEqualTo("bar"); } }