diff --git a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java index 778d041fc..04db3ab99 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java +++ b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java @@ -889,38 +889,6 @@ private CelMutableSource normalizeMacroSource( } } - // Handle replacing the synthetic single-element list inside a `map` or `filter` loop_step. - // - // Example: `[1].map(x, 10)` - // - In the main AST, `loop_step` is `@result + [10]` (where `[10]` is a synthetic LIST - // node wrapping the body `10`). - // - In `macro_calls`, the call is `[1].map(x, 10)`, which only references the inner `10` - // node—NOT the synthetic `[10]` LIST node. - // - // If a mutation replaces the `[10]` LIST node itself (e.g. with `[20]`), the loop above - // won't find `[10]`'s ID in `macro_calls`, so we unwrap `20` from `[20]` into the macro - // call (`[1].map(x, 20)`). - // - // We must first verify that: - // 1. This macro is actually a COMPREHENSION (e.g. `has(msg.f)` is in `macro_calls` too, - // but expands to a SELECT node, not a COMPREHENSION). - // 2. The replaced LIST node is inside *this* comprehension's `loop_step` (not an unrelated - // list replacement elsewhere in the AST, such as folding `[1, 2] + [3, 4]`). - if (exprIdToReplace > 0 - && allExprs.get(callId).getKind().equals(ExprKind.Kind.COMPREHENSION)) { - long replacedId = idGenerator.generate(exprIdToReplace); - CelMutableComprehension comprehension = allExprs.get(callId).comprehension(); - boolean isListExprBeingReplaced = - allExprs.containsKey(replacedId) - && allExprs.get(replacedId).getKind().equals(ExprKind.Kind.LIST) - && CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) - .allNodes() - .anyMatch(node -> node.id() == replacedId); - if (isListExprBeingReplaced) { - unwrapListArgumentsInMacroCallExpr(comprehension, newMacroCallExpr); - } - } - newMacroSource.addMacroCalls(callId, newMacroCallExpr); } @@ -963,57 +931,6 @@ private CelMutableSource normalizeMacroSource( return newMacroSource; } - /** - * Unwraps the arguments in the extraneous list_expr which is present in the AST but does not - * exist in the macro call map. `map`, `filter` are examples of such. - * - *

