Skip to content

Commit eb9055b

Browse files
l46kokcopybara-github
authored andcommitted
Avoid O(N^2) list accumulator copies in ProgramPlanner comprehensions
PiperOrigin-RevId: 996735724
1 parent 40d7f3a commit eb9055b

6 files changed

Lines changed: 143 additions & 40 deletions

File tree

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

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -454,6 +454,7 @@ java_library(
454454
"//common/exceptions:runtime_exception",
455455
"//common/values",
456456
"//runtime:accumulated_unknowns",
457+
"//runtime:concatenated_list_view",
457458
"//runtime:evaluation_exception",
458459
"//runtime:interpretable",
459460
"//runtime:interpreter_util",
@@ -1119,6 +1120,7 @@ cel_android_library(
11191120
"//common/exceptions:runtime_exception",
11201121
"//common/values:values_android",
11211122
"//runtime:accumulated_unknowns_android",
1123+
"//runtime:concatenated_list_view",
11221124
"//runtime:evaluation_exception",
11231125
"//runtime:interpretable_android",
11241126
"//runtime:interpreter_util_android",

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

Lines changed: 22 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -62,16 +62,13 @@ Object evalInternal(GlobalResolver resolver, ExecutionFrame frame) throws CelEva
6262
}
6363
Folder folder = new Folder(resolver, frame, accuInit, accuVar, iterVar, iterVar2);
6464

65-
Object result;
6665
if (iterRangeRaw instanceof Map) {
67-
result = evalMap((Map<?, ?>) iterRangeRaw, folder, frame);
66+
return evalMap((Map<?, ?>) iterRangeRaw, folder, frame);
6867
} else if (iterRangeRaw instanceof Collection) {
69-
result = evalList((Collection<?>) iterRangeRaw, folder, frame);
68+
return evalList((Collection<?>) iterRangeRaw, folder, frame);
7069
} else {
7170
throw new IllegalArgumentException("Unexpected iter_range type: " + iterRangeRaw.getClass());
7271
}
73-
74-
return maybeUnwrapAccumulator(result);
7572
}
7673

7774
private Object evalMap(Map<?, ?> iterRange, Folder folder, ExecutionFrame frame)
@@ -94,16 +91,14 @@ private Object evalMap(Map<?, ?> iterRange, Folder folder, ExecutionFrame frame)
9491
}
9592
boolean cond = (boolean) condResult;
9693
if (!cond) {
97-
folder.computeResult = true;
98-
return result.eval(folder, frame);
94+
break;
9995
}
10096

10197
Object stepResult = loopStep.eval(folder, frame);
10298
folder.accuVal = mergeAccumulator(folder.accuVal, stepResult);
10399
folder.initialized = true;
104100
}
105-
folder.computeResult = true;
106-
return result.eval(folder, frame);
101+
return folder.evalResult(result);
107102
}
108103

109104
private Object evalList(Collection<?> iterRange, Folder folder, ExecutionFrame frame)
@@ -129,17 +124,15 @@ private Object evalList(Collection<?> iterRange, Folder folder, ExecutionFrame f
129124
}
130125
boolean cond = (boolean) condResult;
131126
if (!cond) {
132-
folder.computeResult = true;
133-
return result.eval(folder, frame);
127+
break;
134128
}
135129

136130
Object stepResult = loopStep.eval(folder, frame);
137131
folder.accuVal = mergeAccumulator(folder.accuVal, stepResult);
138132
folder.initialized = true;
139133
index++;
140134
}
141-
folder.computeResult = true;
142-
return result.eval(folder, frame);
135+
return folder.evalResult(result);
143136
}
144137

