diff --git a/bundle/src/main/java/dev/cel/bundle/CelEnvironment.java b/bundle/src/main/java/dev/cel/bundle/CelEnvironment.java index ccbaef61b..6b4684b27 100644 --- a/bundle/src/main/java/dev/cel/bundle/CelEnvironment.java +++ b/bundle/src/main/java/dev/cel/bundle/CelEnvironment.java @@ -85,7 +85,9 @@ public abstract class CelEnvironment { "cel.limit.parse_error_recovery", CelOptions.Builder::maxParseErrorRecoveryLimit, "cel.limit.parse_recursion_depth", - CelOptions.Builder::maxParseRecursionDepth); + CelOptions.Builder::maxParseRecursionDepth, + "cel.limit.expression_node_count", + CelOptions.Builder::maxParseExpressionNodeCount); private static final ImmutableMap FEATURE_HANDLERS = ImmutableMap.of( diff --git a/bundle/src/main/java/dev/cel/bundle/CelEnvironmentExporter.java b/bundle/src/main/java/dev/cel/bundle/CelEnvironmentExporter.java index d233fd36f..6e10edd92 100644 --- a/bundle/src/main/java/dev/cel/bundle/CelEnvironmentExporter.java +++ b/bundle/src/main/java/dev/cel/bundle/CelEnvironmentExporter.java @@ -237,6 +237,11 @@ private void addOptions(CelEnvironment.Builder envBuilder, CelOptions options) { CelEnvironment.Limit.create( "cel.limit.parse_recursion_depth", options.maxParseRecursionDepth())); } + if (options.maxParseExpressionNodeCount() != CelOptions.DEFAULT.maxParseExpressionNodeCount()) { + limits.add( + CelEnvironment.Limit.create( + "cel.limit.expression_node_count", options.maxParseExpressionNodeCount())); + } envBuilder.setLimits(limits.build()); } diff --git a/bundle/src/test/java/dev/cel/bundle/CelEnvironmentExporterTest.java b/bundle/src/test/java/dev/cel/bundle/CelEnvironmentExporterTest.java index ae0de2c18..7560a12aa 100644 --- a/bundle/src/test/java/dev/cel/bundle/CelEnvironmentExporterTest.java +++ b/bundle/src/test/java/dev/cel/bundle/CelEnvironmentExporterTest.java @@ -348,6 +348,7 @@ public void options() { .maxExpressionCodePointSize(100) .maxParseErrorRecoveryLimit(10) .maxParseRecursionDepth(10) + .maxParseExpressionNodeCount(500) .enableQuotedIdentifierSyntax(true) .enableHeterogeneousNumericComparisons(true) .populateMacroCalls(true) @@ -365,6 +366,7 @@ public void options() { .containsExactly( CelEnvironment.Limit.create("cel.limit.expression_code_points", 100), CelEnvironment.Limit.create("cel.limit.parse_error_recovery", 10), - CelEnvironment.Limit.create("cel.limit.parse_recursion_depth", 10)); + CelEnvironment.Limit.create("cel.limit.parse_recursion_depth", 10), + CelEnvironment.Limit.create("cel.limit.expression_node_count", 500)); } } diff --git a/bundle/src/test/java/dev/cel/bundle/CelEnvironmentTest.java b/bundle/src/test/java/dev/cel/bundle/CelEnvironmentTest.java index a5a2f3e6d..a48ea0ff8 100644 --- a/bundle/src/test/java/dev/cel/bundle/CelEnvironmentTest.java +++ b/bundle/src/test/java/dev/cel/bundle/CelEnvironmentTest.java @@ -134,7 +134,8 @@ public void extend_allLimits() throws Exception { .setLimits( CelEnvironment.Limit.create("cel.limit.expression_code_points", 20), CelEnvironment.Limit.create("cel.limit.parse_error_recovery", 10), - CelEnvironment.Limit.create("cel.limit.parse_recursion_depth", 10)) + CelEnvironment.Limit.create("cel.limit.parse_recursion_depth", 10), + CelEnvironment.Limit.create("cel.limit.expression_node_count", 500)) .build(); Cel cel = @@ -147,6 +148,7 @@ public void extend_allLimits() throws Exception { assertThat(checkerOptions.maxExpressionCodePointSize()).isEqualTo(20); assertThat(checkerOptions.maxParseErrorRecoveryLimit()).isEqualTo(10); assertThat(checkerOptions.maxParseRecursionDepth()).isEqualTo(10); + assertThat(checkerOptions.maxParseExpressionNodeCount()).isEqualTo(500); CelAbstractSyntaxTree ast = cel.compile("1 + 2 + 3 + 4 + 5").getAst(); Long result = (Long) cel.createProgram(ast).eval(); @@ -158,6 +160,27 @@ public void extend_allLimits() throws Exception { .contains("expression code point size exceeds limit: size: 21, limit 20"); } + @Test + public void extend_expressionNodeCountLimit() throws Exception { + CelEnvironment environment = + CelEnvironment.newBuilder() + .setLimits(CelEnvironment.Limit.create("cel.limit.expression_node_count", 2)) + .build(); + + Cel cel = + environment.extend( + CelFactory.legacyCelBuilder() + .setStandardMacros(CelStandardMacro.STANDARD_MACROS) + .build(), + CelOptions.DEFAULT); + CelOptions checkerOptions = cel.toCheckerBuilder().options(); + assertThat(checkerOptions.maxParseExpressionNodeCount()).isEqualTo(2); + + CelValidationResult validationResult = cel.compile("1 + 2 + 3"); + assertThat(validationResult.hasError()).isTrue(); + assertThat(validationResult.getErrorString()).contains("expression node limit (2) exceeded"); + } + @Test public void extend_unsupportedFeatureFlag_throws() throws Exception { CelEnvironment environment = diff --git a/common/src/main/java/dev/cel/common/CelOptions.java b/common/src/main/java/dev/cel/common/CelOptions.java index d9c2dd818..c50e40c3f 100644 --- a/common/src/main/java/dev/cel/common/CelOptions.java +++ b/common/src/main/java/dev/cel/common/CelOptions.java @@ -60,6 +60,8 @@ public enum ProtoUnsetFieldOptions { public abstract int maxParseRecursionDepth(); + public abstract int maxParseExpressionNodeCount(); + public abstract boolean populateMacroCalls(); public abstract boolean retainRepeatedUnaryOperators(); @@ -134,6 +136,7 @@ public static Builder newBuilder() { .maxExpressionCodePointSize(100_000) .maxParseErrorRecoveryLimit(30) .maxParseRecursionDepth(250) + .maxParseExpressionNodeCount(100_000) .populateMacroCalls(false) .retainRepeatedUnaryOperators(false) .retainUnbalancedLogicalExpressions(false) @@ -223,6 +226,14 @@ public abstract static class Builder { /** Limit the amount of recursion within parse expressions. */ public abstract Builder maxParseRecursionDepth(int value); + /** + * Set a limit on the number of expression nodes in the abstract syntax tree for the expression. + * This prevents cases where macro expansion results in an AST that is larger than expected from + * the source expression. Once exceeded, the parser will record an error and stop expanding + * macros but continue parsing to report other errors. + */ + public abstract Builder maxParseExpressionNodeCount(int value); + /** Populate macro_calls map in source_info with macro calls parsed from the expression. */ public abstract Builder populateMacroCalls(boolean value); diff --git a/extensions/src/main/java/dev/cel/extensions/CelOptionalLibrary.java b/extensions/src/main/java/dev/cel/extensions/CelOptionalLibrary.java index 8b67d5c79..85ee9c756 100644 --- a/extensions/src/main/java/dev/cel/extensions/CelOptionalLibrary.java +++ b/extensions/src/main/java/dev/cel/extensions/CelOptionalLibrary.java @@ -297,6 +297,7 @@ static CelExtensionLibrary library() { public static final CelOptionalLibrary INSTANCE = CelOptionalLibrary.library().latest(); private static final String UNUSED_ITER_VAR = "#unused"; + private static final String OPTIONAL_MAP_VAR = "@target"; private final int version; private final ImmutableSet functions; @@ -524,21 +525,51 @@ private static Optional expandOptMap( CelExpr mapExpr = checkNotNull(arguments.get(1)); String varName = varIdent.ident().name(); - return Optional.of( + if (target.exprKind().getKind() == CelExpr.ExprKind.Kind.IDENT) { + return Optional.of( + exprFactory.newGlobalCall( + Operator.CONDITIONAL.getFunction(), + exprFactory.newReceiverCall(HAS_VALUE.getFunction(), target), + exprFactory.newGlobalCall( + OPTIONAL_OF.getFunction(), + exprFactory.fold( + UNUSED_ITER_VAR, + exprFactory.newList(), + varName, + exprFactory.newReceiverCall(VALUE.getFunction(), exprFactory.copy(target)), + exprFactory.newBoolLiteral(true), + exprFactory.newIdentifier(varName), + mapExpr)), + exprFactory.newGlobalCall(OPTIONAL_NONE.getFunction()))); + } + + CelExpr localVar = exprFactory.newIdentifier(OPTIONAL_MAP_VAR); + CelExpr localVarCopy = exprFactory.copy(localVar); + CelExpr conditionalExpr = exprFactory.newGlobalCall( Operator.CONDITIONAL.getFunction(), - exprFactory.newReceiverCall(HAS_VALUE.getFunction(), target), + exprFactory.newReceiverCall(HAS_VALUE.getFunction(), localVar), exprFactory.newGlobalCall( OPTIONAL_OF.getFunction(), exprFactory.fold( UNUSED_ITER_VAR, exprFactory.newList(), varName, - exprFactory.newReceiverCall(VALUE.getFunction(), exprFactory.copy(target)), + exprFactory.newReceiverCall(VALUE.getFunction(), localVarCopy), exprFactory.newBoolLiteral(true), exprFactory.newIdentifier(varName), mapExpr)), - exprFactory.newGlobalCall(OPTIONAL_NONE.getFunction()))); + exprFactory.newGlobalCall(OPTIONAL_NONE.getFunction())); + + return Optional.of( + exprFactory.fold( + UNUSED_ITER_VAR, + exprFactory.newList(), + OPTIONAL_MAP_VAR, + target, + exprFactory.newBoolLiteral(false), + exprFactory.newIdentifier(OPTIONAL_MAP_VAR), + conditionalExpr)); } private static Optional expandOptFlatMap( @@ -558,19 +589,47 @@ private static Optional expandOptFlatMap( CelExpr mapExpr = checkNotNull(arguments.get(1)); String varName = varIdent.ident().name(); - return Optional.of( + if (target.exprKind().getKind() == CelExpr.ExprKind.Kind.IDENT) { + return Optional.of( + exprFactory.newGlobalCall( + Operator.CONDITIONAL.getFunction(), + exprFactory.newReceiverCall(HAS_VALUE.getFunction(), target), + exprFactory.fold( + UNUSED_ITER_VAR, + exprFactory.newList(), + varName, + exprFactory.newReceiverCall(VALUE.getFunction(), exprFactory.copy(target)), + exprFactory.newBoolLiteral(true), + exprFactory.newIdentifier(varName), + mapExpr), + exprFactory.newGlobalCall(OPTIONAL_NONE.getFunction()))); + } + + CelExpr localVar = exprFactory.newIdentifier(OPTIONAL_MAP_VAR); + CelExpr localVarCopy = exprFactory.copy(localVar); + CelExpr conditionalExpr = exprFactory.newGlobalCall( Operator.CONDITIONAL.getFunction(), - exprFactory.newReceiverCall(HAS_VALUE.getFunction(), target), + exprFactory.newReceiverCall(HAS_VALUE.getFunction(), localVar), exprFactory.fold( UNUSED_ITER_VAR, exprFactory.newList(), varName, - exprFactory.newReceiverCall(VALUE.getFunction(), exprFactory.copy(target)), + exprFactory.newReceiverCall(VALUE.getFunction(), localVarCopy), exprFactory.newBoolLiteral(true), exprFactory.newIdentifier(varName), mapExpr), - exprFactory.newGlobalCall(OPTIONAL_NONE.getFunction()))); + exprFactory.newGlobalCall(OPTIONAL_NONE.getFunction())); + + return Optional.of( + exprFactory.fold( + UNUSED_ITER_VAR, + exprFactory.newList(), + OPTIONAL_MAP_VAR, + target, + exprFactory.newBoolLiteral(false), + exprFactory.newIdentifier(OPTIONAL_MAP_VAR), + conditionalExpr)); } private static Object indexOptionalMap( diff --git a/extensions/src/test/java/dev/cel/extensions/BUILD.bazel b/extensions/src/test/java/dev/cel/extensions/BUILD.bazel index 9fda186cf..920ba537b 100644 --- a/extensions/src/test/java/dev/cel/extensions/BUILD.bazel +++ b/extensions/src/test/java/dev/cel/extensions/BUILD.bazel @@ -17,6 +17,7 @@ java_library( "//common:compiler_common", "//common:container", "//common:options", + "//common/ast", "//common/exceptions:attribute_not_found", "//common/exceptions:divide_by_zero", "//common/exceptions:index_out_of_bounds", diff --git a/extensions/src/test/java/dev/cel/extensions/CelOptionalLibraryTest.java b/extensions/src/test/java/dev/cel/extensions/CelOptionalLibraryTest.java index 2ba12910f..fab444750 100644 --- a/extensions/src/test/java/dev/cel/extensions/CelOptionalLibraryTest.java +++ b/extensions/src/test/java/dev/cel/extensions/CelOptionalLibraryTest.java @@ -34,6 +34,7 @@ import dev.cel.common.CelOverloadDecl; import dev.cel.common.CelValidationException; import dev.cel.common.CelVarDecl; +import dev.cel.common.ast.CelExpr; import dev.cel.common.types.CelType; import dev.cel.common.types.ListType; import dev.cel.common.types.MapType; @@ -1571,6 +1572,68 @@ public void optionalFlatMapMacro_receiverHasValue_returnsOptionalValue() throws assertThat(result).hasValue(43L); } + @Test + public void optionalMapMacro_simpleTarget_notWrappedInComprehension() throws Exception { + Cel cel = + newCelBuilder() + .addVar("x", OptionalType.create(SimpleType.INT)) + .setResultType(OptionalType.create(SimpleType.INT)) + .build(); + CelAbstractSyntaxTree ast = compile(cel, "x.optMap(y, y + 1)"); + + assertThat(ast.getExpr().exprKind().getKind()).isEqualTo(CelExpr.ExprKind.Kind.CALL); + } + + @Test + public void optionalMapMacro_complexTarget_astWrappedInComprehension() throws Exception { + Cel cel = + newCelBuilder() + .setResultType(OptionalType.create(SimpleType.INT)) + .addVar("msg", StructTypeReference.create(TestAllTypes.getDescriptor().getFullName())) + .build(); + CelAbstractSyntaxTree ast = compile(cel, "msg.?single_int32.optMap(y, y + 1)"); + + assertThat(ast.getExpr().exprKind().getKind()).isEqualTo(CelExpr.ExprKind.Kind.COMPREHENSION); + assertThat(ast.getExpr().comprehension().accuVar()).isEqualTo("@target"); + + Optional result = + (Optional) + cel.createProgram(ast) + .eval(ImmutableMap.of("msg", TestAllTypes.newBuilder().setSingleInt32(42).build())); + assertThat(result).hasValue(43L); + } + + @Test + public void optionalFlatMapMacro_simpleTarget_notWrappedInComprehension() throws Exception { + Cel cel = + newCelBuilder() + .addVar("x", OptionalType.create(SimpleType.INT)) + .setResultType(OptionalType.create(SimpleType.INT)) + .build(); + CelAbstractSyntaxTree ast = compile(cel, "x.optFlatMap(y, optional.of(y + 1))"); + + assertThat(ast.getExpr().exprKind().getKind()).isEqualTo(CelExpr.ExprKind.Kind.CALL); + } + + @Test + public void optionalFlatMapMacro_complexTarget_astWrappedInComprehension() throws Exception { + Cel cel = + newCelBuilder() + .setResultType(OptionalType.create(SimpleType.INT)) + .addVar("msg", StructTypeReference.create(TestAllTypes.getDescriptor().getFullName())) + .build(); + CelAbstractSyntaxTree ast = compile(cel, "msg.?single_int32.optFlatMap(y, optional.of(y + 1))"); + + assertThat(ast.getExpr().exprKind().getKind()).isEqualTo(CelExpr.ExprKind.Kind.COMPREHENSION); + assertThat(ast.getExpr().comprehension().accuVar()).isEqualTo("@target"); + + Optional result = + (Optional) + cel.createProgram(ast) + .eval(ImmutableMap.of("msg", TestAllTypes.newBuilder().setSingleInt32(42).build())); + assertThat(result).hasValue(43L); + } + @Test public void optionalFlatMapMacro_withOptionalOfNonZeroValue_optionalEmptyWhenValueIsZero() throws Exception { diff --git a/parser/src/main/java/dev/cel/parser/Parser.java b/parser/src/main/java/dev/cel/parser/Parser.java index af860e936..240b0af9d 100644 --- a/parser/src/main/java/dev/cel/parser/Parser.java +++ b/parser/src/main/java/dev/cel/parser/Parser.java @@ -153,7 +153,8 @@ static CelValidationResult parse(CelParserImpl parser, CelSource source, CelOpti new ExprFactory( antlrParser, sourceInfo, - options.enableHiddenAccumulatorVar() ? HIDDEN_ACCUMULATOR_NAME : ACCUMULATOR_NAME); + options.enableHiddenAccumulatorVar() ? HIDDEN_ACCUMULATOR_NAME : ACCUMULATOR_NAME, + options.maxParseExpressionNodeCount()); Parser parserImpl = new Parser(parser, options, sourceInfo, exprFactory); ErrorListener errorListener = new ErrorListener(exprFactory); antlrLexer.removeErrorListeners(); @@ -655,6 +656,12 @@ private Optional visitMacro( ImmutableList args, Optional target, CelMacro macro) { + if (exprFactory.isNodeLimitExceeded()) { + return Optional.of( + exprFactory.reportError( + exprFactory.getPosition(expr.id()), + "could not expand macro: expression node limit exceeded")); + } Optional expandedMacro = expandMacro( @@ -1077,16 +1084,20 @@ private static final class ExprFactory extends CelMacroExprFactory { private final ArrayList issues; private final ArrayDeque positions; private final String accumulatorVarName; + private final int maxExpressionNodeCount; + private boolean nodeLimitExceeded; private ExprFactory( org.antlr.v4.runtime.Parser recognizer, CelSource.Builder sourceInfo, - String accumulatorVarName) { + String accumulatorVarName, + int maxExpressionNodeCount) { this.recognizer = recognizer; this.sourceInfo = sourceInfo; this.issues = new ArrayList<>(); this.positions = new ArrayDeque<>(1); // Currently this usually contains at most 1 position. this.accumulatorVarName = accumulatorVarName; + this.maxExpressionNodeCount = maxExpressionNodeCount; } // Implementation of CelExprFactory. @@ -1133,6 +1144,15 @@ private CelExpr reportError(Token token, String message) { return reportError(CelIssue.formatError(getLocation(token), message)); } + @CanIgnoreReturnValue + private CelExpr reportError(int position, String message) { + return reportError(CelIssue.formatError(getLocation(position), message)); + } + + boolean isNodeLimitExceeded() { + return nodeLimitExceeded; + } + // Implementation of CelExprFactory. @Override @@ -1159,17 +1179,17 @@ private int peekPosition() { private long nextExprId(int position) { long exprId = super.nextExprId(); + if (exprId > maxExpressionNodeCount && !nodeLimitExceeded) { + nodeLimitExceeded = true; + reportError( + position, String.format("expression node limit (%d) exceeded", maxExpressionNodeCount)); + } if (position != -1) { sourceInfo.addPositions(exprId, position); } return exprId; } - @Override - public long copyExprId(long id) { - return nextExprId(getPosition(id)); - } - @Override public long nextExprId() { checkState(!positions.isEmpty()); // Should only be called while expanding macros. @@ -1177,6 +1197,11 @@ public long nextExprId() { return nextExprId(peekPosition()); } + @Override + public long copyExprId(long id) { + return nextExprId(getPosition(id)); + } + private List getIssuesList() { return issues; } diff --git a/parser/src/test/java/dev/cel/parser/CelParserImplTest.java b/parser/src/test/java/dev/cel/parser/CelParserImplTest.java index 1e7b44fab..756e97d31 100644 --- a/parser/src/test/java/dev/cel/parser/CelParserImplTest.java +++ b/parser/src/test/java/dev/cel/parser/CelParserImplTest.java @@ -256,6 +256,54 @@ public void parse_exprUnderMaxRecursionLimit_doesNotThrow( assertThat(parseResult.getAst()).isNotNull(); } + @Test + public void parse_nodeLimitExceeded_throws() { + CelParser parser = + CelParserImpl.newBuilder() + .setOptions(CelOptions.newBuilder().maxParseExpressionNodeCount(2).build()) + .build(); + CelValidationResult parseResult = parser.parse("a + b + c"); + + CelValidationException exception = + assertThrows(CelValidationException.class, parseResult::getAst); + assertThat(exception).hasMessageThat().contains("expression node limit (2) exceeded"); + assertThat(exception.getErrors()).hasSize(1); + } + + @Test + public void parse_macroExpansionNodeLimitExceeded_throws() { + CelParser parser = + CelParserImpl.newBuilder() + .setStandardMacros(CelStandardMacro.STANDARD_MACROS) + .setOptions(CelOptions.newBuilder().maxParseExpressionNodeCount(5).build()) + .build(); + CelValidationResult parseResult = parser.parse("[1, 2, 3, 4, 5].map(x, x * 2)"); + + CelValidationException exception = + assertThrows(CelValidationException.class, parseResult::getAst); + assertThat(exception).hasMessageThat().contains("expression node limit (5) exceeded"); + assertThat( + exception.getErrors().stream() + .anyMatch( + issue -> + issue + .getMessage() + .contains("could not expand macro: expression node limit exceeded"))) + .isTrue(); + } + + @Test + public void parse_macroExpansionNodeLimitNotExceeded_success() throws CelValidationException { + CelParser parser = + CelParserImpl.newBuilder() + .setStandardMacros(CelStandardMacro.STANDARD_MACROS) + .setOptions(CelOptions.newBuilder().maxParseExpressionNodeCount(100).build()) + .build(); + CelValidationResult parseResult = parser.parse("[1, 2, 3, 4, 5].map(x, x * 2)"); + assertThat(parseResult.hasError()).isFalse(); + assertThat(parseResult.getAst()).isNotNull(); + } + @Test @TestParameters("{expression: 'A.map(a?b, c)'}") @TestParameters("{expression: 'A.all(a?b, c)'}")