This method inspects the comprehension's accumulator initializer to infer that the list_expr - * solely exists to match the expected result type of the macro call signature. - * - * @param comprehension Comprehension in the main AST to extract the macro call arguments from - * (loop step). - * @param newMacroCallExpr (Output parameter) Modified macro call expression with the call - * arguments unwrapped. - */ - private static void unwrapListArgumentsInMacroCallExpr( - CelMutableComprehension comprehension, CelMutableExpr newMacroCallExpr) { - CelMutableExpr accuInit = comprehension.accuInit(); - if (!accuInit.getKind().equals(ExprKind.Kind.LIST) || !accuInit.list().elements().isEmpty()) { - // Does not contain an extraneous list. - return; - } - - CelMutableExpr loopStepExpr = comprehension.loopStep(); - List loopStepArgs = loopStepExpr.call().args(); - if (loopStepArgs.size() != 2 && loopStepArgs.size() != 3) { - throw new IllegalArgumentException( - String.format( - "Expected exactly 2 or 3 arguments but got %d instead on expr id: %d", - loopStepArgs.size(), loopStepExpr.id())); - } - - CelMutableCall existingMacroCall = newMacroCallExpr.call(); - CelMutableCall newMacroCall = - existingMacroCall.target().isPresent() - ? CelMutableCall.create(existingMacroCall.target().get(), existingMacroCall.function()) - : CelMutableCall.create(existingMacroCall.function()); - newMacroCall.addArgs( - existingMacroCall.args().get(0)); // iter_var is first argument of the call by convention - - CelMutableList extraneousList; - if (loopStepArgs.size() == 2) { - extraneousList = loopStepArgs.get(1).list(); - } else { - newMacroCall.addArgs(loopStepArgs.get(0)); - // For map(x,y,z), z is wrapped in a _+_(@result, [z]) - extraneousList = loopStepArgs.get(1).call().args().get(1).list(); - } - - newMacroCall.addArgs(extraneousList.elements()); - - newMacroCallExpr.setCall(newMacroCall); - } - private CelMutableExpr mutateExpr( ExprIdGenerator idGenerator, CelMutableExpr root, diff --git a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java index 0a168eaab..ef3dc85b6 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java @@ -308,77 +308,6 @@ public void replaceSubtree_macroReplacedWithConstExpr_macroCallCleared() throws assertThat(CEL.createProgram(CEL.check(mutatedAst).getAst()).eval()).isEqualTo(1); } - @Test - @SuppressWarnings("unchecked") // Test only - public void replaceSubtree_replaceExtraneousListCreatedByMacro_unparseSuccess() throws Exception { - // Certain macros such as `map` or `filter` generates an extraneous list_expr in the loop step's - // argument that does not exist in the original expression. - // For example, the loop step of this expression looks like: - // CALL [10] { - // function: _+_ - // args: { - // IDENT [8] { - // name: __result__ - // } - // LIST [9] { - // elements: { - // CONSTANT [5] { value: 1 } - // } - // } - // } - // } - CelAbstractSyntaxTree ast = CEL.compile("[1].map(x, 1)").getAst(); - CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast); - CelMutableAst mutableAst2 = CelMutableAst.fromCelAst(ast); - - // These two mutation are equivalent. - CelAbstractSyntaxTree mutatedAstWithList = - AST_MUTATOR - .replaceSubtree( - mutableAst, - CelMutableExpr.ofList( - CelMutableList.create(CelMutableExpr.ofConstant(CelConstant.ofValue(2L)))), - 9L) - .toParsedAst(); - CelAbstractSyntaxTree mutatedAstWithConstant = - AST_MUTATOR - .replaceSubtree(mutableAst2, CelMutableExpr.ofConstant(CelConstant.ofValue(2L)), 5L) - .toParsedAst(); - - assertThat(CEL_UNPARSER.unparse(mutatedAstWithList)).isEqualTo("[1].map(x, 2)"); - assertThat(CEL_UNPARSER.unparse(mutatedAstWithConstant)).isEqualTo("[1].map(x, 2)"); - assertThat((List) CEL.createProgram(CEL.check(mutatedAstWithList).getAst()).eval()) - .containsExactly(2L); - } - - @Test - @SuppressWarnings("unchecked") // Test only - public void replaceSubtree_replaceExtraneousListCreatedByThreeArgMacro_unparseSuccess() - throws Exception { - CelAbstractSyntaxTree ast = CEL.compile("[1].map(x, true, 1)").getAst(); - CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast); - CelMutableAst mutableAst2 = CelMutableAst.fromCelAst(ast); - - // These two mutation are equivalent. - CelAbstractSyntaxTree mutatedAstWithList = - AST_MUTATOR - .replaceSubtree( - mutableAst, - CelMutableExpr.ofList( - CelMutableList.create(CelMutableExpr.ofConstant(CelConstant.ofValue(2L)))), - 10L) - .toParsedAst(); - CelAbstractSyntaxTree mutatedAstWithConstant = - AST_MUTATOR - .replaceSubtree(mutableAst2, CelMutableExpr.ofConstant(CelConstant.ofValue(2L)), 6L) - .toParsedAst(); - - assertThat(CEL_UNPARSER.unparse(mutatedAstWithList)).isEqualTo("[1].map(x, true, 2)"); - assertThat(CEL_UNPARSER.unparse(mutatedAstWithConstant)).isEqualTo("[1].map(x, true, 2)"); - assertThat((List) CEL.createProgram(CEL.check(mutatedAstWithList).getAst()).eval()) - .containsExactly(2L); - } - @Test public void globalCallExpr_replaceRoot() throws Exception { // Tree shape (brackets are expr IDs): @@ -580,8 +509,19 @@ public void list_replaceElement() throws Exception { } @Test - public void list_replaceSubtreeWithListInAstWithHasMacro_success() throws Exception { - CelAbstractSyntaxTree ast = CEL.compile("has(msg.single_int64) && 1 in [2]").getAst(); + @TestParameters( + "{source: 'has(msg.single_int64) && 1 in [2]', exprIdToReplace: 8," + + " expected: 'has(msg.single_int64) && 1 in [1, 2]'}") + @TestParameters( + "{source: '[1].filter(x, x in [2])', exprIdToReplace: 7," + + " expected: '[1].filter(x, x in [1, 2])'}") + @TestParameters("{source: '[1].map(x, [2])', exprIdToReplace: 5, expected: '[1].map(x, [1, 2])'}") + @TestParameters( + "{source: '[1].map(x, true, [2])', exprIdToReplace: 6," + + " expected: '[1].map(x, true, [1, 2])'}") + public void list_replaceSubtreeWithListInAstWithMacro_success( + String source, long exprIdToReplace, String expected) throws Exception { + CelAbstractSyntaxTree ast = CEL.compile(source).getAst(); CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast); CelMutableExpr foldedList = CelMutableExpr.ofList( @@ -589,12 +529,10 @@ public void list_replaceSubtreeWithListInAstWithHasMacro_success() throws Except CelMutableExpr.ofConstant(CelConstant.ofValue(1)), CelMutableExpr.ofConstant(CelConstant.ofValue(2)))); - // Node 8 is `[2]`; replacing it with a LIST triggers normalizeMacroSource while `has(...)` is - // present in macroCalls as a SELECT node rather than a COMPREHENSION node. CelAbstractSyntaxTree replacedAst = - AST_MUTATOR.replaceSubtree(mutableAst, foldedList, 8).toParsedAst(); + AST_MUTATOR.replaceSubtree(mutableAst, foldedList, exprIdToReplace).toParsedAst(); - assertThat(CEL_UNPARSER.unparse(replacedAst)).isEqualTo("has(msg.single_int64) && 1 in [1, 2]"); + assertThat(CEL_UNPARSER.unparse(replacedAst)).isEqualTo(expected); assertConsistentMacroCalls(replacedAst); }