145138
private static Object mergeAccumulator(@Nullable Object currentAccu, Object newVal) {
@@ -159,16 +152,6 @@ private static Object maybeWrapAccumulator(Object val) {
159152
return val;
160153
}
161154

162-
private static Object maybeUnwrapAccumulator(Object val) {
163-
if (val instanceof ConcatenatedListView) {
164-
return ImmutableList.copyOf((ConcatenatedListView<?>) val);
165-
}
166-
if (val instanceof MutableMapValue) {
167-
return ImmutableMap.copyOf((MutableMapValue) val);
168-
}
169-
return val;
170-
}
171-
172155
private EvalFold(
173156
CelExpr expr,
174157
String accuVar,
@@ -220,7 +203,11 @@ public boolean isLocallyBound(String name) {
220203
if (!initialized) {
221204
initialized = true;
222205
try {
223-
accuVal = maybeWrapAccumulator(accuInit.eval(resolver, frame));
206+
Object initVal = accuInit.eval(resolver, frame);
207+
accuVal =
208+
!computeResult && frame.enableShortCircuiting()
209+
? maybeWrapAccumulator(initVal)
210+
: initVal;
224211
} catch (CelEvaluationException e) {
225212
throw new LazyEvaluationRuntimeException(e);
226213
}
@@ -241,6 +228,17 @@ public boolean isLocallyBound(String name) {
241228
return resolver.resolve(name);
242229
}
243230

231+
private Object evalResult(PlannedInterpretable result) throws CelEvaluationException {
232+
computeResult = true;
233+
// Materialize the mutable accumulator so the result expr never observes in-place mutation.
234+
if (accuVal instanceof ConcatenatedListView) {
235+
accuVal = ImmutableList.copyOf((ConcatenatedListView<?>) accuVal);
236+
} else if (accuVal instanceof MutableMapValue) {
237+
accuVal = ImmutableMap.copyOf((MutableMapValue) accuVal);
238+
}
239+
return result.eval(this, frame);
240+
}
241+
244242
private Folder(
245243
GlobalResolver resolver,
246244
ExecutionFrame frame,

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

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
import dev.cel.runtime.CelEvaluationException;
2424
import dev.cel.runtime.CelResolvedOverload;
2525
import dev.cel.runtime.CelUnknownSet;
26+
import dev.cel.runtime.ConcatenatedListView;
2627
import dev.cel.runtime.GlobalResolver;
2728
import dev.cel.runtime.InterpreterUtil;
2829

@@ -130,7 +131,10 @@ static Object maybeAdaptNonStrictArg(Object val) {
130131
* adapts any public {@link CelUnknownSet} instances into internal {@link AccumulatedUnknowns} for
131132
* AST evaluation.
132133
*/
133-
private static Object convertAndAdaptResult(CelValueConverter valueConverter, Object result) {
134+
static Object convertAndAdaptResult(CelValueConverter valueConverter, Object result) {
135+
if (result instanceof ConcatenatedListView) {
136+
return result;
137+
}
134138
return InterpreterUtil.maybeAdaptToAccumulatedUnknowns(
135139
valueConverter.maybeUnwrap(valueConverter.toRuntimeValue(result)));
136140
}

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

Lines changed: 10 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@
3232
final class ExecutionFrame {
3333

3434
private final int comprehensionIterationLimit;
35+
private final boolean enableShortCircuiting;
3536
private final CelFunctionResolver functionResolver;
3637
private final @Nullable PartialVars partialVars;
3738
private final @Nullable CelEvaluationListener listener;
@@ -45,11 +46,7 @@ static ExecutionFrame create(
4546
@Nullable PartialVars partialVars,
4647
@Nullable CelEvaluationListener listener) {
4748
return new ExecutionFrame(
48-
functionResolver,
49-
getComprehensionMaxIterations(celOptions),
50-
partialVars,
51-
listener,
52-
/* asyncTracker= */ null);
49+
functionResolver, celOptions, partialVars, listener, /* asyncTracker= */ null);
5350
}
5451

5552
static ExecutionFrame createForAsync(
@@ -59,12 +56,7 @@ static ExecutionFrame createForAsync(
5956
@Nullable CelEvaluationListener listener,
6057
AsyncCallStateTracker asyncTracker) {
6158
checkNotNull(asyncTracker, "asyncTracker");
62-
return new ExecutionFrame(
63-
functionResolver,
64-
getComprehensionMaxIterations(celOptions),
65-
partialVars,
66-
listener,
67-
asyncTracker);
59+
return new ExecutionFrame(functionResolver, celOptions, partialVars, listener, asyncTracker);
6860
}
6961

7062
private static int getComprehensionMaxIterations(CelOptions celOptions) {
@@ -109,6 +101,10 @@ AsyncCallStateTracker asyncTracker() {
109101
return asyncTracker;
110102
}
111103

104+
boolean enableShortCircuiting() {
105+
return enableShortCircuiting;
106+
}
107+
112108
Optional<PartialVars> partialVars() {
113109
return Optional.ofNullable(partialVars);
114110
}
@@ -119,11 +115,12 @@ Optional<PartialVars> partialVars() {
119115

120116
private ExecutionFrame(
121117
CelFunctionResolver functionResolver,
122-
int limit,
118+
CelOptions celOptions,
123119
@Nullable PartialVars partialVars,
124120
@Nullable CelEvaluationListener listener,
125121
@Nullable AsyncCallStateTracker asyncTracker) {
126-
this.comprehensionIterationLimit = limit;
122+
this.comprehensionIterationLimit = getComprehensionMaxIterations(celOptions);
123+
this.enableShortCircuiting = celOptions.enableShortCircuiting();
127124
this.functionResolver = functionResolver;
128125
this.partialVars = partialVars;
129126
this.listener = listener;

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

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -198,7 +198,7 @@ static Object applyQualifiers(
198198
}
199199
}
200200

201-
return celValueConverter.maybeUnwrap(celValueConverter.toRuntimeValue(obj));
201+
return EvalHelpers.convertAndAdaptResult(celValueConverter, obj);
202202
}
203203

204204
private static Optional<CelAttributePattern> findPartialMatchingPattern(

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

Lines changed: 103 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -166,7 +166,8 @@ public final class ProgramPlannerTest {
166166
newMemberOverload(
167167
"bytes_concat_bytes", SimpleType.BYTES, SimpleType.BYTES, SimpleType.BYTES)))
168168
.addMessageTypes(TestAllTypes.getDescriptor())
169-
.addLibraries(CelExtensions.optional(), CelExtensions.comprehensions())
169+
.addLibraries(
170+
CelExtensions.optional(), CelExtensions.comprehensions(), CelExtensions.bindings())
170171
.setContainer(CEL_CONTAINER)
171172
.build();
172173

@@ -200,6 +201,14 @@ private static DefaultDispatcher newDispatcher() {
200201
DescriptorTypeResolver.create(TYPE_PROVIDER, CelValueConverter.getDefaultInstance()));
201202
addBindingsToDispatcher(
202203
builder, typeFunction.newFunctionBindings(CEL_OPTIONS, RUNTIME_EQUALITY));
204+
addBindingsToDispatcher(
205+
builder,
206+
CelFunctionBinding.fromOverloads(
207+
"cel.@mapInsert",
208+
CelFunctionBinding.from(
209+
"cel_@mapInsert_map_key_value",
210+
ImmutableList.of(Map.class, Object.class, Object.class),
211+
args -> args[0])));
203212

204213
// Custom functions
205214
addBindingsToDispatcher(
@@ -1002,6 +1011,9 @@ public void plan_select_badPresenceTest_throws() throws Exception {
10021011
@TestParameters("{expression: '[1,2,3].exists(i, v, i >= 0 && v > 0) == true'}")
10031012
@TestParameters("{expression: '[1,2,3].exists(i, v, i < 0 || v < 0) == false'}")
10041013
@TestParameters("{expression: '[1,2,3].map(x, x + 1) == [2,3,4]'}")
1014+
@TestParameters(
1015+
"{expression: 'cel.bind(x, [1, 2], [x + [3], x + [4]]) == [[1, 2, 3], [1, 2, 4]]'}")
1016+
@TestParameters("{expression: 'cel.bind(m, {\"a\": 1}, [m]) == [{\"a\": 1}]'}")
10051017
public void plan_comprehension_lists(String expression) throws Exception {
10061018
CelAbstractSyntaxTree ast = compile(expression);
10071019
Program program = PLANNER.plan(ast);
@@ -1016,6 +1028,7 @@ public void plan_comprehension_lists(String expression) throws Exception {
10161028
@TestParameters("{expression: '{\"a\": 1, \"b\": 2}.exists(k, k == \"c\") == false'}")
10171029
@TestParameters("{expression: '{\"a\": \"b\", \"c\": \"c\"}.exists(k, v, k == v)'}")
10181030
@TestParameters("{expression: '{\"a\": 1, \"b\": 2}.exists(k, v, v == 3) == false'}")
1031+
@TestParameters("{expression: '({}.map(k, k) + [1])[0] == 1'}")
10191032
public void plan_comprehension_maps(String expression) throws Exception {
10201033
CelAbstractSyntaxTree ast = compile(expression);
10211034
Program program = PLANNER.plan(ast);
@@ -2054,6 +2067,95 @@ public void plan_exhaustiveConditional_untakenBranchError_evaluatesSuccessfully(
20542067
assertThat(result).isEqualTo(42L);
20552068
}
20562069

2070+
@Test
2071+
public void plan_exhaustiveComprehension_filterDoesNotMutateUntakenBranch() throws Exception {
2072+
CelAbstractSyntaxTree ast = CEL_COMPILER.compile("[1, 2, 3].filter(x, x > 1)").getAst();
2073+
ProgramPlanner planner =
2074+
newPlannerWithOptions(CelOptions.current().enableShortCircuiting(false).build());
2075+
Program program = planner.plan(ast);
2076+
2077+
Object result = program.eval();
2078+
2079+
assertThat(result).isEqualTo(ImmutableList.of(2L, 3L));
2080+
}
2081+
2082+
@Test
2083+
@TestParameters("{expression: '[1, 2, 3].map(x, x + 1) == [2, 3, 4]'}")
2084+
@TestParameters("{expression: '{\"a\": 1, \"b\": 2}.transformMap(k, v, v) == {}'}")
2085+
@SuppressWarnings("Immutable") // Test only
2086+
public void plan_comprehension_accumulationSkipsIntermediateConversion(String expression)
2087+
throws Exception {
2088+
int[] containerConversions = new int[1];
2089+
CelValueConverter countingConverter =
2090+
new CelValueConverter() {
2091+
@Override
2092+
public Object toRuntimeValue(Object value) {
2093+
if (value instanceof Iterable || value instanceof ImmutableMap) {
2094+
containerConversions[0]++;
2095+
}
2096+
return super.toRuntimeValue(value);
2097+
}
2098+
};
2099+
ProgramPlanner planner =
2100+
ProgramPlanner.newPlanner(
2101+
TYPE_PROVIDER,
2102+
VALUE_PROVIDER,
2103+
newDispatcher(),
2104+
countingConverter,
2105+
CEL_CONTAINER,
2106+
CEL_OPTIONS,
2107+
ImmutableSet.of(),
2108+
RUNTIME_EQUALITY,
2109+
CelAsyncEvaluationOptions.defaultOptions(),
2110+
/* asyncExecutor= */ null);
2111+
Program program = planner.plan(CEL_COMPILER.compile(expression).getAst());
2112+
2113+
boolean result = (boolean) program.eval();
2114+
2115+
assertThat(result).isTrue();
2116+
assertThat(containerConversions[0]).isEqualTo(1);
2117+
}
2118+
2119+
@Test
2120+
public void plan_nonStrictFunction_withErrorArg_adaptsToException() throws Exception {
2121+
CelCompiler compiler =
2122+
CelCompilerFactory.standardCelCompilerBuilder()
2123+
.addFunctionDeclarations(
2124+
newFunctionDeclaration(
2125+
"is_error", newGlobalOverload("is_error_int", SimpleType.BOOL, SimpleType.INT)))
2126+
.build();
2127+
DefaultDispatcher.Builder builder = DefaultDispatcher.newBuilder();
2128+
addBindingsToDispatcher(
2129+
builder,
2130+
CelStandardFunctions.newBuilder()
2131+
.includeFunctions(StandardFunction.DIVIDE)
2132+
.build()
2133+
.newFunctionBindings(RUNTIME_EQUALITY, CEL_OPTIONS));
2134+
builder.addOverload(
2135+
"is_error",
2136+
"is_error_int",
2137+
ImmutableList.of(Long.class),
2138+
/* isStrict= */ false,
2139+
args -> args[0] instanceof Exception);
2140+
ProgramPlanner planner =
2141+
ProgramPlanner.newPlanner(
2142+
TYPE_PROVIDER,
2143+
VALUE_PROVIDER,
2144+
builder.build(),
2145+
CEL_VALUE_CONVERTER,
2146+
CEL_CONTAINER,
2147+
CEL_OPTIONS,
2148+
ImmutableSet.of(),
2149+
RUNTIME_EQUALITY,
2150+
CelAsyncEvaluationOptions.defaultOptions(),
2151+
/* asyncExecutor= */ null);
2152+
Program program = planner.plan(compiler.compile("is_error(1 / 0)").getAst());
2153+
2154+
Object result = program.eval();
2155+
2156+
assertThat(result).isEqualTo(true);
2157+
}
2158+
20572159
private static ProgramPlanner newPlannerWithOptions(CelOptions options) {
20582160
return ProgramPlanner.newPlanner(
20592161
TYPE_PROVIDER,

0 commit comments

Comments
 (0)