Skip to content

Commit 8e6d76f

Browse files
l46kokcopybara-github
authored andcommitted
Fix floating point comparisons involving infinity/NaN for cross-type numeric comparisons
PiperOrigin-RevId: 955066995
1 parent 0bd9173 commit 8e6d76f

9 files changed

Lines changed: 356 additions & 74 deletions

File tree

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

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -129,6 +129,7 @@ java_library(
129129
"//common/ast",
130130
"//common/ast:cel_block",
131131
"//common/types",
132+
"//common/types:cel_types",
132133
"//common/types:type_providers",
133134
"//verifier/axioms",
134135
"@maven//:com_google_errorprone_error_prone_annotations",

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

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616

1717
import com.google.common.collect.ImmutableList;
1818
import com.google.common.collect.ImmutableSet;
19+
import com.google.common.collect.Iterables;
1920
import com.microsoft.z3.ArithExpr;
2021
import com.microsoft.z3.ArrayExpr;
2122
import com.microsoft.z3.BoolExpr;
@@ -37,6 +38,7 @@
3738
import dev.cel.common.types.CelKind;
3839
import dev.cel.common.types.CelType;
3940
import dev.cel.common.types.CelTypeProvider;
41+
import dev.cel.common.types.CelTypes;
4042
import dev.cel.common.types.ListType;
4143
import dev.cel.common.types.MapType;
4244
import dev.cel.common.types.NullableType;
@@ -79,6 +81,7 @@ final class CelAstToZ3Translator {
7981
private static final String EMPTY_MSG_REF_PREFIX = "!empty_msg_ref_";
8082
private static final String EMPTY_LIST_PREFIX = "!empty_list";
8183
private static final String EMPTY_MAP_PREFIX = "!empty_map";
84+
private static final String NULL_VALUE_FIELD = "null_value";
8285
private final Context ctx;
8386
private final CelZ3TypeSystem typeSystem;
8487
private final CelZ3OperatorTranslator operatorTranslator;
@@ -370,6 +373,10 @@ private TranslatedValue translateMap(CelExpr celExpr, CelAbstractSyntaxTree ast)
370373

371374
private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree ast) {
372375
CelExpr.CelStruct createStruct = celExpr.struct();
376+
if (isJsonWkt(createStruct.messageName())) {
377+
return translateJsonWktStruct(celExpr, createStruct, ast);
378+
}
379+
373380
// Bypass SMT when the struct is empty (return the cached SMT default pointer)
374381
if (createStruct.entries().isEmpty()) {
375382
return TranslatedValue.create(
@@ -448,6 +455,59 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a
448455
return TranslatedValue.propagateStrict(ctx, typeSystem, result, celExpr, elementsTv);
449456
}
450457

458+
private static boolean isJsonWkt(String messageName) {
459+
return messageName.equals(CelTypes.VALUE_MESSAGE)
460+
|| messageName.equals(CelTypes.LIST_VALUE_MESSAGE)
461+
|| messageName.equals(CelTypes.STRUCT_MESSAGE);
462+
}
463+
464+
// Concretize JSON WKT unwrapping directly into native Z3 primitives to avoid
465+
// sort incompatibilities (Message == String) and solver performance penalties (quantifiers).
466+
private TranslatedValue translateJsonWktStruct(
467+
CelExpr celExpr, CelExpr.CelStruct createStruct, CelAbstractSyntaxTree ast) {
468+
Expr<?> fallback;
469+
if (createStruct.messageName().equals(CelTypes.VALUE_MESSAGE)) {
470+
fallback = typeSystem.mkNull();
471+
} else if (createStruct.messageName().equals(CelTypes.LIST_VALUE_MESSAGE)) {
472+
fallback = getDefaultValueForType(ListType.create(SimpleType.DYN));
473+
} else {
474+
fallback = getDefaultValueForType(MapType.create(SimpleType.STRING, SimpleType.DYN));
475+
}
476+
477+
if (createStruct.entries().isEmpty()) {
478+
return TranslatedValue.create(fallback, celExpr, typeSystem, ctx.mkFalse());
479+
}
480+
481+
// JSON WKT messages (gp.Struct, gp.Value, gp.ListValue) can only have a single top-level
482+
// field entry in non-empty creation literals (e.g., 'fields' for Struct, 'values' for
483+
// ListValue, or a single 'oneof' field for Value).
484+
CelExpr.CelStruct.Entry entry = Iterables.getOnlyElement(createStruct.entries());
485+
486+
// Translate the value to properly capture approximations and Optionals
487+
TranslatedValue entryTv = translateExpr(entry.value(), ast);
488+
Expr<?> finalVal = entryTv.z3Expr();
489+
490+
boolean isNullValueField =
491+
createStruct.messageName().equals(CelTypes.VALUE_MESSAGE)
492+
&& entry.fieldKey().equals(NULL_VALUE_FIELD);
493+
494+
if (entry.optionalEntry()) {
495+
Expr<?> optRef = typeSystem.getOptionalRef(finalVal);
496+
BoolExpr hasValue = typeSystem.optHasValue(optRef);
497+
498+
Expr<?> unpackedVal = typeSystem.getOptionalValue(optRef);
499+
if (isNullValueField) {
500+
unpackedVal = typeSystem.mkNull();
501+
}
502+
finalVal = ctx.mkITE(hasValue, unpackedVal, fallback);
503+
} else if (isNullValueField) {
504+
finalVal = typeSystem.mkNull();
505+
}
506+
507+
return TranslatedValue.propagateStrict(
508+
ctx, typeSystem, finalVal, celExpr, ImmutableList.of(entryTv));
509+
}
510+
451511
private Expr<?> getDefaultValueForType(CelType type) {
452512
if (type instanceof NullableType) {
453513
return typeSystem.mkNull();

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

Lines changed: 66 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -805,38 +805,83 @@ private Expr<?> buildMapIndex(
805805

806806
private TranslatedValue translateIndex(
807807
List<TranslatedValue> args, CelAbstractSyntaxTree ast, boolean isOptional) {
808-
Expr<?> lhsTrans = args.get(0).z3Expr();
809-
Expr<?> rhsTrans = args.get(1).z3Expr();
810-
811808
TranslatedValue lhs = args.get(0);
812809
TranslatedValue rhs = args.get(1);
813810
CelType lhsType = extractAstTypeOrDefault(lhs, ast);
814811
CelType rhsType = extractAstTypeOrDefault(rhs, ast);
815812

816-
Expr<?> actualValue;
813+
Expr<?> lhsTrans = lhs.z3Expr();
814+
Expr<?> rhsTrans = rhs.z3Expr();
815+
816+
BoolExpr isLhsOpt = ctx.mkFalse();
817+
BoolExpr lhsHasValue = ctx.mkFalse();
818+
BoolExpr shouldEvaluate = ctx.mkTrue();
819+
820+
if (isOptional) {
821+
isLhsOpt = typeSystem.isOptional(lhsTrans);
822+
Expr<?> optRef = typeSystem.getOptionalRef(lhsTrans);
823+
lhsHasValue = typeSystem.optHasValue(optRef);
824+
825+
lhsTrans = ctx.mkITE(isLhsOpt, typeSystem.getOptionalValue(optRef), lhsTrans);
826+
shouldEvaluate = (BoolExpr) ctx.mkITE(isLhsOpt, lhsHasValue, ctx.mkTrue());
827+
828+
if (lhsType instanceof OptionalType) {
829+
lhsType = lhsType.parameters().get(0);
830+
}
831+
}
832+
833+
Expr<?> actualValue =
834+
buildAndConstrainIndex(lhsTrans, rhsTrans, lhsType, rhsType, shouldEvaluate, isOptional);
835+
836+
if (isOptional) {
837+
actualValue =
838+
ctx.mkITE(
839+
ctx.mkAnd(isLhsOpt, ctx.mkNot(lhsHasValue)),
840+
typeSystem.mkOptionalNone(),
841+
actualValue);
842+
}
843+
844+
return TranslatedValue.propagateStrict(ctx, typeSystem, actualValue, args);
845+
}
846+
847+
private Expr<?> buildAndConstrainIndex(
848+
Expr<?> lhsTrans,
849+
Expr<?> rhsTrans,
850+
CelType lhsType,
851+
CelType rhsType,
852+
BoolExpr shouldEvaluate,
853+
boolean isOptional) {
854+
CelType expectedElemType = null;
817855
if (lhsType.kind() == CelKind.LIST && rhsType.kind() == CelKind.INT) {
818-
actualValue = buildListIndex(lhsTrans, rhsTrans, ctx.mkTrue(), isOptional);
819-
constraintSink.accept(
820-
ctx.mkImplies(
821-
ctx.mkNot(typeSystem.isError(actualValue)),
822-
typeConstraintGenerator.apply(actualValue, ((ListType) lhsType).elemType())));
856+
expectedElemType = ((ListType) lhsType).elemType();
823857
} else if (lhsType.kind() == CelKind.MAP) {
824-
actualValue = buildMapIndex(lhsTrans, rhsTrans, ctx.mkTrue(), isOptional);
858+
expectedElemType = ((MapType) lhsType).valueType();
859+
}
860+
861+
if (expectedElemType != null) {
862+
Expr<?> actualValue =
863+
lhsType.kind() == CelKind.LIST
864+
? buildListIndex(lhsTrans, rhsTrans, shouldEvaluate, isOptional)
865+
: buildMapIndex(lhsTrans, rhsTrans, shouldEvaluate, isOptional);
866+
867+
CelType finalType = isOptional ? OptionalType.create(expectedElemType) : expectedElemType;
868+
825869
constraintSink.accept(
826870
ctx.mkImplies(
827-
ctx.mkNot(typeSystem.isError(actualValue)),
828-
typeConstraintGenerator.apply(actualValue, ((MapType) lhsType).valueType())));
829-
} else {
830-
BoolExpr isListGuard = ctx.mkAnd(typeSystem.isList(lhsTrans), typeSystem.isInt(rhsTrans));
831-
BoolExpr isMapGuard = typeSystem.isMap(lhsTrans);
832-
actualValue =
833-
CelZ3TypeSystem.SwitchBuilder.newBuilder(ctx)
834-
.addCase(isListGuard, buildListIndex(lhsTrans, rhsTrans, isListGuard, isOptional))
835-
.addCase(isMapGuard, buildMapIndex(lhsTrans, rhsTrans, isMapGuard, isOptional))
836-
.build(typeSystem.mkError());
871+
ctx.mkAnd(shouldEvaluate, ctx.mkNot(typeSystem.isError(actualValue))),
872+
typeConstraintGenerator.apply(actualValue, finalType)));
873+
874+
return actualValue;
837875
}
838876

839-
return TranslatedValue.propagateStrict(ctx, typeSystem, actualValue, args);
877+
BoolExpr isListGuard =
878+
ctx.mkAnd(shouldEvaluate, typeSystem.isList(lhsTrans), typeSystem.isInt(rhsTrans));
879+
BoolExpr isMapGuard = ctx.mkAnd(shouldEvaluate, typeSystem.isMap(lhsTrans));
880+
881+
return CelZ3TypeSystem.SwitchBuilder.newBuilder(ctx)
882+
.addCase(isListGuard, buildListIndex(lhsTrans, rhsTrans, isListGuard, isOptional))
883+
.addCase(isMapGuard, buildMapIndex(lhsTrans, rhsTrans, isMapGuard, isOptional))
884+
.build(typeSystem.mkError());
840885
}
841886

842887
private TranslatedValue translateConditional(

verifier/src/main/java/dev/cel/verifier/axioms/AxiomHelpers.java

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,9 @@
1616

1717
import com.microsoft.z3.BoolExpr;
1818
import com.microsoft.z3.Context;
19+
import com.microsoft.z3.FPExpr;
1920
import com.microsoft.z3.IntExpr;
21+
import com.microsoft.z3.RealExpr;
2022

2123
/** Helper methods for Z3 axioms operations. */
2224
final class AxiomHelpers {
@@ -47,5 +49,51 @@ static IntExpr mkTruncatedMod(Context ctx, IntExpr a, IntExpr b) {
4749
return (IntExpr) ctx.mkSub(a, ctx.mkMul(mkTruncatedDiv(ctx, a, b), b));
4850
}
4951

52+
/**
53+
* Safe comparison between a Z3 Real (from int/uint) and a Z3 FloatingPoint (double) for {@code
54+
* <}.
55+
*/
56+
static BoolExpr mkRealLtFp(Context ctx, RealExpr real, FPExpr fp) {
57+
return mkSafeFpComparison(ctx, fp, isPosInf(ctx, fp), ctx.mkLt(real, ctx.mkFPToReal(fp)));
58+
}
59+
60+
/**
61+
* Safe comparison between a Z3 FloatingPoint (double) and a Z3 Real (from int/uint) for {@code
62+
* <}.
63+
*/
64+
static BoolExpr mkFpLtReal(Context ctx, FPExpr fp, RealExpr real) {
65+
return mkSafeFpComparison(ctx, fp, isNegInf(ctx, fp), ctx.mkLt(ctx.mkFPToReal(fp), real));
66+
}
67+
68+
/**
69+
* Safe comparison between a Z3 Real (from int/uint) and a Z3 FloatingPoint (double) for {@code
70+
* <=}.
71+
*/
72+
static BoolExpr mkRealLeFp(Context ctx, RealExpr real, FPExpr fp) {
73+
return mkSafeFpComparison(ctx, fp, isPosInf(ctx, fp), ctx.mkLe(real, ctx.mkFPToReal(fp)));
74+
}
75+
76+
/**
77+
* Safe comparison between a Z3 FloatingPoint (double) and a Z3 Real (from int/uint) for {@code
78+
* <=}.
79+
*/
80+
static BoolExpr mkFpLeReal(Context ctx, FPExpr fp, RealExpr real) {
81+
return mkSafeFpComparison(ctx, fp, isNegInf(ctx, fp), ctx.mkLe(ctx.mkFPToReal(fp), real));
82+
}
83+
84+
private static BoolExpr isPosInf(Context ctx, FPExpr fp) {
85+
return ctx.mkAnd(ctx.mkFPIsInfinite(fp), ctx.mkFPIsPositive(fp));
86+
}
87+
88+
private static BoolExpr isNegInf(Context ctx, FPExpr fp) {
89+
return ctx.mkAnd(ctx.mkFPIsInfinite(fp), ctx.mkFPIsNegative(fp));
90+
}
91+
92+
private static BoolExpr mkSafeFpComparison(
93+
Context ctx, FPExpr fp, BoolExpr infCondition, BoolExpr finiteComparison) {
94+
BoolExpr isFinite = ctx.mkNot(ctx.mkOr(ctx.mkFPIsNaN(fp), ctx.mkFPIsInfinite(fp)));
95+
return ctx.mkOr(infCondition, ctx.mkAnd(isFinite, finiteComparison));
96+
}
97+
5098
private AxiomHelpers() {}
5199
}

verifier/src/main/java/dev/cel/verifier/axioms/GreaterAxiom.java

Lines changed: 16 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -88,33 +88,37 @@ final class GreaterAxiom {
8888
(ctx, typeSystem, constraintSink, lhs, rhs) ->
8989
Optional.of(
9090
typeSystem.wrapBool(
91-
ctx.mkGt(
92-
ctx.mkInt2Real(typeSystem.getInt(lhs)),
93-
ctx.mkFPToReal((FPExpr) typeSystem.getDouble(rhs))))))
91+
AxiomHelpers.mkFpLtReal(
92+
ctx,
93+
(FPExpr) typeSystem.getDouble(rhs),
94+
ctx.mkInt2Real(typeSystem.getInt(lhs))))))
9495
.addBinaryOverloadTranslator(
9596
Comparison.GREATER_UINT64_DOUBLE.celOverloadDecl(),
9697
(ctx, typeSystem, constraintSink, lhs, rhs) ->
9798
Optional.of(
9899
typeSystem.wrapBool(
99-
ctx.mkGt(
100-
ctx.mkInt2Real(typeSystem.getUint(lhs)),
101-
ctx.mkFPToReal((FPExpr) typeSystem.getDouble(rhs))))))
100+
AxiomHelpers.mkFpLtReal(
101+
ctx,
102+
(FPExpr) typeSystem.getDouble(rhs),
103+
ctx.mkInt2Real(typeSystem.getUint(lhs))))))
102104
.addBinaryOverloadTranslator(
103105
Comparison.GREATER_DOUBLE_INT64.celOverloadDecl(),
104106
(ctx, typeSystem, constraintSink, lhs, rhs) ->
105107
Optional.of(
106108
typeSystem.wrapBool(
107-
ctx.mkGt(
108-
ctx.mkFPToReal((FPExpr) typeSystem.getDouble(lhs)),
109-
ctx.mkInt2Real(typeSystem.getInt(rhs))))))
109+
AxiomHelpers.mkRealLtFp(
110+
ctx,
111+
ctx.mkInt2Real(typeSystem.getInt(rhs)),
112+
(FPExpr) typeSystem.getDouble(lhs)))))
110113
.addBinaryOverloadTranslator(
111114
Comparison.GREATER_DOUBLE_UINT64.celOverloadDecl(),
112115
(ctx, typeSystem, constraintSink, lhs, rhs) ->
113116
Optional.of(
114117
typeSystem.wrapBool(
115-
ctx.mkGt(
116-
ctx.mkFPToReal((FPExpr) typeSystem.getDouble(lhs)),
117-
ctx.mkInt2Real(typeSystem.getUint(rhs))))))
118+
AxiomHelpers.mkRealLtFp(
119+
ctx,
120+
ctx.mkInt2Real(typeSystem.getUint(rhs)),
121+
(FPExpr) typeSystem.getDouble(lhs)))))
118122
.addBinaryOverloadTranslator(
119123
Comparison.GREATER_INT64_UINT64.celOverloadDecl(),
120124
(ctx, typeSystem, constraintSink, lhs, rhs) ->

verifier/src/main/java/dev/cel/verifier/axioms/GreaterEqualsAxiom.java

Lines changed: 16 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -88,33 +88,37 @@ final class GreaterEqualsAxiom {
8888
(ctx, typeSystem, constraintSink, lhs, rhs) ->
8989
Optional.of(
9090
typeSystem.wrapBool(
91-
ctx.mkGe(
92-
ctx.mkInt2Real(typeSystem.getInt(lhs)),
93-
ctx.mkFPToReal((FPExpr) typeSystem.getDouble(rhs))))))
91+
AxiomHelpers.mkFpLeReal(
92+
ctx,
93+
(FPExpr) typeSystem.getDouble(rhs),
94+
ctx.mkInt2Real(typeSystem.getInt(lhs))))))
9495
.addBinaryOverloadTranslator(
9596
Comparison.GREATER_EQUALS_UINT64_DOUBLE.celOverloadDecl(),
9697
(ctx, typeSystem, constraintSink, lhs, rhs) ->
9798
Optional.of(
9899
typeSystem.wrapBool(
99-
ctx.mkGe(
100-
ctx.mkInt2Real(typeSystem.getUint(lhs)),
101-
ctx.mkFPToReal((FPExpr) typeSystem.getDouble(rhs))))))
100+
AxiomHelpers.mkFpLeReal(
101+
ctx,
102+
(FPExpr) typeSystem.getDouble(rhs),
103+
ctx.mkInt2Real(typeSystem.getUint(lhs))))))
102104
.addBinaryOverloadTranslator(
103105
Comparison.GREATER_EQUALS_DOUBLE_INT64.celOverloadDecl(),
104106
(ctx, typeSystem, constraintSink, lhs, rhs) ->
105107
Optional.of(
106108
typeSystem.wrapBool(
107-
ctx.mkGe(
108-
ctx.mkFPToReal((FPExpr) typeSystem.getDouble(lhs)),
109-
ctx.mkInt2Real(typeSystem.getInt(rhs))))))
109+
AxiomHelpers.mkRealLeFp(
110+
ctx,
111+
ctx.mkInt2Real(typeSystem.getInt(rhs)),
112+
(FPExpr) typeSystem.getDouble(lhs)))))
110113
.addBinaryOverloadTranslator(
111114
Comparison.GREATER_EQUALS_DOUBLE_UINT64.celOverloadDecl(),
112115
(ctx, typeSystem, constraintSink, lhs, rhs) ->
113116
Optional.of(
114117
typeSystem.wrapBool(
115-
ctx.mkGe(
116-
ctx.mkFPToReal((FPExpr) typeSystem.getDouble(lhs)),
117-
ctx.mkInt2Real(typeSystem.getUint(rhs))))))
118+
AxiomHelpers.mkRealLeFp(
119+
ctx,
120+
ctx.mkInt2Real(typeSystem.getUint(rhs)),
121+
(FPExpr) typeSystem.getDouble(lhs)))))
118122
.addBinaryOverloadTranslator(
119123
Comparison.GREATER_EQUALS_INT64_UINT64.celOverloadDecl(),
120124
(ctx, typeSystem, constraintSink, lhs, rhs) ->

0 commit comments

Comments
 (0)