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
Original file line number Diff line number Diff line change
Expand Up @@ -118,7 +118,9 @@ private static void hashAst(CelExpr expr, @Nullable Scope scope, HasherContext c
break;
case LIST:
context.hasher.putInt(expr.list().elements().size());
for (CelExpr elem : expr.list().elements()) {
for (int i = 0; i < expr.list().elements().size(); i++) {
CelExpr elem = expr.list().elements().get(i);
context.hasher.putBoolean(expr.list().optionalIndices().contains(i));
hashAst(elem, scope, context);
}
break;
Expand Down
57 changes: 47 additions & 10 deletions verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java
Original file line number Diff line number Diff line change
Expand Up @@ -284,11 +284,24 @@ private TranslatedValue translateList(CelExpr celExpr, CelAbstractSyntaxTree ast
// check to a trivial identity check (e.g., `list_ref_0 == list_ref_0`).
if (listRef == null) {
SeqExpr seq = ctx.mkEmptySeq(ctx.mkSeqSort(typeSystem.celValueSort()));
for (CelExpr element : createList.elements()) {
ImmutableList<Integer> optionalIndices = createList.optionalIndices();
ImmutableList<CelExpr> elements = createList.elements();
for (int i = 0; i < elements.size(); i++) {
CelExpr element = elements.get(i);
TranslatedValue elem = translateExpr(element, ast);
elementsTv.add(elem);

seq = typeSystem.mkConcatSafe(seq, ctx.mkUnit(elem.z3Expr()));
if (optionalIndices.contains(i)) {
Expr<?> optRef = typeSystem.getOptionalRef(elem.z3Expr());
seq =
(SeqExpr)
ctx.mkITE(
typeSystem.optHasValue(optRef),
typeSystem.mkConcatSafe(seq, ctx.mkUnit(typeSystem.getOptionalValue(optRef))),
seq);
} else {
seq = typeSystem.mkConcatSafe(seq, ctx.mkUnit(elem.z3Expr()));
}
}
listRef = typeSystem.mkListRefConst(LIST_REF_PREFIX);
typeConstraints.add(ctx.mkEq(typeSystem.getSeq(listRef), seq));
Expand Down Expand Up @@ -318,12 +331,24 @@ private TranslatedValue translateMap(CelExpr celExpr, CelAbstractSyntaxTree ast)
Expr<?> value = valueTv.z3Expr();
elementsTv.add(valueTv);

Expr<?> finalValue = value;
BoolExpr finalPresence = ctx.mkTrue();
if (entryAst.optionalEntry()) {
Expr<?> optRef = typeSystem.getOptionalRef(value);
finalPresence = typeSystem.optHasValue(optRef);
finalValue = typeSystem.getOptionalValue(optRef);
}

BoolExpr keyAlreadyPresent = (BoolExpr) ctx.mkSelect(mapPresence, key);
BoolExpr shouldInsertKey = ctx.mkAnd(ctx.mkNot(keyAlreadyPresent), finalPresence);
keysSeq =
ctx.mkITE(keyAlreadyPresent, keysSeq, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key)));
ctx.mkITE(shouldInsertKey, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key)), keysSeq);

mapValues = ctx.mkStore(mapValues, key, value);
mapPresence = ctx.mkStore(mapPresence, key, ctx.mkTrue());
mapValues =
(ArrayExpr) ctx.mkITE(finalPresence, ctx.mkStore(mapValues, key, finalValue), mapValues);
mapPresence =
(ArrayExpr)
ctx.mkITE(finalPresence, ctx.mkStore(mapPresence, key, ctx.mkTrue()), mapPresence);
}

typeConstraints.add(ctx.mkEq(typeSystem.getMapValues(mapRef), mapValues));
Expand Down Expand Up @@ -371,6 +396,14 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a
.orElseGet(() -> extractAstTypeOrDefault(ast, entryAst.value().id()));
Expr<?> defaultVal = getDefaultValueForType(fieldType);

Expr<?> finalValue = value;
BoolExpr optionalHasValue = ctx.mkTrue();
if (entryAst.optionalEntry()) {
Expr<?> optRef = typeSystem.getOptionalRef(value);
optionalHasValue = typeSystem.optHasValue(optRef);
finalValue = typeSystem.getOptionalValue(optRef);
}

// Canonicalization Trick:
//
// We avoid storing explicit default values (e.g. `single_int32: 0`)
Expand All @@ -379,11 +412,13 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a
// (`msg1 == msg2`) to work without using quantifiers (which avoids MBQI loops).
// Because proto3 singular primitives do not have field presence, we also skip setting
// `msgPresence`.
BoolExpr shouldBypass =
fieldType.kind().isPrimitive() ? ctx.mkEq(value, defaultVal) : ctx.mkFalse();
BoolExpr isDefaultPrimitive =
fieldType.kind().isPrimitive() ? ctx.mkEq(finalValue, defaultVal) : ctx.mkFalse();

BoolExpr shouldBypass = ctx.mkOr(ctx.mkNot(optionalHasValue), isDefaultPrimitive);

msgValues =
(ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, value));
(ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, finalValue));

msgPresence =
(ArrayExpr)
Expand Down Expand Up @@ -655,7 +690,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
List<Expr<?>> allRangeElems = new ArrayList<>();

