Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions runtime/src/main/java/dev/cel/runtime/planner/BUILD.bazel
Original file line number Diff line number Diff line change
Expand Up @@ -730,6 +730,7 @@ cel_android_library(
"//runtime:resolved_overload_android",
"//runtime:runtime_equality_android",
"@maven//:com_google_errorprone_error_prone_annotations",
"@maven//:com_google_guava_guava",
"@maven//:org_jspecify_jspecify",
"@maven_android//:com_google_guava_guava",
],
Expand Down
28 changes: 22 additions & 6 deletions runtime/src/main/java/dev/cel/runtime/planner/EvalFold.java
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@ final class EvalFold extends PlannedInterpretable {
private final PlannedInterpretable condition;
private final PlannedInterpretable loopStep;
private final PlannedInterpretable result;
private final boolean mutableAccu;

static EvalFold create(
CelExpr expr,
Expand All @@ -49,9 +50,19 @@ static EvalFold create(
PlannedInterpretable iterRange,
PlannedInterpretable loopCondition,
PlannedInterpretable loopStep,
PlannedInterpretable result) {
PlannedInterpretable result,
boolean mutableAccu) {
return new EvalFold(
expr, accuVar, accuInit, iterVar, iterVar2, iterRange, loopCondition, loopStep, result);
expr,
accuVar,
accuInit,
iterVar,
iterVar2,
iterRange,
loopCondition,
loopStep,
result,
mutableAccu);
}

@Override
Expand All @@ -60,7 +71,7 @@ Object evalInternal(GlobalResolver resolver, ExecutionFrame frame) throws CelEva
if (iterRangeRaw instanceof AccumulatedUnknowns) {
return iterRangeRaw;
}
Folder folder = new Folder(resolver, frame, accuInit, accuVar, iterVar, iterVar2);
Folder folder = new Folder(resolver, frame, accuInit, accuVar, iterVar, iterVar2, mutableAccu);

if (iterRangeRaw instanceof Map) {
return evalMap((Map<?, ?>) iterRangeRaw, folder, frame);
Expand Down Expand Up @@ -161,7 +172,8 @@ private EvalFold(
PlannedInterpretable iterRange,
PlannedInterpretable condition,
PlannedInterpretable loopStep,
PlannedInterpretable result) {
PlannedInterpretable result,
boolean mutableAccu) {
super(expr);
this.accuVar = accuVar;
this.accuInit = accuInit;
Expand All @@ -171,6 +183,7 @@ private EvalFold(
this.condition = condition;
this.loopStep = loopStep;
this.result = result;
this.mutableAccu = mutableAccu;
}

private static final class Folder implements ActivationWrapper {
Expand All @@ -180,6 +193,7 @@ private static final class Folder implements ActivationWrapper {
private final String accuVar;
private final String iterVar;
private final String iterVar2;
private final boolean mutableAccu;

private @Nullable Object iterVarVal;
private @Nullable Object iterVar2Val;
Expand All @@ -205,7 +219,7 @@ public boolean isLocallyBound(String name) {
try {
Object initVal = accuInit.eval(resolver, frame);
accuVal =
!computeResult && frame.enableShortCircuiting()
mutableAccu && !computeResult && frame.enableShortCircuiting()
? maybeWrapAccumulator(initVal)
: initVal;
} catch (CelEvaluationException e) {
Expand Down Expand Up @@ -245,13 +259,15 @@ private Folder(
PlannedInterpretable accuInit,
String accuVar,
String iterVar,
String iterVar2) {
String iterVar2,
boolean mutableAccu) {
this.resolver = resolver;
this.frame = frame;
this.accuInit = accuInit;
this.accuVar = accuVar;
this.iterVar = iterVar;
this.iterVar2 = iterVar2;
this.mutableAccu = mutableAccu;
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
import static com.google.common.base.Preconditions.checkNotNull;

import com.google.auto.value.AutoValue;
import com.google.common.annotations.VisibleForTesting;
import com.google.common.base.Strings;
import com.google.common.collect.ImmutableList;
import com.google.common.collect.ImmutableMap;
Expand All @@ -39,6 +40,7 @@
import dev.cel.common.ast.CelExpr.CelSelect;
import dev.cel.common.ast.CelExpr.CelStruct;
import dev.cel.common.ast.CelExpr.CelStruct.Entry;
import dev.cel.common.ast.CelExpr.ExprKind.Kind;
import dev.cel.common.ast.CelReference;
import dev.cel.common.exceptions.CelOverloadNotFoundException;
import dev.cel.common.types.CelKind;
Expand Down Expand Up @@ -68,6 +70,8 @@
@Immutable
@Internal
public final class ProgramPlanner {
private static final String MAP_INSERT_FUNCTION = "cel.@mapInsert";

private final CelTypeProvider typeProvider;
private final CelValueProvider valueProvider;
private final DefaultDispatcher dispatcher;
Expand Down Expand Up @@ -510,7 +514,45 @@ private PlannedInterpretable planComprehension(CelExpr expr, PlannerContext ctx)
iterRange,
loopCondition,
loopStep,
result);
result,
isMutableAccuSafe(comprehension));
}

@VisibleForTesting
static boolean isMutableAccuSafe(CelComprehension comprehension) {
// '@'-prefixed names cannot be written in CEL source, so user expressions cannot alias them.
return comprehension.accuVar().startsWith("@") && isStandardMacroShape(comprehension);
}

@VisibleForTesting
static boolean isStandardMacroShape(CelComprehension comprehension) {
String accuVar = comprehension.accuVar();
if (accuVar.equals(comprehension.iterVar())
|| accuVar.equals(comprehension.iterVar2())
|| !isIdent(comprehension.result(), accuVar)) {
return false;
}
CelCall step = comprehension.loopStep().callOrDefault();
// filter() wraps the accumulation in a ternary: `cond ? accu + [elem] : accu`.
if (step.function().equals(Operator.CONDITIONAL.getFunction())) {
step = step.args().get(1).callOrDefault();
}
if (step.args().isEmpty() || !isIdent(step.args().get(0), accuVar)) {
return false;
}
CelExpr accuInit = comprehension.accuInit();
if (step.function().equals(Operator.ADD.getFunction())) {
return accuInit.getKind() == Kind.LIST
&& accuInit.list().elements().isEmpty()
&& step.args().get(1).listOrDefault().elements().size() == 1;
}
return step.function().equals(MAP_INSERT_FUNCTION)
&& !comprehension.iterVar2().isEmpty()
&& accuInit.getKind() == Kind.MAP;
}

private static boolean isIdent(CelExpr expr, String name) {
return expr.identOrDefault().name().equals(name);
}

/**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
import dev.cel.common.CelSource;
import dev.cel.common.ast.CelConstant;
import dev.cel.common.ast.CelExpr;
import dev.cel.common.ast.CelExpr.CelComprehension;
import dev.cel.common.exceptions.CelDivideByZeroException;
import dev.cel.common.exceptions.CelInvalidArgumentException;
import dev.cel.common.internal.CelDescriptorPool;
Expand Down Expand Up @@ -70,6 +71,7 @@
import dev.cel.expr.conformance.proto3.TestAllTypes;
import dev.cel.expr.conformance.proto3.TestAllTypes.NestedMessage;
import dev.cel.extensions.CelExtensions;
import dev.cel.parser.CelMacro;
import dev.cel.parser.CelStandardMacro;
import dev.cel.runtime.CelAsyncEvaluationOptions;
import dev.cel.runtime.CelAttribute;
Expand Down Expand Up @@ -135,9 +137,47 @@ public final class ProgramPlannerTest {
CelAsyncEvaluationOptions.defaultOptions(),
/* asyncExecutor= */ null);

// Raw comprehensions: compre(iterVar, accuVar, iterRange, accuInit, loopCondition, loopStep,
// result), its two-variable form compre2(iterVar, iterVar2, accuVar, ...), and insert(map, ...)
// for cel.@mapInsert, which cannot be written in CEL source.
private static final ImmutableList<CelMacro> COMPREHENSION_TEST_MACROS =
ImmutableList.of(
CelMacro.newGlobalMacro(
"compre",
7,
(factory, unused, args) ->
Optional.of(
factory.fold(
args.get(0).ident().name(),
args.get(2),
args.get(1).ident().name(),
args.get(3),
args.get(4),
args.get(5),
args.get(6)))),
CelMacro.newGlobalMacro(
"compre2",
8,
(factory, unused, args) ->
Optional.of(
factory.fold(
args.get(0).ident().name(),
args.get(1).ident().name(),
args.get(3),
args.get(2).ident().name(),
args.get(4),
args.get(5),
args.get(6),
args.get(7)))),
CelMacro.newGlobalVarArgMacro(
"insert",
(factory, unused, args) ->
Optional.of(factory.newGlobalCall("cel.@mapInsert", args))));

private static final CelCompiler CEL_COMPILER =
CelCompilerFactory.standardCelCompilerBuilder()
.setStandardMacros(CelStandardMacro.STANDARD_MACROS)
.addMacros(COMPREHENSION_TEST_MACROS)
.addFunctionDeclarations(
newFunctionDeclaration(
"late_bound_func",
Expand Down Expand Up @@ -191,6 +231,7 @@ private static DefaultDispatcher newDispatcher() {
StandardFunction.DIVIDE,
StandardFunction.EQUALS,
StandardFunction.NOT_STRICTLY_FALSE,
StandardFunction.SIZE,
StandardFunction.DYN)
.build();
addBindingsToDispatcher(
Expand Down Expand Up @@ -1014,6 +1055,13 @@ public void plan_select_badPresenceTest_throws() throws Exception {
@TestParameters(
"{expression: 'cel.bind(x, [1, 2], [x + [3], x + [4]]) == [[1, 2, 3], [1, 2, 4]]'}")
@TestParameters("{expression: 'cel.bind(m, {\"a\": 1}, [m]) == [{\"a\": 1}]'}")
@TestParameters(
"{expression: 'compre(i, acc, [1, 2], [], true, acc + [size(acc + [0])], acc) == [1, 2]'}")
@TestParameters(
"{expression: 'compre(i, acc, [1, 2], [10], true, acc + [acc[0]], acc) == [10, 10, 10]'}")
@TestParameters(
"{expression: 'compre(i, acc, [1, 2, 3], [], size(acc + [i]) < 3, acc + [i], acc)"
+ " == [1, 2]'}")
public void plan_comprehension_lists(String expression) throws Exception {
CelAbstractSyntaxTree ast = compile(expression);
Program program = PLANNER.plan(ast);
Expand All @@ -1038,6 +1086,48 @@ public void plan_comprehension_maps(String expression) throws Exception {
assertThat(result).isTrue();
}

@Test
@TestParameters("{expression: '[1].map(x, x)', safe: true}")
@TestParameters("{expression: '[1].filter(x, x > 0)', safe: true}")
@TestParameters("{expression: '{1: 2}.transformMap(k, v, v)', safe: true}")
@TestParameters("{expression: '[1].exists(x, x > 0)', safe: false}")
@TestParameters("{expression: 'compre(i, acc, [1], [], true, acc + [i], acc)', safe: false}")
public void isMutableAccuSafe(String expression, boolean safe) throws Exception {
CelComprehension comprehension =
CEL_COMPILER.parse(expression).getAst().getExpr().comprehension();

assertThat(ProgramPlanner.isMutableAccuSafe(comprehension)).isEqualTo(safe);
}

@Test
@TestParameters("{expression: 'compre(i, a, [1], [], true, a + [i], a)', expected: true}")
@TestParameters(
"{expression: 'compre(i, a, [1], [], true, i > 0 ? a + [i] : a, a)', expected: true}")
@TestParameters(
"{expression: 'compre2(k, v, a, {1: 2}, {}, true, insert(a, k, v), a)', expected: true}")
@TestParameters("{expression: 'compre(a, a, [1], [], true, a + [a], a)', expected: false}")
@TestParameters(
"{expression: 'compre2(k, a, a, {1: 2}, {}, true, insert(a, k, a), a)', expected: false}")
@TestParameters("{expression: 'compre(i, a, [1], [], true, a + [i], [a])', expected: false}")
@TestParameters("{expression: 'compre(i, a, [1], [], true, a, a)', expected: false}")
@TestParameters("{expression: 'compre(i, a, [1], [], true, i + [i], a)', expected: false}")
@TestParameters("{expression: 'compre(i, a, [1], [0], true, a + [i], a)', expected: false}")
@TestParameters("{expression: 'compre(i, a, [1], dyn([]), true, a + [i], a)', expected: false}")
@TestParameters("{expression: 'compre(i, a, [1], [], true, a + a, a)', expected: false}")
@TestParameters("{expression: 'compre(i, a, [1], [], true, a + [i, i], a)', expected: false}")
@TestParameters(
"{expression: 'compre(i, a, [1], {}, true, insert(a, i, i), a)', expected: false}")
@TestParameters(
"{expression: 'compre2(k, v, a, {1: 2}, [], true, insert(a, k, v), a)', expected: false}")
@TestParameters(
"{expression: 'compre2(k, v, a, {1: 2}, {}, true, foo(a, k, v), a)', expected: false}")
public void isStandardMacroShape(String expression, boolean expected) throws Exception {
CelComprehension comprehension =
CEL_COMPILER.parse(expression).getAst().getExpr().comprehension();

assertThat(ProgramPlanner.isStandardMacroShape(comprehension)).isEqualTo(expected);
}

@Test
@TestParameters("{expression: '[1, 2, 3, 4, 5, 6].map(x, x)'}")
@TestParameters("{expression: '[1, 2, 3].map(x, [1, 2].map(y, x + y))'}")
Expand Down
Loading