diff --git a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java index 2e3027dcc..778d041fc 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java +++ b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java @@ -889,14 +889,35 @@ private CelMutableSource normalizeMacroSource( } } - if (exprIdToReplace > 0) { + // 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); + && allExprs.get(replacedId).getKind().equals(ExprKind.Kind.LIST) + && CelNavigableMutableExpr.fromExpr(comprehension.loopStep()) + .allNodes() + .anyMatch(node -> node.id() == replacedId); if (isListExprBeingReplaced) { - unwrapListArgumentsInMacroCallExpr( - allExprs.get(callId).comprehension(), newMacroCallExpr); + unwrapListArgumentsInMacroCallExpr(comprehension, newMacroCallExpr); } } diff --git a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java index 2138ff142..0a168eaab 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java @@ -579,6 +579,25 @@ public void list_replaceElement() throws Exception { assertThat(CEL_UNPARSER.unparse(result.toParsedAst())).isEqualTo("[2, 3, 5]"); } + @Test + public void list_replaceSubtreeWithListInAstWithHasMacro_success() throws Exception { + CelAbstractSyntaxTree ast = CEL.compile("has(msg.single_int64) && 1 in [2]").getAst(); + CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast); + CelMutableExpr foldedList = + CelMutableExpr.ofList( + CelMutableList.create( + 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(); + + assertThat(CEL_UNPARSER.unparse(replacedAst)).isEqualTo("has(msg.single_int64) && 1 in [1, 2]"); + assertConsistentMacroCalls(replacedAst); + } + @Test public void struct_replaceValue() throws Exception { // Tree shape (brackets are expr IDs):