// For statically known list/map literals, unroll them exactly.
if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST) {
if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST
&& iterRangeExpr.list().optionalIndices().isEmpty()) {
ImmutableList<CelExpr> elements = iterRangeExpr.list().elements();
for (int i = 0; i < elements.size(); i++) {
TranslatedValue valueTv = translateExpr(elements.get(i), ast);
Expand All @@ -664,7 +700,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
iterationElements.add(new IterationElement(typeSystem.mkInt(i), value));
allRangeElems.add(value);
}
} else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP) {
} else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP
&& iterRangeExpr.map().entries().stream().noneMatch(CelExpr.CelMap.Entry::optionalEntry)) {
for (CelExpr.CelMap.Entry entry : iterRangeExpr.map().entries()) {
TranslatedValue keyTv = translateExpr(entry.key(), ast);
Expr<?> key = keyTv.z3Expr();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -497,6 +497,11 @@ private BoolExpr getDynamicNumericEquality(Expr<?> z3Expr0, Expr<?> z3Expr1) {
.build(ctx.mkFalse());
}

private boolean hasOptionalElements(TranslatedValue arg) {
return arg.isLiteral(ExprKind.Kind.LIST)
&& !arg.celExpr().get().list().optionalIndices().isEmpty();
}

private BoolExpr unrollListEquality(
TranslatedValue listA, TranslatedValue listB, CelAbstractSyntaxTree ast) {
CelExpr literalListAst =
Expand Down Expand Up @@ -544,7 +549,9 @@ private TranslatedValue translateEquality(
equality = getNumericEquality(arg0, arg1, ast);
} else if (type0.kind() == CelKind.LIST
&& type1.kind() == CelKind.LIST
&& (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST))) {
&& (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST))
&& !hasOptionalElements(arg0)
&& !hasOptionalElements(arg1)) {
equality = unrollListEquality(arg0, arg1, ast);
} else if (isStaticallyKnown(type0) && isStaticallyKnown(type1)) {
equality = typeSystem.getStructuralEquality(z3Arg0, z3Arg1);
Expand All @@ -554,7 +561,9 @@ private TranslatedValue translateEquality(

// Check if one side is an explicit LIST that we can unroll
BoolExpr structuralEq = typeSystem.getStructuralEquality(z3Arg0, z3Arg1);
if (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST)) {
if ((arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST))
&& !hasOptionalElements(arg0)
&& !hasOptionalElements(arg1)) {
structuralEq =
(BoolExpr)
ctx.mkITE(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@
import dev.cel.common.ast.CelExpr.CelCall;
import dev.cel.common.types.ListType;
import dev.cel.common.types.MapType;
import dev.cel.common.types.OptionalType;
import dev.cel.common.types.ProtoMessageTypeProvider;
import dev.cel.common.types.SimpleType;
import dev.cel.common.types.StructTypeReference;
Expand Down Expand Up @@ -100,6 +101,7 @@ public final class CelVerifierZ3ImplTest {
.addVar("dyn_map", MapType.create(SimpleType.DYN, SimpleType.DYN))
.addVar("dyn_var", SimpleType.DYN)
.addVar("dyn_var2", SimpleType.DYN)
.addVar("opt_var", OptionalType.create(SimpleType.INT))
.addVar("string_int_map", MapType.create(SimpleType.STRING, SimpleType.INT))
.addVar("bytes_val", SimpleType.BYTES)
.addVar(
Expand Down Expand Up @@ -1420,7 +1422,15 @@ private enum EquivalenceTestCase {
"has(dyn({'a': 1}).a) && has(dyn(TestAllTypes{single_int32: 1}).single_int32)"),
DYNAMIC_INDEXING_TYPE_MISMATCH(
"type(request) == type(1) && request[1] == 1 && request[2] == 2",
"type(request) == type(1) && 1 / 0 == 1 && request[2] == 2");
"type(request) == type(1) && 1 / 0 == 1 && request[2] == 2"),
OPTIONAL_PRUNE_LIST_LITERAL("[1, ?optional.of(3)]", "[1,3]"),
OPTIONAL_PRUNE_LIST_NONE("[?optional.none(), ?opt_var]", "[?opt_var]"),
OPTIONAL_PRUNE_MAP_NONE("{?1: optional.none()}", "{}"),
OPTIONAL_PRUNE_STRUCT_LIST(
"TestAllTypes{?repeated_int32: optional.of([1, 2])}",
"cel.expr.conformance.proto3.TestAllTypes{repeated_int32: [1, 2]}"),
OPTIONAL_PRUNE_LIST_EQUALITY("[?optional.none(), 1] == [1]", "true"),
OPTIONAL_PRUNE_LIST_COMPREHENSION("[1, ?optional.none()].all(x, x > 0)", "true");

private final String exprA;
private final String exprB;
Expand Down Expand Up @@ -1458,11 +1468,13 @@ private enum EquivalenceViolationTestCase {
HETEROGENEOUS_FIELD_SELECTION(
"test_all_types.single_int32 == 10", "test_all_types.single_int64 == 10"),
STRUCT_VARIABLE_NOT_EQUIVALENT_TO_DEFAULT("test_all_types == TestAllTypes{}", "true"),
OPTIONAL_INVALID_PRUNE_OPT_VAR("[1, ?opt_var]", "[1]"),
CROSS_TYPE_NUMERIC_INEQUALITY_INT_DOUBLE("request == 1.0", "request == 2.0 || request == 1"),
CROSS_TYPE_SYMBOLIC_INEQUALITY_INT_UINT("dyn(x) == dyn(u)", "false"),
CROSS_TYPE_SYMBOLIC_INEQUALITY_UINT_INT("dyn(u) == dyn(x)", "false"),
OPTIONAL_OR_VALUE_VIOLATION("optional.of(x).orValue(y)", "y"),
OPTIONAL_VALUE_VIOLATION("optional.of(x).value()", "y"),
LIST_OPTIONAL_ELEMENTS_COLLISION("[1, ?opt_var]", "[1, opt_var]"),
CROSS_NUMERIC_EQUALITY_INT_DYN_VIOLATION("1 == request", "false");

final String exprA;
Expand Down
Loading