Skip to content

Commit 68fcce7

Browse files
l46kokcopybara-github
authored andcommitted
Restrict ProgramPlanner mutable comprehension accumulators to standard macro shapes
Port of cel-expr/cel-go#1545 PiperOrigin-RevId: 996738501
1 parent a8a98a2 commit 68fcce7

4 files changed

Lines changed: 156 additions & 7 deletions

File tree

‎runtime/src/main/java/dev/cel/runtime/planner/BUILD.bazel‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -730,6 +730,7 @@ cel_android_library(
730730
"//runtime:resolved_overload_android",
731731
"//runtime:runtime_equality_android",
732732
"@maven//:com_google_errorprone_error_prone_annotations",
733+
"@maven//:com_google_guava_guava",
733734
"@maven//:org_jspecify_jspecify",
734735
"@maven_android//:com_google_guava_guava",
735736
],

‎runtime/src/main/java/dev/cel/runtime/planner/EvalFold.java‎

Lines changed: 22 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,7 @@ final class EvalFold extends PlannedInterpretable {
3939
private final PlannedInterpretable condition;
4040
private final PlannedInterpretable loopStep;
4141
private final PlannedInterpretable result;
42+
private final boolean mutableAccu;
4243

4344
static EvalFold create(
4445
CelExpr expr,
@@ -49,9 +50,19 @@ static EvalFold create(
4950
PlannedInterpretable iterRange,
5051
PlannedInterpretable loopCondition,
5152
PlannedInterpretable loopStep,
52-
PlannedInterpretable result) {
53+
PlannedInterpretable result,
54+
boolean mutableAccu) {
5355
return new EvalFold(
54-
expr, accuVar, accuInit, iterVar, iterVar2, iterRange, loopCondition, loopStep, result);
56+
expr,
57+
accuVar,
58+
accuInit,
59+
iterVar,
60+
iterVar2,
61+
iterRange,
62+
loopCondition,
63+
loopStep,
64+
result,
65+
mutableAccu);
5566
}
5667

5768
@Override
@@ -60,7 +71,7 @@ Object evalInternal(GlobalResolver resolver, ExecutionFrame frame) throws CelEva
6071
if (iterRangeRaw instanceof AccumulatedUnknowns) {
6172
return iterRangeRaw;
6273
}
63-
Folder folder = new Folder(resolver, frame, accuInit, accuVar, iterVar, iterVar2);
74+
Folder folder = new Folder(resolver, frame, accuInit, accuVar, iterVar, iterVar2, mutableAccu);
6475

6576
if (iterRangeRaw instanceof Map) {
6677
return evalMap((Map<?, ?>) iterRangeRaw, folder, frame);
@@ -161,7 +172,8 @@ private EvalFold(
161172
PlannedInterpretable iterRange,
162173
PlannedInterpretable condition,
163174
PlannedInterpretable loopStep,
164-
PlannedInterpretable result) {
175+
PlannedInterpretable result,
176+
boolean mutableAccu) {
165177
super(expr);
166178
this.accuVar = accuVar;
167179
this.accuInit = accuInit;
@@ -171,6 +183,7 @@ private EvalFold(
171183
this.condition = condition;
172184
this.loopStep = loopStep;
173185
this.result = result;
186+
this.mutableAccu = mutableAccu;
174187
}
175188

176189
private static final class Folder implements ActivationWrapper {
@@ -180,6 +193,7 @@ private static final class Folder implements ActivationWrapper {
180193
private final String accuVar;
181194
private final String iterVar;
182195
private final String iterVar2;
196+
private final boolean mutableAccu;
183197

184198
private @Nullable Object iterVarVal;
185199
private @Nullable Object iterVar2Val;
@@ -205,7 +219,7 @@ public boolean isLocallyBound(String name) {
205219
try {
206220
Object initVal = accuInit.eval(resolver, frame);
207221
accuVal =
208-
!computeResult && frame.enableShortCircuiting()
222+
mutableAccu && !computeResult && frame.enableShortCircuiting()
209223
? maybeWrapAccumulator(initVal)
210224
: initVal;
211225
} catch (CelEvaluationException e) {
@@ -245,13 +259,15 @@ private Folder(
245259
PlannedInterpretable accuInit,
246260
String accuVar,
247261
String iterVar,
248-
String iterVar2) {
262+
String iterVar2,
263+
boolean mutableAccu) {
249264
this.resolver = resolver;
250265
this.frame = frame;
251266
this.accuInit = accuInit;
252267
this.accuVar = accuVar;
253268
this.iterVar = iterVar;
254269
this.iterVar2 = iterVar2;
270+
this.mutableAccu = mutableAccu;
255271
}
256272
}
257273

‎runtime/src/main/java/dev/cel/runtime/planner/ProgramPlanner.java‎

Lines changed: 43 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
import static com.google.common.base.Preconditions.checkNotNull;
1818

1919
import com.google.auto.value.AutoValue;
20+
import com.google.common.annotations.VisibleForTesting;
2021
import com.google.common.base.Strings;
2122
import com.google.common.collect.ImmutableList;
2223
import com.google.common.collect.ImmutableMap;
@@ -39,6 +40,7 @@
3940
import dev.cel.common.ast.CelExpr.CelSelect;
4041
import dev.cel.common.ast.CelExpr.CelStruct;
4142
import dev.cel.common.ast.CelExpr.CelStruct.Entry;
43+
import dev.cel.common.ast.CelExpr.ExprKind.Kind;
4244
import dev.cel.common.ast.CelReference;
4345
import dev.cel.common.exceptions.CelOverloadNotFoundException;
4446
import dev.cel.common.types.CelKind;
@@ -68,6 +70,8 @@
6870
@Immutable
6971
@Internal
7072
public final class ProgramPlanner {
73+
private static final String MAP_INSERT_FUNCTION = "cel.@mapInsert";
74+
7175
private final CelTypeProvider typeProvider;
7276
private final CelValueProvider valueProvider;
7377
private final DefaultDispatcher dispatcher;
@@ -510,7 +514,45 @@ private PlannedInterpretable planComprehension(CelExpr expr, PlannerContext ctx)
510514
iterRange,
511515
loopCondition,
512516
loopStep,
513-
result);
517+
result,
518+
isMutableAccuSafe(comprehension));
519+
}
520+
521+
@VisibleForTesting
522+
static boolean isMutableAccuSafe(CelComprehension comprehension) {
523+
// '@'-prefixed names cannot be written in CEL source, so user expressions cannot alias them.
524+
return comprehension.accuVar().startsWith("@") && isStandardMacroShape(comprehension);
525+
}
526+
527+
@VisibleForTesting
528+
static boolean isStandardMacroShape(CelComprehension comprehension) {
529+
String accuVar = comprehension.accuVar();
530+
if (accuVar.equals(comprehension.iterVar())
531+
|| accuVar.equals(comprehension.iterVar2())
532+
|| !isIdent(comprehension.result(), accuVar)) {
533+
return false;
534+
}
535+
CelCall step = comprehension.loopStep().callOrDefault();
536+
// filter() wraps the accumulation in a ternary: `cond ? accu + [elem] : accu`.
537+
if (step.function().equals(Operator.CONDITIONAL.getFunction())) {
538+
step = step.args().get(1).callOrDefault();
539+
}
540+
if (step.args().isEmpty() || !isIdent(step.args().get(0), accuVar)) {
541+
return false;
542+
}
543+
CelExpr accuInit = comprehension.accuInit();
544+
if (step.function().equals(Operator.ADD.getFunction())) {
545+
return accuInit.getKind() == Kind.LIST
546+
&& accuInit.list().elements().isEmpty()
547+
&& step.args().get(1).listOrDefault().elements().size() == 1;
548+
}
549+
return step.function().equals(MAP_INSERT_FUNCTION)
550+
&& !comprehension.iterVar2().isEmpty()
551+
&& accuInit.getKind() == Kind.MAP;
552+
}
553+
554+
private static boolean isIdent(CelExpr expr, String name) {
555+
return expr.identOrDefault().name().equals(name);
514556
}
515557

516558
/**

‎runtime/src/test/java/dev/cel/runtime/planner/ProgramPlannerTest.java‎

Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,7 @@
4141
import dev.cel.common.CelSource;
4242
import dev.cel.common.ast.CelConstant;
4343
import dev.cel.common.ast.CelExpr;
44+
import dev.cel.common.ast.CelExpr.CelComprehension;
4445
import dev.cel.common.exceptions.CelDivideByZeroException;
4546
import dev.cel.common.exceptions.CelInvalidArgumentException;
4647
import dev.cel.common.internal.CelDescriptorPool;
@@ -70,6 +71,7 @@
7071
import dev.cel.expr.conformance.proto3.TestAllTypes;
7172
import dev.cel.expr.conformance.proto3.TestAllTypes.NestedMessage;
7273
import dev.cel.extensions.CelExtensions;
74+
import dev.cel.parser.CelMacro;
7375
import dev.cel.parser.CelStandardMacro;
7476
import dev.cel.runtime.CelAsyncEvaluationOptions;
7577
import dev.cel.runtime.CelAttribute;
@@ -135,9 +137,47 @@ public final class ProgramPlannerTest {
135137
CelAsyncEvaluationOptions.defaultOptions(),
136138
/* asyncExecutor= */ null);
137139

140+
// Raw comprehensions: compre(iterVar, accuVar, iterRange, accuInit, loopCondition, loopStep,
141+
// result), its two-variable form compre2(iterVar, iterVar2, accuVar, ...), and insert(map, ...)
142+
// for cel.@mapInsert, which cannot be written in CEL source.
143+
private static final ImmutableList<CelMacro> COMPREHENSION_TEST_MACROS =
144+
ImmutableList.of(
145+
CelMacro.newGlobalMacro(
146+
"compre",
147+
7,
148+
(factory, unused, args) ->
149+
Optional.of(
150+
factory.fold(
151+
args.get(0).ident().name(),
152+
args.get(2),
153+
args.get(1).ident().name(),
154+
args.get(3),
155+
args.get(4),
156+
args.get(5),
157+
args.get(6)))),
158+
CelMacro.newGlobalMacro(
159+
"compre2",
160+
8,
161+
(factory, unused, args) ->
162+
Optional.of(
163+
factory.fold(
164+
args.get(0).ident().name(),
165+
args.get(1).ident().name(),
166+
args.get(3),
167+
args.get(2).ident().name(),
168+
args.get(4),
169+
args.get(5),
170+
args.get(6),
171+
args.get(7)))),
172+
CelMacro.newGlobalVarArgMacro(
173+
"insert",
174+
(factory, unused, args) ->
175+
Optional.of(factory.newGlobalCall("cel.@mapInsert", args))));
176+
138177
private static final CelCompiler CEL_COMPILER =
139178
CelCompilerFactory.standardCelCompilerBuilder()
140179
.setStandardMacros(CelStandardMacro.STANDARD_MACROS)
180+
.addMacros(COMPREHENSION_TEST_MACROS)
141181
.addFunctionDeclarations(
142182
newFunctionDeclaration(
143183
"late_bound_func",
@@ -191,6 +231,7 @@ private static DefaultDispatcher newDispatcher() {
191231
StandardFunction.DIVIDE,
192232
StandardFunction.EQUALS,
193233
StandardFunction.NOT_STRICTLY_FALSE,
234+
StandardFunction.SIZE,
194235
StandardFunction.DYN)
195236
.build();
196237
addBindingsToDispatcher(
@@ -1014,6 +1055,13 @@ public void plan_select_badPresenceTest_throws() throws Exception {
10141055
@TestParameters(
10151056
"{expression: 'cel.bind(x, [1, 2], [x + [3], x + [4]]) == [[1, 2, 3], [1, 2, 4]]'}")
10161057
@TestParameters("{expression: 'cel.bind(m, {\"a\": 1}, [m]) == [{\"a\": 1}]'}")
1058+
@TestParameters(
1059+
"{expression: 'compre(i, acc, [1, 2], [], true, acc + [size(acc + [0])], acc) == [1, 2]'}")
1060+
@TestParameters(
1061+
"{expression: 'compre(i, acc, [1, 2], [10], true, acc + [acc[0]], acc) == [10, 10, 10]'}")
1062+
@TestParameters(
1063+
"{expression: 'compre(i, acc, [1, 2, 3], [], size(acc + [i]) < 3, acc + [i], acc)"
1064+
+ " == [1, 2]'}")
10171065
public void plan_comprehension_lists(String expression) throws Exception {
10181066
CelAbstractSyntaxTree ast = compile(expression);
10191067
Program program = PLANNER.plan(ast);
@@ -1038,6 +1086,48 @@ public void plan_comprehension_maps(String expression) throws Exception {
10381086
assertThat(result).isTrue();
10391087
}
10401088

1089+
@Test
1090+
@TestParameters("{expression: '[1].map(x, x)', safe: true}")
1091+
@TestParameters("{expression: '[1].filter(x, x > 0)', safe: true}")
1092+
@TestParameters("{expression: '{1: 2}.transformMap(k, v, v)', safe: true}")
1093+
@TestParameters("{expression: '[1].exists(x, x > 0)', safe: false}")
1094+
@TestParameters("{expression: 'compre(i, acc, [1], [], true, acc + [i], acc)', safe: false}")
1095+
public void isMutableAccuSafe(String expression, boolean safe) throws Exception {
1096+
CelComprehension comprehension =
1097+
CEL_COMPILER.parse(expression).getAst().getExpr().comprehension();
1098+
1099+
assertThat(ProgramPlanner.isMutableAccuSafe(comprehension)).isEqualTo(safe);
1100+
}
1101+
1102+
@Test
1103+
@TestParameters("{expression: 'compre(i, a, [1], [], true, a + [i], a)', expected: true}")
1104+
@TestParameters(
1105+
"{expression: 'compre(i, a, [1], [], true, i > 0 ? a + [i] : a, a)', expected: true}")
1106+
@TestParameters(
1107+
"{expression: 'compre2(k, v, a, {1: 2}, {}, true, insert(a, k, v), a)', expected: true}")
1108+
@TestParameters("{expression: 'compre(a, a, [1], [], true, a + [a], a)', expected: false}")
1109+
@TestParameters(
1110+
"{expression: 'compre2(k, a, a, {1: 2}, {}, true, insert(a, k, a), a)', expected: false}")
1111+
@TestParameters("{expression: 'compre(i, a, [1], [], true, a + [i], [a])', expected: false}")
1112+
@TestParameters("{expression: 'compre(i, a, [1], [], true, a, a)', expected: false}")
1113+
@TestParameters("{expression: 'compre(i, a, [1], [], true, i + [i], a)', expected: false}")
1114+
@TestParameters("{expression: 'compre(i, a, [1], [0], true, a + [i], a)', expected: false}")
1115+
@TestParameters("{expression: 'compre(i, a, [1], dyn([]), true, a + [i], a)', expected: false}")
1116+
@TestParameters("{expression: 'compre(i, a, [1], [], true, a + a, a)', expected: false}")
1117+
@TestParameters("{expression: 'compre(i, a, [1], [], true, a + [i, i], a)', expected: false}")
1118+
@TestParameters(
1119+
"{expression: 'compre(i, a, [1], {}, true, insert(a, i, i), a)', expected: false}")
1120+
@TestParameters(
1121+
"{expression: 'compre2(k, v, a, {1: 2}, [], true, insert(a, k, v), a)', expected: false}")
1122+
@TestParameters(
1123+
"{expression: 'compre2(k, v, a, {1: 2}, {}, true, foo(a, k, v), a)', expected: false}")
1124+
public void isStandardMacroShape(String expression, boolean expected) throws Exception {
1125+
CelComprehension comprehension =
1126+
CEL_COMPILER.parse(expression).getAst().getExpr().comprehension();
1127+
1128+
assertThat(ProgramPlanner.isStandardMacroShape(comprehension)).isEqualTo(expected);
1129+
}
1130+
10411131
@Test
10421132
@TestParameters("{expression: '[1, 2, 3, 4, 5, 6].map(x, x)'}")
10431133
@TestParameters("{expression: '[1, 2, 3].map(x, [1, 2].map(y, x + y))'}")

0 commit comments

Comments
 (0)