Skip to content

Commit dfe2ef2

Browse files
l46kokcopybara-github
authored andcommitted
Add canonicalization for two-variable comprehensions
PiperOrigin-RevId: 961063857
1 parent d459cc9 commit dfe2ef2

13 files changed

Lines changed: 1110 additions & 464 deletions

verifier/src/main/java/dev/cel/verifier/BUILD.bazel

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -145,6 +145,7 @@ java_library(
145145
"//optimizer:ast_optimizer",
146146
"//optimizer:mutable_ast",
147147
"@maven//:com_google_guava_guava",
148+
"@maven//:org_jspecify_jspecify",
148149
],
149150
)
150151

verifier/src/main/java/dev/cel/verifier/CanonicalizationOptimizer.java

Lines changed: 703 additions & 426 deletions
Large diffs are not rendered by default.

verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -215,9 +215,9 @@ private static final class HasherContext {
215215

216216
private static final class Scope {
217217
final String varName;
218-
final Scope parent;
218+
final @Nullable Scope parent;
219219

220-
Scope(String varName, Scope parent) {
220+
Scope(String varName, @Nullable Scope parent) {
221221
this.varName = varName;
222222
this.parent = parent;
223223
}

verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java

Lines changed: 20 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1202,24 +1202,37 @@ private TranslatedValue reduceAllOrExists(
12021202
}
12031203

12041204
private static boolean isAllMacro(CelComprehension comp) {
1205-
return isBooleanAccuInit(comp, true) && isNotStrictlyFalseLoopCondition(comp);
1205+
return isBooleanAccuInit(comp, true)
1206+
&& isNotStrictlyFalseLoopCondition(comp, /* expectNot= */ false);
12061207
}
12071208

12081209
private static boolean isExistsMacro(CelComprehension comp) {
1209-
return isBooleanAccuInit(comp, false) && isNotStrictlyFalseLoopCondition(comp);
1210+
return isBooleanAccuInit(comp, false)
1211+
&& isNotStrictlyFalseLoopCondition(comp, /* expectNot= */ true);
12101212
}
12111213

12121214
private static boolean isBooleanAccuInit(CelComprehension comp, boolean expectedValue) {
12131215
return comp.accuInit().constantOrDefault().getKind() == CelConstant.Kind.BOOLEAN_VALUE
12141216
&& comp.accuInit().constant().booleanValue() == expectedValue;
12151217
}
12161218

1217-
private static boolean isNotStrictlyFalseLoopCondition(CelComprehension comp) {
1219+
private static boolean isNotStrictlyFalseLoopCondition(CelComprehension comp, boolean expectNot) {
12181220
CelExpr.CelCall call = comp.loopCondition().callOrDefault();
1219-
return (call.function().equals(Operator.NOT_STRICTLY_FALSE.getFunction())
1220-
|| call.function().equals(Operator.OLD_NOT_STRICTLY_FALSE.getFunction()))
1221-
&& call.args().size() == 1
1222-
&& call.args().get(0).identOrDefault().name().equals(comp.accuVar());
1221+
if ((!call.function().equals(Operator.NOT_STRICTLY_FALSE.getFunction())
1222+
&& !call.function().equals(Operator.OLD_NOT_STRICTLY_FALSE.getFunction()))
1223+
|| call.args().size() != 1) {
1224+
return false;
1225+
}
1226+
CelExpr arg = call.args().get(0);
1227+
if (expectNot) {
1228+
CelExpr.CelCall notCall = arg.callOrDefault();
1229+
if (!notCall.function().equals(Operator.LOGICAL_NOT.getFunction())
1230+
|| notCall.args().size() != 1) {
1231+
return false;
1232+
}
1233+
arg = notCall.args().get(0);
1234+
}
1235+
return arg.identOrDefault().name().equals(comp.accuVar());
12231236
}
12241237

12251238
private BoolExpr createTypeConstraint(Expr<?> val, long exprId, CelAbstractSyntaxTree ast) {

verifier/src/main/java/dev/cel/verifier/CelZ3CounterexampleGenerator.java

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -88,7 +88,7 @@ private static String formatExpr(
8888
} else if (decl.equals(typeSystem.timestampCons().ConstructorDecl())) {
8989
return "timestamp(" + formatExpr(ctx, typeSystem, model, expr.getArgs()[0]) + ")";
9090
} else if (decl.equals(typeSystem.durationCons().ConstructorDecl())) {
91-
return "duration(" + formatExpr(ctx, typeSystem, model, expr.getArgs()[0]) + ")";
91+
return "duration('" + formatExpr(ctx, typeSystem, model, expr.getArgs()[0]) + "s')";
9292
} else if (decl.equals(typeSystem.uintCons().ConstructorDecl())) {
9393
return formatExpr(ctx, typeSystem, model, expr.getArgs()[0]) + "u";
9494
} else if (decl.equals(typeSystem.boolCons().ConstructorDecl())) {

verifier/src/main/java/dev/cel/verifier/tools/CelVerifierRepl.java

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -314,15 +314,21 @@ private static void printHelp(String topic, PrintStream out) {
314314
out.println("Declares a variable in the REPL session with a specific type.");
315315
out.println();
316316
out.println("Supported Types:");
317-
out.println(" - Primitive types: int, uint, string, bool, double, bytes");
318-
out.println(" - List types: list<T> (e.g., list<int>, list<string>)");
319-
out.println(" - Map types: map<K,V> (e.g., map<string,int>, map<string,string>)");
317+
out.println(" - Primitive types: int, uint, string, bool, double, bytes, dyn");
318+
out.println(" - Well-known types: timestamp, duration");
319+
out.println(" - List types: list<T> (e.g., list<int>, list<string>)");
320+
out.println(" - Map types: map<K,V> (e.g., map<string,int>, map<string,string>)");
321+
out.println(" - Optional types: optional<T> (e.g., optional<string>, optional<int>)");
322+
out.println(" - Protobuf types: coming soon");
320323
out.println();
321324
out.println("Examples:");
322325
out.println(" cel-verifier> :var role string");
323326
out.println(" cel-verifier> :var port int");
324327
out.println(" cel-verifier> :var scores map<string,int>");
325328
out.println(" cel-verifier> :var tags list<string>");
329+
out.println(" cel-verifier> :var created_at timestamp");
330+
out.println(" cel-verifier> :var timeout duration");
331+
out.println(" cel-verifier> :var opt_flag optional<bool>");
326332
break;
327333
case "unknown":
328334
out.println("Command: :unknown <identifier>");

verifier/src/main/java/dev/cel/verifier/tools/VerificationOptions.java

Lines changed: 30 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
import dev.cel.common.types.CelType;
2222
import dev.cel.common.types.ListType;
2323
import dev.cel.common.types.MapType;
24+
import dev.cel.common.types.OptionalType;
2425
import dev.cel.common.types.SimpleType;
2526
import java.time.Duration;
2627
import java.util.ArrayList;
@@ -150,13 +151,15 @@ static CelType parseCelType(String typeStr) {
150151
Preconditions.checkNotNull(typeStr, "Type string cannot be null.");
151152
String str = typeStr.trim().toLowerCase(Locale.US);
152153

153-
if (str.startsWith("list<") && str.endsWith(">")) {
154+
if ((str.startsWith("list<") && str.endsWith(">"))
155+
|| (str.startsWith("list(") && str.endsWith(")"))) {
154156
String inner = str.substring(5, str.length() - 1).trim();
155157
CelType elemType = parseCelType(inner);
156158
return ListType.create(elemType);
157159
}
158160

159-
if (str.startsWith("map<") && str.endsWith(">")) {
161+
if ((str.startsWith("map<") && str.endsWith(">"))
162+
|| (str.startsWith("map(") && str.endsWith(")"))) {
160163
String inner = str.substring(4, str.length() - 1).trim();
161164
List<String> parts = splitGenericArgs(inner);
162165
if (parts.size() != 2) {
@@ -170,6 +173,20 @@ static CelType parseCelType(String typeStr) {
170173
return MapType.create(keyType, valueType);
171174
}
172175

176+
if ((str.startsWith("optional<") && str.endsWith(">"))
177+
|| (str.startsWith("optional(") && str.endsWith(")"))) {
178+
String inner = str.substring(9, str.length() - 1).trim();
179+
CelType elemType = parseCelType(inner);
180+
return OptionalType.create(elemType);
181+
}
182+
183+
if ((str.startsWith("optional_type<") && str.endsWith(">"))
184+
|| (str.startsWith("optional_type(") && str.endsWith(")"))) {
185+
String inner = str.substring(14, str.length() - 1).trim();
186+
CelType elemType = parseCelType(inner);
187+
return OptionalType.create(elemType);
188+
}
189+
173190
switch (str) {
174191
case "int":
175192
return SimpleType.INT;
@@ -187,12 +204,19 @@ static CelType parseCelType(String typeStr) {
187204
return SimpleType.BYTES;
188205
case "dyn":
189206
return SimpleType.DYN;
207+
case "timestamp":
208+
case "google.protobuf.timestamp":
209+
return SimpleType.TIMESTAMP;
210+
case "duration":
211+
case "google.protobuf.duration":
212+
return SimpleType.DURATION;
190213
default:
214+
// TODO: Support protobuf message types (coming soon).
191215
throw new IllegalArgumentException(
192216
"Unsupported type for CLI variable declaration: '"
193217
+ typeStr
194-
+ "'. Supported types: int, uint, string, bool, double, bytes, dyn, list<T>, map<K,"
195-
+ " V>.");
218+
+ "'. Supported types: int, uint, string, bool, double, bytes, dyn, timestamp,"
219+
+ " duration, list<T>, map<K, V>, optional<T>.");
196220
}
197221
}
198222

@@ -202,10 +226,10 @@ private static List<String> splitGenericArgs(String inner) {
202226
StringBuilder current = new StringBuilder();
203227
for (int i = 0; i < inner.length(); i++) {
204228
char c = inner.charAt(i);
205-
if (c == '<') {
229+
if (c == '<' || c == '(') {
206230
depth++;
207231
current.append(c);
208-
} else if (c == '>') {
232+
} else if (c == '>' || c == ')') {
209233
depth--;
210234
current.append(c);
211235
} else if (c == ',' && depth == 0) {

verifier/src/test/java/dev/cel/verifier/CanonicalizationOptimizerTest.java

Lines changed: 88 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@
2727
import dev.cel.common.CelContainer;
2828
import dev.cel.common.CelMutableAst;
2929
import dev.cel.common.CelOptions;
30+
import dev.cel.common.ast.CelExpr.ExprKind.Kind;
3031
import dev.cel.common.ast.CelMutableExpr;
3132
import dev.cel.common.ast.CelMutableExpr.CelMutableCall;
3233
import dev.cel.common.types.ListType;
@@ -263,16 +264,25 @@ private enum CanonicalizationTestCase {
263264
TWO_VAR_EXISTS_COMMUTATIVE_AND(
264265
"string_int_map.exists(k, v, v == 1 && k == 'foo')",
265266
"string_int_map.exists(k, v, k == \"foo\" && v == 1)"),
267+
TWO_VAR_EXISTS_COMMUTATIVE_AND_REVERSE_ALPHABETICAL_VARS(
268+
"string_int_map.exists(z_key, a_val, a_val == 1 && z_key == 'foo')",
269+
"string_int_map.exists(z_key, a_val, z_key == \"foo\" && a_val == 1)"),
266270
TWO_VAR_ALL_COMMUTATIVE_OR(
267271
"string_int_map.all(k, v, v == 1 || k == 'foo')",
268272
"string_int_map.all(k, v, k == \"foo\" || v == 1)"),
273+
TWO_VAR_ALL_COMMUTATIVE_OR_REVERSE_ALPHABETICAL_VARS(
274+
"string_int_map.all(z_key, a_val, a_val == 1 || z_key == 'foo')",
275+
"string_int_map.all(z_key, a_val, z_key == \"foo\" || a_val == 1)"),
269276
TWO_VAR_EXISTS_SYMMETRIC_EQUALITY(
270277
"string_int_map.exists(k, v, v == 1)", "string_int_map.exists(k, v, v == 1)"),
271278
TWO_VAR_ALL_SYMMETRIC_INEQUALITY(
272279
"string_int_map.all(k, v, v != 0)", "string_int_map.all(k, v, v != 0)"),
273280
TWO_VAR_EXISTS_INT_STRING_MAP(
274281
"int_string_map.exists(k, v, v == 'bar' && k == 1)",
275282
"int_string_map.exists(k, v, k == 1 && v == \"bar\")"),
283+
TWO_VAR_EXISTS_INT_STRING_MAP_REVERSE_ALPHABETICAL_VARS(
284+
"int_string_map.exists(z_key, a_val, a_val == 'bar' && z_key == 1)",
285+
"int_string_map.exists(z_key, a_val, z_key == 1 && a_val == \"bar\")"),
276286
TWO_VAR_ALL_INT_STRING_MAP(
277287
"!int_string_map.all(k, v, k == 1 || v == 'bar')", "k != 1 && v != \"bar\""),
278288
TWO_VAR_EXISTS_LIST_INDEX_VALUE(
@@ -285,10 +295,10 @@ private enum CanonicalizationTestCase {
285295
"int_list.all(i, v, v == 100 || i == 0)", "int_list.all(i, v, i == 0 || v == 100)"),
286296
TWO_VAR_NESTED_COMPREHENSIONS(
287297
"string_int_map.exists(k, v, k == 'foo' && int_list.all(i, e, e == v && i == 0))",
288-
"string_int_map.exists(k, v, k == \"foo\" && int_list.all(i, e, e == v && i == 0))"),
298+
"string_int_map.exists(k, v, k == \"foo\" && int_list.all(i, e, v == e && i == 0))"),
289299
DE_MORGAN_2VAR_NESTED_COMPREHENSIONS(
290300
"string_int_map.exists(k, v, k == 'foo' && !int_list.exists(i, e, e == v))",
291-
"string_int_map.exists(k, v, k == \"foo\" && e != v)"),
301+
"string_int_map.exists(k, v, k == \"foo\" && v != e)"),
292302
TWO_VAR_COMPREHENSION_WITH_OPTIONALS(
293303
"!string_int_map.exists(k, v, optional.of(v).hasValue() && k == 'foo')",
294304
"!optional.of(v).hasValue() || k != \"foo\""),
@@ -304,18 +314,18 @@ private enum CanonicalizationTestCase {
304314
// Extension Coverage - cel.bind Macro
305315
CEL_BIND_COMMUTATIVE_AND(
306316
"cel.bind(x, int_var + 10, 1 == x && 2 == int_var2)",
307-
"cel.bind(x, int_var + 10, int_var2 == 2 && x == 1)"),
317+
"cel.bind(x, int_var + 10, x == 1 && int_var2 == 2)"),
308318
CEL_BIND_COMMUTATIVE_OR(
309319
"cel.bind(x, int_var + 10, 1 == x || 2 == int_var2)",
310-
"cel.bind(x, int_var + 10, int_var2 == 2 || x == 1)"),
320+
"cel.bind(x, int_var + 10, x == 1 || int_var2 == 2)"),
311321
CEL_BIND_SYMMETRIC_EQUALITY(
312322
"cel.bind(x, int_var + 10, 20 == x)", "cel.bind(x, int_var + 10, x == 20)"),
313323
CEL_BIND_NESTED(
314324
"cel.bind(x, int_var + 10, cel.bind(y, int_var2 + 20, 2 == y && 1 == x))",
315325
"cel.bind(x, int_var + 10, cel.bind(y, int_var2 + 20, x == 1 && y == 2))"),
316326
CEL_BIND_DE_MORGAN(
317327
"cel.bind(x, int_var == 1, !(2 == int_var2 && x == true))",
318-
"cel.bind(x, int_var == 1, int_var2 != 2 || x != true)"),
328+
"cel.bind(x, int_var == 1, x != true || int_var2 != 2)"),
319329

320330
// Nested Lists, Maps, and Structs
321331
NESTED_LIST_EQUALITY_SYMMETRY(
@@ -464,7 +474,49 @@ private enum CanonicalizationTestCase {
464474
IDENT_INEQUALITY_SYMMETRY(
465475
"dyn_b != dyn_a || dyn_d != dyn_c", "dyn_a != dyn_b || dyn_c != dyn_d"),
466476
IDENT_SAME_NAME_DIFFERENT_OPERATORS(
467-
"dyn_a != dyn_b && dyn_a == dyn_b", "dyn_a != dyn_b && dyn_a == dyn_b");
477+
"dyn_a != dyn_b && dyn_a == dyn_b", "dyn_a != dyn_b && dyn_a == dyn_b"),
478+
479+
// Comprehension Sorting & Structure Comparison (iterRange, accuInit, loopStep, iterVar2)
480+
COMPREHENSIONS_DIFFERENT_ITER_RANGE_EQUALITY(
481+
"[2, 3].all(x, x > 0) == [1, 2].all(x, x > 0)",
482+
"[1, 2].all(x, x > 0) == [2, 3].all(x, x > 0)"),
483+
COMPREHENSIONS_DIFFERENT_ITER_RANGE_AND(
484+
"[2, 3].all(x, x > 0) && [1, 2].all(x, x > 0)",
485+
"[1, 2].all(x, x > 0) && [2, 3].all(x, x > 0)"),
486+
COMPREHENSIONS_DIFFERENT_PREDICATES_AND(
487+
"[1, 2].all(x, x > 10) && [1, 2].all(x, x > 0)",
488+
"[1, 2].all(x, x > 0) && [1, 2].all(x, x > 10)"),
489+
COMPREHENSIONS_EXISTS_VS_ALL_AND(
490+
"[1, 2].all(x, x == 1) && [1, 2].exists(x, x == 1)",
491+
"[1, 2].exists(x, x == 1) && [1, 2].all(x, x == 1)"),
492+
COMPREHENSIONS_ONE_VAR_VS_TWO_VAR_AND(
493+
"string_int_map.all(k, v, v > 0) && string_int_map.all(k, k == 'a')",
494+
"string_int_map.all(k, k == \"a\") && string_int_map.all(k, v, v > 0)"),
495+
496+
// Macro Scope Coverage (filter, map, exists_one, optMap, optFlatMap)
497+
FILTER_MACRO_PREDICATE_ORDER(
498+
"int_list.filter(x, x > 10 && x > 0)", "int_list.filter(x, x > 0 && x > 10)"),
499+
MAP_MACRO_PREDICATE_ORDER(
500+
"int_list.map(x, x == 2 && x == 1)", "int_list.map(x, x == 1 && x == 2)"),
501+
EXISTS_ONE_MACRO_PREDICATE_ORDER(
502+
"int_list.exists_one(x, x > 10 && x > 0)", "int_list.exists_one(x, x > 0 && x > 10)"),
503+
OPT_MAP_MACRO_PREDICATE_ORDER(
504+
"optional.of(int_var).optMap(x, x == 2 && x == 1)",
505+
"optional.of(int_var).optMap(x, x == 1 && x == 2)"),
506+
OPT_FLAT_MAP_MACRO_PREDICATE_ORDER(
507+
"optional.of(int_var).optFlatMap(x, optional.of(x == 2 && x == 1))",
508+
"optional.of(int_var).optFlatMap(x, optional.of(x == 1 && x == 2))"),
509+
510+
// Literal & Constant Comparator Branches
511+
CONST_UINT_SYMMETRIC_EQUALITY("20u == 10u", "10u == 20u"),
512+
CONST_DOUBLE_SYMMETRIC_EQUALITY("2.5 == 1.5", "1.5 == 2.5"),
513+
CONST_BYTES_SYMMETRIC_EQUALITY(
514+
"b'xyz' == b'abc'", "b\"\\141\\142\\143\" == b\"\\170\\171\\172\""),
515+
MAP_DIFFERENT_KEYS_EQUALITY("{'b': 1} == {'a': 1}", "{\"a\": 1} == {\"b\": 1}"),
516+
MAP_DIFFERENT_VALUES_EQUALITY("{'a': 2} == {'a': 1}", "{\"a\": 1} == {\"a\": 2}"),
517+
LIST_DIFFERENT_ELEMENTS_EQUALITY("[2, 1] == [1, 2]", "[1, 2] == [2, 1]"),
518+
SELECT_DIFFERENT_FIELDS_EQUALITY(
519+
"msg2.single_int64 == msg.single_int64", "msg.single_int64 == msg2.single_int64");
468520

469521
private final String input;
470522
private final String expected;
@@ -563,4 +615,34 @@ public void optimize_customMacroWithExistsStructure_notCanonicalized() throws Ex
563615
.optimizedAst();
564616
assertThat(UNPARSER.unparse(optimizedAst)).isEqualTo("!int_list.my_custom_exists(e, e == 1)");
565617
}
618+
619+
@Test
620+
public void optimize_comprehensionWithoutMacroCalls_deMorganSucceeds() throws Exception {
621+
CelAbstractSyntaxTree ast = CEL.compile("!int_list.exists(e, e == 1)").getAst();
622+
CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast);
623+
mutableAst.source().getMacroCalls().clear();
624+
625+
CelAbstractSyntaxTree optimizedAst =
626+
CanonicalizationOptimizer.newInstance(CanonicalizationOptions.newBuilder().build())
627+
.optimize(mutableAst.toParsedAst(), CEL)
628+
.optimizedAst();
629+
assertThat(optimizedAst.getExpr().getKind()).isEqualTo(Kind.COMPREHENSION);
630+
assertThat(optimizedAst.getExpr().comprehension().accuInit().constant().booleanValue())
631+
.isTrue();
632+
}
633+
634+
@Test
635+
public void optimize_comprehensionAllWithoutMacroCalls_deMorganSucceeds() throws Exception {
636+
CelAbstractSyntaxTree ast = CEL.compile("!int_list.all(e, e == 1)").getAst();
637+
CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast);
638+
mutableAst.source().getMacroCalls().clear();
639+
640+
CelAbstractSyntaxTree optimizedAst =
641+
CanonicalizationOptimizer.newInstance(CanonicalizationOptions.newBuilder().build())
642+
.optimize(mutableAst.toParsedAst(), CEL)
643+
.optimizedAst();
644+
assertThat(optimizedAst.getExpr().getKind()).isEqualTo(Kind.COMPREHENSION);
645+
assertThat(optimizedAst.getExpr().comprehension().accuInit().constant().booleanValue())
646+
.isFalse();
647+
}
566648
}

0 commit comments

Comments
 (0)